bert_simpson / forgebench /code /eval /align_v2.py
Ronaldo-GOAT's picture
Add files using upload-large-folder tool
4c3d957 verified
Raw History Blame Contribute Delete
9.44 kB
#!/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()