bert_simpson / forgebench /code /eval /rot_table.py
Ronaldo-GOAT's picture
forgebench: final scorer tree (shard support), aggregate_results.py, latest Omni VGGT drivers (v2/v3), README updates; no seeds
a3752bd verified
Raw History Blame Contribute Delete
5.33 kB
"""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()