#!/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()