File size: 5,325 Bytes
a3752bd
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
"""Dump the FULL per-rotation candidate table for every (object,view,model).

For all 24 cube rotations (x each plane-strip variant) records
  extent_err, raw candidate score f1@2 vs gt_vis, post-ICP score,
  and the achieved final quality (F1_gt block + visible F1@1) for both the
  R-only pose and the R+T_icp pose.

This makes any candidate-selection rule evaluable offline, without re-running
the search: apply the rule to the table, look up what it would have produced.

Usage: python rot_table.py --root .../cat --selection ... --out table.json
"""

from __future__ import annotations

import argparse
import json
import os
from concurrent.futures import ProcessPoolExecutor
from pathlib import Path

import numpy as np
import trimesh
from scipy import ndimage

from faithfulness import canonicalize, voxelize_points
from evaluate_mv2 import load_gt_mv
from align import (octahedral_rotations, strip_support_plane, _f1_at,
                   _mesh_extents, _extent_mismatch, _voxelize_canon_mesh)
from evaluate import load_pred_mesh
from ssinp_eval2 import gt_voxel_block

ROTS = octahedral_rotations()
NS = 400_000


def work(args):
    root, s, views, models = args
    root = Path(root)
    out = []
    for view in views:
        cams = [s["cam"]] if view == "oneview" else [s["cam"], s["cam2"]]
        gt_vis, gt_full, free, mesh_c = load_gt_mv(s["clip"], s, cams, n=64)
        gt_dil = ndimage.maximum_filter(gt_full, size=3)
        ext_gt = _mesh_extents(mesh_c)
        tgt = mesh_c.sample(8000)

        def final(pts):
            occ = voxelize_points(pts, 64)
            b = gt_voxel_block(occ, gt_full, gt_dil)
            return round(b["f1_gt"], 4), round(_f1_at(gt_vis, occ, 1), 4)

        for model in models:
            d = root / f"{view}_gen" / model
            glb = d / f"{s['object']}.glb"
            if not glb.exists():
                continue
            row = {"object": s["object"], "view": view, "model": model,
                   "ext_gt": [round(float(x), 4) for x in ext_gt],
                   "cands": []}
            try:
                raw = load_pred_mesh(glb)
                st = strip_support_plane(raw)
                variants = [(0, raw)] + ([(1, st)] if st is not raw else [])
                for vtag, mv in variants:
                    ext_pr = _mesh_extents(mv)
                    base = mv.copy()
                    base.vertices = base.vertices - (base.vertices.min(0)
                                                     + base.vertices.max(0)) / 2
                    for i, R in enumerate(ROTS):
                        m = base.copy()
                        m.vertices = m.vertices @ R.T
                        f1r = _f1_at(gt_ref := gt_vis,
                                     _voxelize_canon_mesh(m, 64, 50_000))
                        v, _, _ = canonicalize(m.vertices)
                        m.vertices = v
                        pr = np.asarray(m.sample(NS))
                        fin_r = final(pr)
                        src, _, _ = canonicalize(m.sample(8000))
                        try:
                            T, _, _ = trimesh.registration.icp(
                                src, tgt, max_iterations=50)
                            f1i = _f1_at(gt_ref, voxelize_points(
                                (T @ np.c_[src, np.ones(len(src))].T).T[:, :3],
                                64))
                            fin_i = final((T @ np.c_[pr, np.ones(len(pr))].T
                                           ).T[:, :3])
                            det = float(np.linalg.det(T[:3, :3]))
                        except Exception:
                            f1i, fin_i, det = 0.0, (0.0, 0.0), 0.0
                        row["cands"].append(
                            {"v": vtag, "rot": i,
                             "err": round(_extent_mismatch(R, ext_pr, ext_gt), 4),
                             "f1_raw": round(f1r, 4), "f1_icp": round(f1i, 4),
                             "det_icp": round(det, 4),
                             "fin_raw": fin_r, "fin_icp": fin_i})
            except Exception as e:
                row["error"] = repr(e)
            out.append(row)
    return out


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("--root", type=Path, required=True)
    ap.add_argument("--selection", type=Path, required=True)
    ap.add_argument("--views", nargs="+", default=["oneview", "twoview"])
    ap.add_argument("--models", nargs="+",
                    default=["trellis", "reconviagen", "sam3d_gtdepth"])
    ap.add_argument("--objects", nargs="+")
    ap.add_argument("--out", type=Path, required=True)
    ap.add_argument("--jobs", type=int, default=48)
    args = ap.parse_args()

    sel = json.loads(args.selection.read_text())["selections"]
    if args.objects:
        sel = [s for s in sel if s["object"] in set(args.objects)]
    tasks = [(str(args.root), s, args.views, args.models) for s in sel]
    rows = []
    with ProcessPoolExecutor(args.jobs) as ex:
        for i, r in enumerate(ex.map(work, tasks)):
            rows += r
            print(f"[{i+1}/{len(tasks)}] {r[0]['object']}", flush=True)
    args.out.write_text(json.dumps(rows))
    print("wrote", args.out, len(rows))


if __name__ == "__main__":
    os.environ.setdefault("OMP_NUM_THREADS", "2")
    main()