File size: 9,442 Bytes
4c3d957
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
#!/usr/bin/env python
"""Align v2: register a canonical-frame baseline mesh to the GT canonical frame.

Fixes over align.py / align_baselines.py (kept untouched for reproducibility):
  1. SCALE: the prediction is normalised to the GT's own max bbox extent
     (FORGE3DBench GT spans 0.9, Toys4K/Omni 1.0). align.py always used 1.0,
     which made every FORGE3DBench baseline 11% too large.
  2. DENSE COARSE SEARCH: 24 cube rotations + N uniformly random rotations
     (default 2000), each scored by voxel F1@2 on a 64^3 grid (same F1 as
     align.py). align.py only tried the 24 cube rotations, which misses
     off-axis yaws.
  3. RIGID ICP (no scale, no reflection) on the top-k rotations, picked with
     non-max suppression so the k seeds are >= 20 deg apart.
  4. INPUT-VIEW SILHOUETTE TIE-BREAK: among candidates whose post-ICP F1 is
     within `tie_eps` of the best, pick the one whose projection through the
     known input camera best matches the GT projection (IoU). Resolves
     symmetric / box-like ambiguities using the camera the input came from.

For methods whose output is already pixel-aligned (Pixal3D, Cupid) do NOT use
this file: they are placed with the input camera instead.

Usage (texture preserved; skip-existing):
  python align_v2.py --pred-dir D --gt-dir EXP/renders --out D_alignv2 \
      --selection EXP/selection.json [--shard i --nshards n]
Self-test: python align_v2.py --selftest EXP
"""
from __future__ import annotations

import argparse
import json
import sys
from pathlib import Path

import numpy as np
import trimesh
from scipy import ndimage
from scipy.spatial.transform import Rotation

sys.path.insert(0, str(Path(__file__).resolve().parent))
from faithfulness import voxelize_points  # noqa: E402
from align import octahedral_rotations, _f1_at  # noqa: E402


def _bbox(v):
    lo, hi = v.min(0), v.max(0)
    return (lo + hi) / 2.0, float((hi - lo).max())


def candidate_rotations(n_rand: int, seed: int = 0):
    rots = octahedral_rotations()
    if n_rand:
        rots += list(Rotation.random(n_rand, random_state=seed).as_matrix())
    return rots


class Similarity:
    """x -> T @ [ ((x - c0) @ R.T - c1) * s ; 1 ]   (T = rigid ICP, 4x4)."""

    def __init__(self, c0, R, c1, s, T=None):
        self.c0, self.R, self.c1, self.s = c0, R, c1, s
        self.T = np.eye(4) if T is None else T

    def pre(self, x):
        return ((x - self.c0) @ self.R.T - self.c1) * self.s

    def __call__(self, x):
        y = self.pre(x)
        return (self.T @ np.c_[y, np.ones(len(y))].T).T[:, :3]

    def to_json(self):
        return {k: np.asarray(getattr(self, k)).tolist() for k in ("c0", "R", "c1", "s", "T")}


def _normalise(verts, R, target_ext):
    c0, _ = _bbox(verts)
    rv = (verts - c0) @ R.T
    c1, ext = _bbox(rv)
    return Similarity(c0, R, c1, target_ext / max(ext, 1e-9))


def _silhouette(pts, cam, res=128):
    """Point-splat silhouette of pts through the input camera (OpenCV)."""
    if cam is None:
        return None
    w2c = np.linalg.inv(cam["c2w"])
    pc = (w2c[:3, :3] @ pts.T).T + w2c[:3, 3]
    z = pc[:, 2]
    ok = z > 1e-6
    sc = res / cam["res"]
    u = (cam["fx"] * pc[ok, 0] / z[ok] + cam["cx"]) * sc
    v = (cam["fy"] * pc[ok, 1] / z[ok] + cam["cy"]) * sc
    m = np.zeros((res, res), bool)
    k = (u >= 0) & (u < res) & (v >= 0) & (v < res)
    m[v[k].astype(int), u[k].astype(int)] = True
    m = ndimage.binary_closing(m, iterations=2)
    return ndimage.binary_fill_holes(m)


def _iou(a, b):
    if a is None or b is None:
        return 0.0
    u = (a | b).sum()
    return float((a & b).sum() / u) if u else 0.0


def load_cam(npz_path: Path):
    if not npz_path.exists():
        return None
    z = np.load(npz_path)
    res = float(z["res"]) if "res" in z.files else 512.0
    return dict(fx=float(z["fx"]), fy=float(z["fy"]), cx=float(z["cx"]),
                cy=float(z["cy"]), c2w=np.asarray(z["c2w_cv"], float), res=res)


