"""Align a model-output mesh to the GT canonical frame. Search space (deliberately small, per protocol): the 24 cube ("side") orientations, pre-filtered by axis-extent ordering — a candidate survives only if rotating the prediction's bbox extents roughly matches the GT extents (longest axis to longest axis, etc.). Each survivor is scored by voxel F1@2 against the GT surface grid; the best is optionally ICP-refined (trimesh point-to-point) and re-scored. Returns both raw-best and ICP-refined transforms + scores so the report can show both, and saves candidate renders for the visual verifier. Self-test: python metrics/align.py (recovers known rotations of a GLB) """ from __future__ import annotations import itertools import json from dataclasses import dataclass, field import numpy as np import trimesh from faithfulness import canonicalize, voxelize_points def octahedral_rotations() -> list[np.ndarray]: """All 24 rotation matrices of the cube group (det=+1).""" mats = [] for perm in itertools.permutations(range(3)): for signs in itertools.product((1, -1), repeat=3): R = np.zeros((3, 3)) for i, (p, s) in enumerate(zip(perm, signs)): R[i, p] = s if np.isclose(np.linalg.det(R), 1.0): mats.append(R) assert len(mats) == 24 return mats def _extent_mismatch(R: np.ndarray, ext_pred: np.ndarray, ext_gt: np.ndarray) -> float: """Relative mismatch between rotated prediction extents and GT extents.""" rot_ext = np.abs(R) @ ext_pred # cube rotation permutes extents return float(np.max(np.abs(rot_ext - ext_gt) / np.maximum(ext_gt, 1e-9))) def _mesh_extents(mesh: trimesh.Trimesh) -> np.ndarray: return mesh.vertices.max(0) - mesh.vertices.min(0) def _voxelize_canon_mesh(mesh: trimesh.Trimesh, n: int, samples: int) -> np.ndarray: pts = mesh.sample(samples) pts, _, _ = canonicalize(pts) # own-bbox canonicalization return voxelize_points(pts, n) def _f1_at(gt: np.ndarray, pred: np.ndarray, r: int = 2) -> float: from scipy import ndimage gt_d = ndimage.maximum_filter(gt, size=2 * r + 1) pr_d = ndimage.maximum_filter(pred, size=2 * r + 1) prec = float(gt_d[pred].mean()) if pred.any() else 0.0 rec = float(pr_d[gt].mean()) if gt.any() else 0.0 return 2 * prec * rec / (prec + rec) if prec + rec > 0 else 0.0 @dataclass class AlignResult: R_raw: np.ndarray = None # best cube rotation f1_raw: float = 0.0 T_icp: np.ndarray = None # 4x4 refinement AFTER R_raw (canon space) f1_icp: float = 0.0 candidates: list = field(default_factory=list) # (extent_err, f1) per R extent_err: float = 0.0 # extent mismatch of the chosen rotation def best_f1(self) -> float: return max(self.f1_raw, self.f1_icp) def strip_support_plane(mesh: trimesh.Trimesh, angle_deg: float = 15.0, plane_tol: float = 0.02, min_area_frac: float = 0.25, min_shrink: float = 0.30) -> trimesh.Trimesh: """Remove a large flat sheet fused to the object (e.g. a generated ground/support plane), which otherwise corrupts own-bbox canonicalization. A candidate sheet = faces whose normal is within `angle_deg` of one axis and whose plane position clusters (area-weighted histogram peak), with total area >= min_area_frac of the mesh. The strip is ACCEPTED only if removing it shrinks the bbox by >= min_shrink in some axis — a real object face (e.g. a box side) leaves the bbox unchanged and is kept. """ fn = mesh.face_normals area = mesh.area_faces total = float(area.sum()) ext = _mesh_extents(mesh) cos = np.cos(np.deg2rad(angle_deg)) centers = mesh.triangles.mean(axis=1) best = None for ax in range(3): aligned = np.abs(fn[:, ax]) > cos if not aligned.any(): continue c = centers[aligned, ax] hist, edges = np.histogram(c, bins=64, weights=area[aligned]) pos = (edges[hist.argmax()] + edges[hist.argmax() + 1]) / 2 near = aligned & (np.abs(centers[:, ax] - pos) < plane_tol * max(ext[ax], 1e-9)) a = float(area[near].sum()) if a >= min_area_frac * total and (best is None or a > best[0]): best = (a, near, ax) if best is None: return mesh a, near, ax = best m = mesh.submesh([np.nonzero(~near)[0]], append=True) parts = m.split(only_watertight=False) if len(parts) > 1: # cutting the sheet out of an open shell shatters the object into # several parts (e.g. can wall quarters) plus plane residue. The # residue is thin ALONG THE PLANE AXIS; the object parts span it. span = max(float(_mesh_extents(p)[ax]) for p in parts) keep = [p for p in parts if _mesh_extents(p)[ax] >= 0.05 * span] if not keep: return mesh m = trimesh.util.concatenate(keep) if m.is_empty or len(m.vertices) < 16: return mesh shrink = 1.0 - _mesh_extents(m) / np.maximum(ext, 1e-9) if shrink.max() < min_shrink: return mesh # was a real face, keep intact return m def align_mesh(pred_mesh: trimesh.Trimesh, gt_ref: np.ndarray, gt_mesh_canon: trimesh.Trimesh, n: int = 64, samples: int = 200_000, extent_tol: float = 0.35, extent_slack: float = 0.20, icp_raw_margin: float = 0.25, icp: bool = True, icp_topk: int = 8) -> AlignResult: """Find the cube orientation (+ optional ICP) aligning pred to GT canon. pred_mesh : model output, arbitrary canonical pose (its own frame) gt_ref : (n,n,n) bool reference voxels for SCORING candidates. Use gt_vis: the gt_full objective is degenerate for slab/box shapes (all 24 rotations score ~equal and the argmax can be a wrong side); gt_vis is one-sided and discriminates. This is best-case alignment for the faithfulness metric by construction. gt_mesh_canon : GT mesh already in canonical coords (for ICP target) extent_tol : keep cube rotations whose extent mismatch <= tol. extent_slack : a HARD tol prune can delete the only correct rotation: _extent_mismatch is a max over axes, so for objects whose two short GT axes are similar (a van: 0.39 x 1.00 x 0.43) a permutation that is wrong on BOTH short axes by a moderate amount beats the right one that is wrong on a single axis by more. So also keep every rotation within `extent_slack` of the best achievable extent error and let the score decide. icp_raw_margin: the post-ICP score may only move the answer into a basin whose raw score is within this margin of the best one (see the shrink comment below); 0.25 is far outside the <0.10 "flat" regime this stage was introduced for. """ res = AlignResult() ext_gt = _mesh_extents(gt_mesh_canon) ext_pr = _mesh_extents(pred_mesh) rots = octahedral_rotations() errs = [_extent_mismatch(R, ext_pr, ext_gt) for R in rots] thr = max(extent_tol, min(errs) + extent_slack) keep = [i for i, e in enumerate(errs) if e <= thr] or list(range(len(rots))) base = pred_mesh.copy() base.vertices = base.vertices - (base.vertices.min(0) + base.vertices.max(0)) / 2.0 scored = [] for i in keep: R = rots[i] m = base.copy() m.vertices = m.vertices @ R.T occ = _voxelize_canon_mesh(m, n, samples // 4) f1 = _f1_at(gt_ref, occ) scored.append((f1, i)) res.candidates.append({"rot": i, "extent_err": round(errs[i], 3), "f1@2": round(f1, 4)}) if f1 > res.f1_raw: res.f1_raw, res.R_raw = f1, R res.extent_err = errs[i] if res.R_raw is None: return res if icp: # raw per-rotation F1 is nearly flat for slab/box shapes — the raw # argmax can sit in the wrong basin. ICP-refine the top-k rotations # and pick by post-ICP score instead. scored.sort(reverse=True) spread = scored[0][0] - scored[-1][0] k = len(scored) if spread < 0.10 else icp_topk # flat -> try all pool = [c for c in scored[:k] if scored[0][0] - c[0] <= icp_raw_margin] or [scored[0]] tgt = gt_mesh_canon.sample(8000) best_icp_sel = -np.inf icp_win = None # (sel, f1_icp, f1_raw, R, T, err) for f1_r, i in pool: R = rots[i] m = base.copy() m.vertices = m.vertices @ R.T src, _, _ = canonicalize(m.sample(8000)) try: T, _, _ = trimesh.registration.icp(src, tgt, max_iterations=50) except Exception: continue src_h = np.c_[src, np.ones(len(src))] occ = voxelize_points((T @ src_h.T).T[:, :3], n) f1_i = _f1_at(gt_ref, occ) # trimesh's ICP fits SCALE too, and shrinking the prediction # inside the dilated GT shell inflates precision -> the post-ICP # score saturates and goes nearly flat across basins (a van: # 0.79 for the UPSIDE-DOWN rotation vs 0.75 for the upright one, # while the raw scores are 0.44 vs 0.65). Ranking basins on it # alone therefore picks poses the discriminative raw score # clearly rejects. Damp the score by the shrink it needed, so a # basin can only win on genuine agreement, not on shrinking. shrink = min(1.0, abs(np.linalg.det(T[:3, :3])) ** (1 / 3)) sel_i = f1_i * shrink if sel_i > best_icp_sel: best_icp_sel = sel_i icp_win = (sel_i, f1_i, f1_r, R, T, errs[i]) if icp_win is not None: # the raw winner is always inside scored[:k], so the ICP winner # won the ranking against it -> adopt its basin, keeping R and T # consistent (apply_alignment applies T on top of R; the old code # could apply a T fitted in a different basin). sel_i, f1_i, f1_r, R, T, e = icp_win res.f1_icp, res.T_icp = f1_i, T res.R_raw, res.f1_raw, res.extent_err = R, f1_r, e return res def apply_alignment(pred_mesh: trimesh.Trimesh, res: AlignResult, use_icp: bool = True) -> trimesh.Trimesh: """Return pred_mesh mapped into GT canonical coords by the found align.""" m = pred_mesh.copy() m.vertices = m.vertices - (m.vertices.min(0) + m.vertices.max(0)) / 2.0 m.vertices = m.vertices @ res.R_raw.T v, _, _ = canonicalize(m.vertices) m.vertices = v if use_icp and res.T_icp is not None and res.f1_icp >= res.f1_raw: vh = np.c_[m.vertices, np.ones(len(m.vertices))] m.vertices = (res.T_icp @ vh.T).T[:, :3] return m # ---------------------------------------------------------------- self-test def _self_test(): from pathlib import Path from gt_loader import load_gt, _load_json sel = _load_json(Path("../.debug/exp_faithfulness/selection.json")) s = next(x for x in sel["selections"] if x["object"] == "wooden_foo_dog") g = load_gt(sel["clip"], s) gt_occ = voxelize_points(canonicalize(g.mesh_canon.sample(200_000), np.zeros(3), 1.0)[0], 64) rng = np.random.default_rng(0) rots = octahedral_rotations() ok = True for trial, R_true in enumerate([rots[7], rots[15]]): m = g.mesh_canon.copy() m.vertices = m.vertices @ R_true.T res = align_mesh(m, gt_occ, g.mesh_canon) rec = np.allclose(res.R_raw @ R_true, np.eye(3)) print(f"cube rot {trial}: f1_raw={res.f1_raw:.3f} " f"f1_icp={res.f1_icp:.3f} exact_inverse={rec} " f"n_cand={len(res.candidates)}") ok &= res.f1_raw > 0.98 # non-cube perturbation: 10 deg yaw on top of a cube rot -> ICP recovers a = np.deg2rad(10) Rz = np.array([[np.cos(a), -np.sin(a), 0], [np.sin(a), np.cos(a), 0], [0, 0, 1]]) m = g.mesh_canon.copy() m.vertices = m.vertices @ (Rz @ rots[7]).T res = align_mesh(m, gt_occ, g.mesh_canon) print(f"cube+10deg: f1_raw={res.f1_raw:.3f} f1_icp={res.f1_icp:.3f}") ok &= res.f1_icp > res.f1_raw and res.f1_icp > 0.95 assert ok, "align self-test failed" print("ALL ALIGN SELF-TESTS PASSED") if __name__ == "__main__": _self_test()