File size: 5,325 Bytes
a3752bd | 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 | """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()
|