def align_v2(pred: trimesh.Trimesh, gt: trimesh.Trimesh, cam=None, n=64,
             n_rand=2000, topk=8, nms_deg=20.0, tie_eps=0.03, seed=0):
    gt_ext = _bbox(np.asarray(gt.vertices))[1]
    gt_occ = voxelize_points(np.asarray(gt.sample(200_000)), n)
    gt_pts = np.asarray(gt.sample(8000))
    gt_sil = _silhouette(np.asarray(gt.sample(30_000)), cam)
    V = np.asarray(pred.vertices, float)
    P = np.asarray(pred.sample(20_000), float)
    P_icp, P_sil = P[:8000], P

    rots = candidate_rotations(n_rand, seed)
    scored = []
    for i, R in enumerate(rots):
        S = _normalise(V, R, gt_ext)
        scored.append((_f1_at(gt_occ, voxelize_points(S.pre(P), n)), i, S))
    scored.sort(key=lambda t: -t[0])

    seeds = []
    for f1, i, S in scored:
        if all(np.degrees(Rotation.from_matrix(S.R @ s[2].R.T).magnitude()) >= nms_deg
               for s in seeds):
            seeds.append((f1, i, S))
        if len(seeds) >= topk:
            break

    cands = []
    for f1_raw, i, S in seeds:
        best = (f1_raw, S)
        try:
            T, _, _ = trimesh.registration.icp(S.pre(P_icp), gt_pts, max_iterations=50,
                                               scale=False, reflection=False)
            S2 = Similarity(S.c0, S.R, S.c1, S.s, T)
            f1_icp = _f1_at(gt_occ, voxelize_points(S2(P), n))
            if f1_icp >= f1_raw:
                best = (f1_icp, S2)
        except Exception:
            pass
        f1, Sb = best
        cands.append(dict(rot=i, f1_raw=f1_raw, f1=f1, S=Sb,
                          sil=_iou(_silhouette(Sb(P_sil), cam), gt_sil)))

    top = max(c["f1"] for c in cands)
    pool = [c for c in cands if c["f1"] >= top - tie_eps]
    win = max(pool, key=lambda c: (c["sil"], c["f1"]))
    cube_best = max(f for f, i, _ in scored if i < 24)
    return win, dict(best_f1=round(win["f1"], 4), f1_raw=round(win["f1_raw"], 4),
                     sil_iou=round(win["sil"], 4), rot=win["rot"],
                     cube24_raw=round(cube_best, 4), top_f1=round(top, 4),
                     n_tied=len(pool), gt_ext=round(gt_ext, 4),
                     cands=[{k: round(c[k], 4) if isinstance(c[k], float) else c[k]
                             for k in ("rot", "f1_raw", "f1", "sil")} for c in cands])


def resolve_gt(gt_dir: Path, obj: str):
    for c in (gt_dir / f"{obj}_canon.glb", gt_dir / obj / "mesh.glb", gt_dir / f"{obj}.glb"):
        if c.exists():
            return c
    return None


def geom(m):
    return trimesh.Trimesh(np.asarray(m.vertices), np.asarray(m.faces), process=False)


def selftest(exp: Path, k=12):
    """Random rotation + scale + translation of GT must be undone (F1 ~ 1)."""
    rng = np.random.default_rng(0)
    objs = sorted(p.name[:-len("_canon.glb")] for p in (exp / "renders").glob("*_canon.glb"))
    objs = [objs[j] for j in rng.choice(len(objs), k, replace=False)]
    f1s = []
    for o in objs:
        gt = trimesh.load(str(exp / "renders" / f"{o}_canon.glb"), force="mesh", process=False)
        R = Rotation.random(random_state=int(rng.integers(1e9))).as_matrix()
        m = geom(gt)
        m.vertices = (np.asarray(m.vertices) @ R.T) * rng.uniform(0.3, 3) + rng.normal(size=3)
        win, rep = align_v2(m, gt, load_cam(exp / "renders" / f"{o}_front.npz"))
        f1s.append(rep["best_f1"])
        print(f"selftest {o}: best_f1={rep['best_f1']:.3f} cube24_raw={rep['cube24_raw']:.3f} sil={rep['sil_iou']:.3f}", flush=True)
    print(f"SELFTEST mean F1 {np.mean(f1s):.3f} min {np.min(f1s):.3f}")


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("--pred-dir", type=Path)
    ap.add_argument("--gt-dir", type=Path)
    ap.add_argument("--out", type=Path)
    ap.add_argument("--selection", type=Path)
    ap.add_argument("--n-rand", type=int, default=2000)
    ap.add_argument("--shard", type=int, default=0)
    ap.add_argument("--nshards", type=int, default=1)
    ap.add_argument("--selftest", type=Path)
    a = ap.parse_args()
    if a.selftest:
        return selftest(a.selftest)
    a.out.mkdir(parents=True, exist_ok=True)
    objs = [s["object"] for s in json.loads(a.selection.read_text())["selections"]][a.shard::a.nshards]
    for i, obj in enumerate(objs, 1):
        outp, repp = a.out / f"{obj}.glb", a.out / f"{obj}.align.json"
        if outp.exists() and repp.exists():
            print(f"[{i}/{len(objs)}] {obj} SKIP", flush=True); continue
        pp, gp = a.pred_dir / f"{obj}.glb", resolve_gt(a.gt_dir, obj)
        if not pp.exists() or gp is None:
            print(f"[{i}/{len(objs)}] {obj} MISSING {'pred' if not pp.exists() else 'gt'}", flush=True); continue
        try:
            tex = trimesh.load(str(pp), force="mesh", process=False)
            gt = trimesh.load(str(gp), force="mesh", process=False)
            win, rep = align_v2(geom(tex), gt, load_cam(a.gt_dir / f"{obj}_front.npz"), n_rand=a.n_rand)
            tex.vertices = win["S"](np.asarray(tex.vertices, float))
            tmp = outp.with_suffix(".tmp.glb"); tex.export(str(tmp)); tmp.rename(outp)
            rep["transform"] = win["S"].to_json()
            repp.write_text(json.dumps(rep))
            print(f"[{i}/{len(objs)}] {obj} OK best_f1={rep['best_f1']:.3f} (cube24_raw={rep['cube24_raw']:.3f} "
                  f"sil={rep['sil_iou']:.3f} tied={rep['n_tied']})", flush=True)
        except Exception as e:
            print(f"[{i}/{len(objs)}] {obj} FAIL {type(e).__name__}: {e}", flush=True)


if __name__ == "__main__":
    main()