bert_simpson / forgebench /code /ours /batch_appforce_sam3d.py
Ronaldo-GOAT's picture
Add files using upload-large-folder tool
4c3d957 verified
Raw History Blame Contribute Delete
53.9 kB
#!/usr/bin/env python
"""SAM3D APPEARANCE-FORCING generation driver (the core novel method).
Produces a textured mesh for each object with TWO methods (selected by --method),
sharing identical code so the paired (4) vs (6) comparison is apples-to-apples:
(4) sam3d_geom STAGE-1 geometry forcing (z1fwd) -> coords;
STAGE-2 = the model's DEFAULT image-conditioned flow from
pure noise. No appearance forcing.
(6) sam3d_geom_dino same STAGE-1; STAGE-2 = APPEARANCE FORCING: the visible
voxels' stage-2 initial-noise rows are SEEDED with the
SLAT-encoded real DINOv2 features, then the flow runs
normally (SEED-ONLY, no re-injection). Invisible rows
stay noise.
Both stages use the PRODUCTION SAM3D sampler (pipe.sample_sparse_structure /
pipe.sample_slat), so schedule / CFG / prune+downsample are exactly the model's
own (ss_rescale_t=3, ss_cfg_strength=7, slat_rescale_t=3 -- inference_pipeline.py
:81-115). We only intercept the generators' `_generate_noise` to seed voxels --
the same monkeypatch pattern the repo itself uses in
sample_slat_multi_view_weighted (inference_pipeline.py:1493-1512).
WHY seed-only via _generate_noise works: both generators
(ShortCut for stage-1, FlowMatching for stage-2) draw x_0 EXACTLY ONCE at the top
of generate_iter and never re-inject (shortcut/model.py:generate_iter,
flow_matching/model.py:generate_iter). So replacing x_0's visible rows and then
integrating the flow forward (t:0->1) is precisely "seed-only forcing".
STAGE-1 z1fwd (== port_m5_sam3d.py --sub-scale 0, the validated ablation): seed
the SS shape-stream's initial latent with z1 = ss_encoder(visible-occupancy grid)
and run the SS flow forward with no velocity subtraction. Pose streams stay noise
(they are inert for the shape stream: protect_modality_list=["shape"]).
STAGE-2 appearance forcing recipe (verbatim from the PASSED round-trip
scripts/appforce_roundtrip_sam3d.py):
* voxel set = the FINAL coords handed to sample_slat (after prune+downsample)
* visibility from GT-render depth (project centers with inv(c2w_cv), OpenCV,
tol 0.02), union over the input views
* DINOv2 dinov2_vitl14_reg, crop->518, premult-alpha on black, ImageNet-norm,
feats = x_prenorm[:, num_register_tokens+1:] -> (1024,37,37) RAW, bilinear
grid_sample at each voxel's projected pixel IN THE CROP FRAME, averaged over
views where visible
* ENCODER INPUT = VISIBLE-ONLY subset (the round-trip winner), scatter the
encoded latent back into the full-coord noise tensor by coordinate
* NORMALIZE with the pipeline's slat_mean/std (pipeline.yaml values, i.e.
pipe.slat_mean / pipe.slat_std -- NOT the dead inference_utils SLAT_MEAN
constants) because the flow lives in normalized latent space (sample_slat
de-normalizes with `slat * slat_std + slat_mean` afterwards).
EXPORT (per the env constraint verified by the round-trip): to_glb texture baking
is UNAVAILABLE in this env (utils3d==1.7 dropped the rasterizer; inria gsplat not
installed). So we export the MESH DECODER's NATIVE per-vertex colors
(vertex_attrs[:, :3]) as a vertex-colored GLB, in the canonical [-0.5,0.5] frame
with NO rotation (the round-trip confirmed the raw decoder frame aligns with the
input c2w_cv). Mesh is decoded from a FLOAT32 latent outside autocast (flexicubes
index_add_ needs fp32). Both methods use this identical representation -> the
paired comparison is fair.
CLI (matches the existing batch drivers)
----------------------------------------
python batch_appforce_sam3d.py --method {sam3d_geom|sam3d_geom_dino}
--selection SEL.json --inputs INPUTDIR --exp EXPDIR --out OUTDIR
--views {1|2} [--schedule {seed|reinject}] [--seed 42] [--limit N]
[--gpu 4] [--shard i --nshards n]
--schedule (sam3d_geom_dino only): seed = SEED-ONLY (default, unchanged); reinject
= RePaint HARD re-injection that PINS the visible rows to the on-path value of the
encoded observation at every Euler step -- for the OOD case where the default
hallucinates and the encode ceiling exceeds it. Interpolant (SAM3D t:0->1):
x_t = (1-(1-sigma_min)t)*eps + t*z ; after each step (state at t_next):
x[vis] = (1-(1-sigma_min)*t_next)*eps_fixed + t_next*z_forced (eps_fixed drawn once)
SAM3D's SLAT stage is a dense (bs, N_coords, 8) tensor over a SINGLE shared
coords array (bs is multi-VIEW of one object), so cross-object batching is not
possible without rewriting the generator. Throughput = CO-LOCATED per-object
processes: launch n shards (--shard i --nshards n), e.g. 2 per GPU on GPUs 4/5.
EXPDIR must contain renders/<obj>_{front,side}.npz, renders/<obj>_canon.glb,
inputs/, npz_1v/<obj>/da3_output.npz (1-view) and npz_2v/... (2-view).
Output: OUTDIR/<obj>.glb (vertex-colored mesh). Idempotent (skips existing),
per-object try/except, atomic writes, per-object timing.
Re-execs itself into /lp-dev/jonghoon/mv-mesh/envs/mv-sam3d.
"""
import os
import sys
REPO = "/lp-dev/jonghoon/mv-mesh/mv-sam3d"
ENV = "/lp-dev/jonghoon/mv-mesh/envs/mv-sam3d"
ENV_PY = os.path.join(ENV, "bin", "python")
HF = "/lp-dev/jonghoon/mv-mesh/hf_cache"
METRICS_DIR = os.path.dirname(os.path.abspath(__file__))
_ENV_VARS = {
"PYTHONUNBUFFERED": "1",
"OMP_NUM_THREADS": "4",
"MKL_NUM_THREADS": "4",
"CONDA_PREFIX": ENV, # notebook/inference.py does CUDA_HOME=CONDA_PREFIX
"HF_HOME": HF,
"HUGGINGFACE_HUB_CACHE": HF,
"HF_HUB_CACHE": HF,
"TORCH_HOME": "/lp-dev/jonghoon/mv-mesh/torch_hub",
"PYOPENGL_PLATFORM": "egl",
}
def _parse_gpu_from_argv():
for i, a in enumerate(sys.argv):
if a == "--gpu" and i + 1 < len(sys.argv):
return sys.argv[i + 1]
if a.startswith("--gpu="):
return a.split("=", 1)[1]
return "4"
def _reexec_in_env():
env = dict(os.environ)
for k, v in _ENV_VARS.items():
env[k] = v
env["PATH"] = os.path.join(ENV, "bin") + os.pathsep + env.get("PATH", "")
if not env.get("CUDA_VISIBLE_DEVICES"):
env["CUDA_VISIBLE_DEVICES"] = _parse_gpu_from_argv()
env["_APPFORCE_SAM3D_INENV"] = "1"
print(f"[af] re-exec in {ENV_PY} (CUDA={env.get('CUDA_VISIBLE_DEVICES')})",
flush=True)
os.execve(ENV_PY, [ENV_PY, os.path.abspath(__file__)] + sys.argv[1:], env)
if not os.environ.get("_APPFORCE_SAM3D_INENV"):
if not os.environ.get("CUDA_VISIBLE_DEVICES"):
os.environ["CUDA_VISIBLE_DEVICES"] = _parse_gpu_from_argv()
_reexec_in_env()
for _k, _v in _ENV_VARS.items():
os.environ.setdefault(_k, _v)
import argparse # noqa: E402
import glob # noqa: E402
import json # noqa: E402
import time # noqa: E402
import traceback # noqa: E402
from concurrent.futures import ThreadPoolExecutor # noqa: E402
from pathlib import Path # noqa: E402
os.chdir(REPO)
sys.path.insert(0, REPO)
import numpy as np # noqa: E402
import torch # noqa: E402
import torch.nn.functional as F # noqa: E402
from PIL import Image # noqa: E402
import trimesh # noqa: E402
from scipy import ndimage # noqa: E402
# light metrics helpers (no trellis / no heavy deps)
sys.path.insert(0, METRICS_DIR)
from faithfulness import voxelize_points # noqa: E402
from gt_loader import carve_free_space_depth # noqa: E402
torch.set_grad_enabled(False)
DEVICE = "cuda"
N_PATCH = 518 // 14 # 37
VOX = 64
TOL = 0.02 # visibility depth tolerance (~1.3 voxel widths)
DTYPE = torch.float16
SLAT_ENC_YAML = glob.glob(
f"{HF}/models--facebook--sam-3d-objects/snapshots/*/checkpoints/slat_encoder.yaml")
SLAT_ENC_CKPT = glob.glob(
f"{HF}/models--facebook--sam-3d-objects/snapshots/*/checkpoints/slat_encoder.ckpt")
# --------------------------------------------------------------------------- #
# GT / occupancy helpers (copied verbatim from evaluate_synth.load_gt_synth so
# we do not import that module's heavy trellis-adjacent deps).
# --------------------------------------------------------------------------- #
def load_gt_synth(exp: Path, obj: str, views, n: int = 64,
mesh_samples: int = 1_000_000, mask_erode: int = 2):
"""-> (gt_vis, gt_full, free) on the shared canonical 64^3 grid."""
mesh_c = trimesh.load(exp / "renders" / f"{obj}_canon.glb", force="mesh")
surf, _ = trimesh.sample.sample_surface(mesh_c, mesh_samples, seed=0)
gt_full = voxelize_points(np.asarray(surf), n)
gt_vis = np.zeros((n, n, n), bool)
free = np.zeros((n, n, n), bool)
eye4 = np.eye(4)
zero3 = np.zeros(3)
for view in views:
z = np.load(exp / "renders" / f"{obj}_{view}.npz")
d = z["depth_mm"]
K = {k: float(z[k]) for k in ("fx", "fy", "cx", "cy")}
c2w = z["c2w_cv"]
px = d > 0
if mask_erode:
px &= ndimage.binary_erosion(px, iterations=mask_erode)
ys, xs = np.nonzero(px)
zz = d[ys, xs].astype(np.float64) / 1000.0
pc = np.stack([(xs - K["cx"]) / K["fx"] * zz,
(ys - K["cy"]) / K["fy"] * zz, zz,
np.ones_like(zz)], 1)
pts = (c2w @ pc.T).T[:, :3]
near = (np.abs(pts) <= 0.5 + 0.012).all(1)
pts = np.clip(pts[near], -0.5, 0.5 - 1e-9)
vis_i = voxelize_points(pts, n)
gt_vis |= vis_i
free |= carve_free_space_depth(d, K, c2w, eye4, zero3, 1.0, vis_i, n)
free &= ~ndimage.maximum_filter(gt_vis, size=3)
# --- OPT-IN OCCUPANCY OVERRIDE (adapter-supplied, gated; method unchanged) ---
# If the HO3D adapter requested it (AF_OCC=hull) and dropped a precomputed
# occupancy grid at exp/hull_occ.npz, seed z1 from THAT (a space-carved
# visual hull built in build_recon_ours.build_hull_occ) instead of the depth
# union above. Pure data-load: no method/sampler logic changes, and the
# default path (flag unset / file absent, e.g. toys4k) is byte-identical.
if os.environ.get("AF_OCC") == "hull":
hp = Path(exp) / "hull_occ.npz"
if hp.is_file():
hull = np.load(hp)["occ"].astype(bool)
if hull.shape == gt_vis.shape and hull.any():
gt_vis = hull
free = np.zeros_like(free) # hull is a solid; no depth carving
print(f"[af] OCC OVERRIDE: hull_occ.npz nvox={int(hull.sum())}",
flush=True)
return gt_vis, gt_full, free
def build_input_grid(gt_vis, free, noise, rng):
"""visible -> 1, known-empty -> 0, unknown -> Bernoulli(0.5) or 0.
(copied from ssinp_infer.build_input_grid). z1fwd uses noise='zeros'."""
occ = np.zeros_like(gt_vis, dtype=np.float32)
occ[gt_vis] = 1.0
unknown = ~gt_vis & ~free
if noise == "bernoulli":
occ[unknown] = (rng.random(int(unknown.sum())) < 0.5).astype(np.float32)
elif noise == "zeros":
pass
else:
raise ValueError(noise)
return occ
# --------------------------------------------------------------------------- #
# camera / visibility / DINO helpers (verbatim recipe from the PASSED round-trip)
# --------------------------------------------------------------------------- #
def load_view(exp, obj, view):
z = np.load(f"{exp}/renders/{obj}_{view}.npz")
return dict(depth_mm=z['depth_mm'].astype(np.float64) / 1000.0,
fx=float(z['fx']), fy=float(z['fy']), cx=float(z['cx']), cy=float(z['cy']),
c2w=z['c2w_cv'].astype(np.float64), bbox=z['bbox'].astype(int),
res=int(z['res']))
def project_visible(centers, v):
"""world voxel centers -> (u,vv,z,visible) in the view's OpenCV camera."""
c2w = v['c2w']; w2c = np.linalg.inv(c2w)
R, t = w2c[:3, :3], w2c[:3, 3]
xc = centers @ R.T + t
z = xc[:, 2]
u = v['fx'] * xc[:, 0] / z + v['cx']
vv = v['fy'] * xc[:, 1] / z + v['cy']
res = v['res']
ui = np.round(u).astype(int); vi = np.round(vv).astype(int)
inframe = (z > 0) & (ui >= 0) & (ui < res) & (vi >= 0) & (vi < res)
dep = np.zeros(len(centers))
dep[inframe] = v['depth_mm'][vi[inframe], ui[inframe]]
visible = inframe & (dep > 0) & (np.abs(z - dep) < TOL)
return u, vv, z, visible
def dino_features(dino, dino_norm, png_path):
"""(1,1024,37,37) RAW x_prenorm patch tokens from the input crop."""
im = Image.open(png_path).convert('RGBA').resize((518, 518), Image.Resampling.LANCZOS)
a = np.array(im).astype(np.float32) / 255.0
_bg = float(os.environ.get('AF_BG', '0.0')) # E1 ablation: composite bg (0=black=orig)
_al = a[:, :, 3:4]
rgb = a[:, :, :3] * _al + _bg * (1.0 - _al) # premult alpha on _bg (default black)
x = torch.from_numpy(rgb).permute(2, 0, 1).float()
x = dino_norm(x).unsqueeze(0).cuda()
feats = dino(x, is_training=True)
pt = feats['x_prenorm'][:, dino.num_register_tokens + 1:]
return pt.permute(0, 2, 1).reshape(1, 1024, N_PATCH, N_PATCH)
def sample_feats(patchtokens, u_full, vv_full, v):
"""bilinear-sample DINO features at voxel projections, in the CROP frame."""
y0, y1, x0, x1 = v['bbox']
side = x1 - x0
un = (u_full - x0 + 0.5) / side * 2 - 1
vn = (vv_full - y0 + 0.5) / side * 2 - 1
uv = torch.from_numpy(np.stack([un, vn], -1)).float().cuda().view(1, -1, 1, 2)
f = F.grid_sample(patchtokens, uv, mode='bilinear', align_corners=False)
return f.squeeze(-1).squeeze(0).permute(1, 0) # (N,1024)
# --------------------------------------------------------------------------- #
# model loading
# --------------------------------------------------------------------------- #
def load_pretrained(yaml_path, ckpt_path):
"""InferencePipeline.instantiate_and_load_from_pretrained, standalone."""
from hydra.utils import instantiate
from omegaconf import OmegaConf
from sam3d_objects.model.io import load_model_from_checkpoint
cfg = OmegaConf.load(yaml_path)
if "pretrained_ckpt_path" in cfg:
del cfg["pretrained_ckpt_path"]
model = instantiate(cfg)
model = load_model_from_checkpoint(
model, ckpt_path, strict=True, device="cpu", freeze=True, eval=True,
state_dict_key=None)
return model.to(DEVICE).eval()
def find_ss_encoder():
cands = sorted(glob.glob(
f"{HF}/models--facebook--sam-3d-objects/snapshots/*/checkpoints/ss_encoder.ckpt"))
cands += ["/data/nvidia/gripper_augmentator/checkpoints/sam3d/hf/ss_encoder.ckpt"]
for c in cands:
if os.path.isfile(c) and os.path.getsize(c) > 1_000_000:
return c
raise FileNotFoundError("ss_encoder.ckpt not found")
def load_ss_encoder():
from sam3d_objects.model.backbone.tdfy_dit.models.sparse_structure_vae import (
SparseStructureEncoder)
ep = find_ss_encoder()
print(f"[af] ss_encoder <- {ep} ({os.path.getsize(ep)} bytes)", flush=True)
esd = torch.load(ep, map_location="cpu")
if isinstance(esd, dict) and "state_dict" in esd:
esd = esd["state_dict"]
enc = SparseStructureEncoder(
in_channels=1, latent_channels=8, num_res_blocks=2,
num_res_blocks_middle=2, channels=[32, 128, 512],
use_fp16=False).eval().to(DEVICE)
enc.load_state_dict({k: v.float() for k, v in esd.items()}, strict=True)
return enc
def grid_to_shape_latent(ss_enc, occ):
"""64^3 occupancy -> (1,4096,8) shape latent (posterior MEAN).
Inverse of the pipeline's decode reshape (inference_pipeline.py:813-817)."""
g = torch.from_numpy(occ.astype(np.float32))[None, None].to(DEVICE)
z = ss_enc(g, sample_posterior=False).float() # (1,8,16,16,16)
return z.view(1, 8, 4096).permute(0, 2, 1).contiguous()
class Models:
"""Everything loaded ONCE."""
def __init__(self):
# SAM3D production pipeline
sys.path.insert(0, os.path.join(REPO, "notebook"))
from inference import Inference # flat import (repo convention)
cfg = "checkpoints/hf/pipeline.yaml"
print(f"[af] building Inference({cfg}) ...", flush=True)
self.pipe = Inference(cfg, compile=False)._pipeline
try:
self.pipe.rendering_engine = "pytorch3d"
except Exception:
pass
# SLAT encoder (gated repo) + SS encoder (for z1fwd)
assert SLAT_ENC_YAML and SLAT_ENC_CKPT, "slat_encoder ckpt/yaml not found in HF cache"
print(f"[af] slat_encoder <- {SLAT_ENC_CKPT[0]}", flush=True)
self.slat_enc = load_pretrained(SLAT_ENC_YAML[0], SLAT_ENC_CKPT[0])
self.ss_enc = load_ss_encoder()
# DINOv2
print("[af] loading DINOv2 (dinov2_vitl14_reg) ...", flush=True)
from torchvision import transforms
self.dino = torch.hub.load('facebookresearch/dinov2', 'dinov2_vitl14_reg').eval().cuda()
self.dino_norm = transforms.Normalize(mean=[0.485, 0.456, 0.406],
std=[0.229, 0.224, 0.225])
# normalization constants (pipeline.yaml values)
self.slat_mean = self.pipe.slat_mean.float().cpu().numpy() # (8,)
self.slat_std = self.pipe.slat_std.float().cpu().numpy()
print(f"[af] slat_mean/std from pipeline.yaml: mean[0]={self.slat_mean[0]:.4f} "
f"std[0]={self.slat_std[0]:.4f}", flush=True)
print("[af] models ready", flush=True)
# --------------------------------------------------------------------------- #
# appearance seed: build the normalized SLAT latent for VISIBLE voxels of `coords`
# --------------------------------------------------------------------------- #
def _coord_hash(xyz):
xyz = np.asarray(xyz, dtype=np.int64)
return (xyz[:, 0] * VOX + xyz[:, 1]) * VOX + xyz[:, 2]
def build_appearance_seed(M: Models, coords: torch.Tensor, dino_pts, vinfos, views):
"""-> (z_norm_full (N,8) float32, visible (N,) bool).
coords: (N,4) [batch,x,y,z] int on cuda -- the FINAL stage-1 coords.
dino_pts[view]: (1,1024,37,37) RAW patch tokens for that view.
Only visible rows of z_norm_full are meaningful; the rest are zeros.
"""
from sam3d_objects.model.backbone.tdfy_dit.modules import sparse as sp
idx = coords[:, 1:].detach().cpu().numpy().astype(np.int64) # (N,3)
centers = (idx.astype(np.float64) + 0.5) / VOX - 0.5
N = len(idx)
# VIEW0-PRIORITY (anchor-single MV mode): a voxel VISIBLE from view0 (== the
# SV reference) takes view0's DINO feature ALONE -- auxiliary views NEVER
# overwrite the reference anchor on shared voxels (spec: extra views are
# additive, filling only voxels view0 cannot see). This keeps the
# co-visible-voxel seed byte-identical to SV, so the MV mesh stays a strict
# superset of SV instead of drifting from averaged aux features. Auxiliary
# views contribute (averaged) ONLY on voxels view0 does not see.
# With 1 view this is a no-op == the original union path.
prio = os.environ.get("AF_VIEW0_PRIORITY") == "1" and len(views) > 1
if prio:
v0 = views[0]
u0, vv0, z0, vis0 = project_visible(centers, vinfos[v0])
f0 = sample_feats(dino_pts[v0], u0, vv0, vinfos[v0]) # (N,1024)
vis0_t = torch.from_numpy(vis0).cuda()
feat_vis = torch.zeros(N, 1024, device=DEVICE)
feat_vis[vis0_t] = f0[vis0_t]
aux_sum = torch.zeros(N, 1024, device=DEVICE)
aux_cnt = torch.zeros(N, device=DEVICE)
for view in views[1:]:
u, vv, z, vis = project_visible(centers, vinfos[view])
f = sample_feats(dino_pts[view], u, vv, vinfos[view])
m = torch.from_numpy((vis & ~vis0).astype(np.float32)).cuda()
aux_sum += f * m[:, None]
aux_cnt += m
aux_only = (aux_cnt > 0)
feat_vis[aux_only] = aux_sum[aux_only] / aux_cnt[aux_only][:, None]
vis_bool = vis0_t | aux_only
visible = vis_bool.cpu().numpy()
nz = vis_bool
n_v0 = int(vis0_t.sum()); n_aux = int((aux_only & ~vis0_t).sum())
print(f"[af] view0-priority appearance: view0_vox={n_v0} aux_only_vox={n_aux}",
flush=True)
else:
feat_sum = torch.zeros(N, 1024, device=DEVICE)
vis_count = torch.zeros(N, device=DEVICE)
for view in views:
v = vinfos[view]
u, vv, z, vis = project_visible(centers, v)
f = sample_feats(dino_pts[view], u, vv, v) # (N,1024)
vt = torch.from_numpy(vis.astype(np.float32)).cuda()
feat_sum += f * vt[:, None]
vis_count += vt
visible = (vis_count > 0).cpu().numpy()
nz = vis_count > 0
feat_vis = torch.zeros(N, 1024, device=DEVICE)
feat_vis[nz] = feat_sum[nz] / vis_count[nz][:, None]
z_norm_full = np.zeros((N, 8), np.float32)
vis_rows = np.nonzero(visible)[0]
if len(vis_rows) == 0:
return z_norm_full, visible
# encode the VISIBLE-ONLY subset (round-trip winner)
vis_idx = idx[vis_rows] # (Nv,3)
st_coords = torch.cat([torch.zeros(len(vis_idx), 1, dtype=torch.int32, device="cpu"),
torch.from_numpy(vis_idx).int()], dim=1)
st = sp.SparseTensor(feats=feat_vis[nz].float().cpu(), coords=st_coords).to(DEVICE)
with torch.autocast(device_type="cuda", dtype=DTYPE):
z_enc = M.slat_enc(st, sample_posterior=False)
zc = z_enc.coords.detach().cpu().numpy() # (Nv,4)
zf = z_enc.feats.float().detach().cpu().numpy() # (Nv,8)
# match encoder-output rows back to input rows BY COORDINATE (the sparse
# encoder may reorder), then normalize with pipeline slat_mean/std.
out_map = {h: i for i, h in enumerate(_coord_hash(zc[:, 1:]))}
want = _coord_hash(vis_idx)
sel = np.array([out_map[h] for h in want], dtype=np.int64)
z_sel = zf[sel] # (Nv,8) aligned to vis_rows
z_norm_full[vis_rows] = (z_sel - M.slat_mean[None]) / M.slat_std[None]
return z_norm_full, visible
# --------------------------------------------------------------------------- #
# vertex-colored mesh export (native decoder colors; canonical frame, no rotation)
#
# Split into a GPU DECODE step (main thread) and a CPU-bound FINALIZE step
# (trimesh vertex-color build + atomic GLB write), so the finalize can be
# pipelined onto a background worker while the GPU starts the next sample.
# --------------------------------------------------------------------------- #
def decode_mesh_arrays(M: Models, slat):
"""MAIN-THREAD GPU work: decode the mesh from a FLOAT32 latent (outside
autocast; flexicubes index_add_ needs fp32) and pull the vertex/face/attr
arrays to CPU numpy. Returns (verts, faces, va) -- everything the CPU-bound
finalize step needs, with NO further GPU dependency, so it is safe to hand
off to a worker thread."""
_t = time.time()
slat_f = slat.replace(slat.feats.float())
with torch.no_grad():
mesh = M.pipe.models["slat_decoder_mesh"](slat_f)[0]
if os.environ.get("AF_TIMING"):
torch.cuda.synchronize()
print(f"[af][time] mesh-decode(GPU) {time.time() - _t:.2f}s", flush=True)
verts = mesh.vertices.detach().cpu().numpy()
faces = mesh.faces.detach().cpu().numpy()
va = mesh.vertex_attrs.detach().cpu().numpy()
return verts, faces, va
def finalize_glb(verts, faces, va, out_path):
"""CPU-BOUND (thread-safe, no CUDA): build the vertex-colored trimesh from
the decoded arrays and atomically write the GLB in the canonical [-0.5,0.5]
frame. Byte-identical to the original synchronous export."""
_t = time.time()
rgb = np.clip(va[:, :3], 0.0, 1.0)
vc = np.concatenate([(rgb * 255).astype(np.uint8),
np.full((len(rgb), 1), 255, np.uint8)], axis=1)
tm = trimesh.Trimesh(vertices=verts, faces=faces, vertex_colors=vc,
process=False)
tmp = out_path + ".tmp.glb"
tm.export(tmp)
os.replace(tmp, out_path)
if os.environ.get("AF_TIMING"):
print(f"[af][time] finalize(CPU trimesh+GLB write) {time.time() - _t:.2f}s "
f"verts={len(verts)}", flush=True)
return len(verts), len(faces)
def export_vertex_colored_glb(M: Models, slat, out_path):
"""Synchronous decode + export (kept for backward-compat / A-B reference)."""
verts, faces, va = decode_mesh_arrays(M, slat)
return finalize_glb(verts, faces, va, out_path)
# --------------------------------------------------------------------------- #
# per-object forcing (one object at a time)
# --------------------------------------------------------------------------- #
def run_object_naive(M: Models, obj, exp, inputs_dir, npz_dir, views_tags, seed):
"""NAIVE default inference: the model's OWN image-conditioned pipeline with
NO geometry forcing (no z1fwd) and NO appearance forcing. Runs the stock
sample_sparse_structure (single or multi-view) from image conditioning, then
the stock sample_slat, then decodes the SAME vertex-colored mesh. The output
lives in the MODEL'S OWN frame/scale (not GT-depth grounded) -> must be
aligned to GT before eval, exactly like ReconViaGen. Returns
(verts, faces, va, 0, N)."""
pipe = M.pipe
multiview = len(views_tags) > 1
imgs = []
for t in views_tags:
p = os.path.join(inputs_dir, f"{obj}_{t}.png")
if not os.path.isfile(p):
raise FileNotFoundError(p)
imgs.append(np.array(Image.open(p).convert("RGBA")))
pms = np.load(os.path.join(npz_dir, obj, "da3_output.npz"))["pointmaps_sam3d"]
if pms.shape[0] < len(views_tags):
raise ValueError(f"{obj}: pointmaps {pms.shape} < views {len(views_tags)}")
pm_tensors = [torch.from_numpy(pms[i]).float() for i in range(len(views_tags))]
with pipe.device:
if multiview:
ss_input_dicts, slat_input_dicts = [], []
for img, pm in zip(imgs, pm_tensors):
pil = Image.fromarray(img)
pmd = pipe.compute_pointmap(pil, pointmap=pm)
ss_input_dicts.append(
pipe.preprocess_image(pil, pipe.ss_preprocessor, pointmap=pmd["pointmap"]))
slat_input_dicts.append(
pipe.preprocess_image(pil, pipe.slat_preprocessor))
else:
img = imgs[0]
pmd = pipe.compute_pointmap(img, pm_tensors[0])
ss_input_dict = pipe.preprocess_image(img, pipe.ss_preprocessor,
pointmap=pmd["pointmap"])
slat_input_dict = pipe.preprocess_image(img, pipe.slat_preprocessor)
# STAGE 1: default SS sampler (image-conditioned, NO seeded noise)
torch.manual_seed(seed)
if multiview:
ss_ret = pipe.sample_sparse_structure_multi_view(
ss_input_dicts, mode="multidiffusion", ss_weighting=False)
else:
ss_ret = pipe.sample_sparse_structure(ss_input_dict)
coords = ss_ret["coords"]
N = coords.shape[0]
# STAGE 2: default SLAT sampler (NO appearance seed)
if multiview:
slat = pipe.sample_slat_multi_view(
slat_input_dicts, coords, mode="multidiffusion")
else:
slat = pipe.sample_slat(slat_input_dict, coords)
verts, faces, va = decode_mesh_arrays(M, slat)
print(f"[af] NAIVE coords={N} verts={len(verts)} faces={len(faces)}",
flush=True)
return verts, faces, va, 0, N
def run_object_weighted_mv(M: Models, obj, exp, inputs_dir, npz_dir, views_tags, seed):
"""OFFICIAL FLAGSHIP weighted multi-view inference (SAM3D's INTENDED MV path).
Calls the repo's OWN weighted-fusion samplers -- the exact same functions that
inference_pipeline.py:run_multi_view invokes internally for the weighted path
(which run_inference_weighted.py:run_weighted_inference is the CLI wrapper for):
* STAGE 1: pipe.sample_sparse_structure_multi_view(..., ss_weighting=True,
ss_entropy_layer=9, ss_entropy_alpha=30.0, ss_warmup_steps=1)
-> attention-entropy WEIGHTED shape fusion (inference_pipeline.py:1695 &
:1009; entropy weights computed at :1091).
* STAGE 2: pipe.sample_slat_multi_view_weighted(..., weighting_config=cfg)
with cfg = the FLAGSHIP WeightingConfig from run_weighted_inference
(weight_source="entropy", entropy_alpha=30.0, attention_layer=6,
attention_step=0, min_weight=0.001) -> per-latent entropy WEIGHTED texture
fusion (inference_pipeline.py:1763 & :1283; two-pass warmup->weighted main).
This is DISTINCT from run_object_naive, which forces the UNWEIGHTED path
(sample_sparse_structure_multi_view(ss_weighting=False) +
sample_slat_multi_view(simple average)). Everything else -- inputs, GT-depth
pointmap conditioning (da3_output.npz["pointmaps_sam3d"]), preprocessing, and
the decoder-native vertex-colored canonical export -- is IDENTICAL to
run_object_naive so the naive-vs-weighted comparison is apples-to-apples.
weight_source="entropy" needs NO camera extrinsics (npz has none), matching the
flagship default. Single-view is not weighted (the model disables weighting for
1 view, run_inference_weighted.py:2727-2730), so this is 2v-only; 1v official ==
naive 1v. Returns (verts, faces, va, 0, N)."""
from sam3d_objects.utils.latent_weighting import WeightingConfig
pipe = M.pipe
assert len(views_tags) > 1, "weighted_mv is multi-view only (1v == naive 1v)"
imgs = []
for t in views_tags:
p = os.path.join(inputs_dir, f"{obj}_{t}.png")
if not os.path.isfile(p):
raise FileNotFoundError(p)
imgs.append(np.array(Image.open(p).convert("RGBA")))
pms = np.load(os.path.join(npz_dir, obj, "da3_output.npz"))["pointmaps_sam3d"]
if pms.shape[0] < len(views_tags):
raise ValueError(f"{obj}: pointmaps {pms.shape} < views {len(views_tags)}")
pm_tensors = [torch.from_numpy(pms[i]).float() for i in range(len(views_tags))]
# FLAGSHIP stage-2 weighting config (verbatim from run_weighted_inference
# defaults: stage2_weight_source/entropy_alpha/attention_layer/step/min_weight)
weighting_config = WeightingConfig(
weight_source="entropy",
use_entropy=True,
entropy_alpha=30.0,
attention_layer=6,
attention_step=0,
min_weight=0.001,
)
with pipe.device:
ss_input_dicts, slat_input_dicts = [], []
for img, pm in zip(imgs, pm_tensors):
pil = Image.fromarray(img)
pmd = pipe.compute_pointmap(pil, pointmap=pm)
ss_input_dicts.append(
pipe.preprocess_image(pil, pipe.ss_preprocessor, pointmap=pmd["pointmap"]))
slat_input_dicts.append(
pipe.preprocess_image(pil, pipe.slat_preprocessor))
# STAGE 1: official WEIGHTED sparse-structure fusion (ss_weighting=True)
torch.manual_seed(seed)
ss_ret = pipe.sample_sparse_structure_multi_view(
ss_input_dicts, mode="multidiffusion",
ss_weighting=True, ss_entropy_layer=9, ss_entropy_alpha=30.0,
ss_warmup_steps=1)
coords = ss_ret["coords"]
N = coords.shape[0]
# STAGE 2: official WEIGHTED SLAT fusion (entropy weighting_config)
slat, weight_manager = pipe.sample_slat_multi_view_weighted(
slat_input_dicts, coords, weighting_config=weighting_config)
verts, faces, va = decode_mesh_arrays(M, slat)
print(f"[af] WEIGHTED-MV coords={N} verts={len(verts)} faces={len(faces)}",
flush=True)
return verts, faces, va, 0, N
def run_object(M: Models, obj, method, exp, inputs_dir, npz_dir, views_tags,
seed, schedule="seed", reinject_until=1.0, steps=0):
"""Drive stage-1 (z1fwd) + stage-2 (default | appearance-forced) for one
object and decode the vertex-colored mesh. Returns
(verts, faces, va, n_vis, N) -- the GPU inference + decode result. The
CPU-bound trimesh build + GLB write is left to the caller (finalize_glb),
so it can be pipelined with the next sample's GPU inference.
schedule (only relevant for method=sam3d_geom_dino):
seed SEED-ONLY: scatter z_forced into the visible rows of the t=0
initial-noise tensor; the flow then runs untouched. DEFAULT and
byte-identical to the original behaviour.
reinject HARD RE-INJECTION (RePaint) for the OOD case: after EVERY Euler
step (state now at t_next) overwrite the visible rows with the
on-path value for z_forced using a FIXED eps drawn once:
x[vis] = (1-(1-sigma_min)*t_next)*eps_fixed + t_next*z_forced
(SAM3D t:0 noise -> 1 data, interpolant
x_t=(1-(1-sigma_min)t)*eps + t*z, flow_matching/model.py:116-127).
At t=0 visible rows = eps_fixed (pure noise, correct start); at
t=1 visible rows = z_forced exactly (sigma_min=0). Invisible rows
integrate normally. Implemented by monkeypatching the shipped
FlowMatching solver's Euler `step` (per-visible-row override, which
the whole-tensor noise_override hook cannot express)."""
pipe = M.pipe
# ---- ANCHOR-SINGLE MV MODE (AF_COND_SINGLE=1) ---------------------------
# DECOUPLE the SLAT/SS IMAGE-CONDITIONING views (`cond_tags`) from the
# geometry+DINO FORCING views (`force_tags`). Under multi-image SLAT
# (sample_slat_multi_view multidiffusion) the strong single reference view's
# texture/latent is DILUTED by the weaker auxiliary views -> MV mesh worse
# than SV for some objects (sugar_box/mug ADDS 100->32). When AF_COND_SINGLE
# is set, the SS + SLAT samplers run in SINGLE-IMAGE mode anchored on view0
# (`force_tags[0]` == front == the exact SV reference view), while the extra
# views are still used ONLY as ADDITIVE constraints: multi-view DINO
# appearance forcing fills voxels view0 cannot see (build_appearance_seed
# over force_tags) and the geometry seed z1 stays the view0 depth shell
# (ref-occ). This makes MV a strict SUPERSET of SV (>= guaranteed): view0's
# SLAT anchor is never overwritten, extra views only add occluded appearance.
force_tags = list(views_tags) # geometry + DINO forcing
cond_single = os.environ.get("AF_COND_SINGLE") == "1"
cond_tags = [force_tags[0]] if cond_single else list(views_tags) # SS/SLAT cond
multiview = len(cond_tags) > 1
# ---- inputs / pointmaps (CONDITIONING views only) -----------------------
imgs = []
for t in cond_tags:
p = os.path.join(inputs_dir, f"{obj}_{t}.png")
if not os.path.isfile(p):
raise FileNotFoundError(p)
imgs.append(np.array(Image.open(p).convert("RGBA")))
pms = np.load(os.path.join(npz_dir, obj, "da3_output.npz"))["pointmaps_sam3d"]
# pointmaps are stored in force_tags order; view0 (front) is row 0, so the
# single-cond anchor picks row 0 (the SV reference pointmap).
if pms.shape[0] < len(cond_tags):
raise ValueError(f"{obj}: pointmaps {pms.shape} < cond {len(cond_tags)}")
pm_tensors = [torch.from_numpy(pms[i]).float() for i in range(len(cond_tags))]
if cond_single:
print(f"[af] ANCHOR-SINGLE: SS/SLAT cond=[{cond_tags[0]}] (view0), "
f"DINO+geom forcing over {force_tags}", flush=True)
# ---- STAGE-1 geometry forcing (z1fwd): build z1 = ss_enc(visible grid) ---
# force_tags for the depth union; overridden by ref-occ (view0 shell) when
# AF_OCC=hull -> geometry seed identical to SV.
gt_vis, gt_full, free = load_gt_synth(Path(exp), obj, force_tags, n=VOX)
rng = np.random.default_rng(seed)
occ = build_input_grid(gt_vis, free, "zeros", rng)
z1 = grid_to_shape_latent(M.ss_enc, occ) # (1,4096,8) float32
# ---- per-view DINO tokens + view infos (ALL force_tags for appearance) ---
dino_pts = {t: dino_features(M.dino, M.dino_norm,
os.path.join(inputs_dir, f"{obj}_{t}.png"))
for t in force_tags}
vinfos = {t: load_view(exp, obj, t) for t in force_tags}
ss_gen = pipe.models["ss_generator"]
slat_gen = pipe.models["slat_generator"]
orig_ss_noise = ss_gen._generate_noise
orig_slat_noise = slat_gen._generate_noise
def ss_noise_seeded(x_shape, x_device):
out = orig_ss_noise(x_shape, x_device)
if isinstance(out, dict) and "shape" in out:
out["shape"] = z1.to(x_device).to(out["shape"].dtype)
return out
# -------- preprocess (mirror InferencePipelinePointMap.run bodies) -------
with pipe.device:
if multiview:
pipe.merge_image_and_mask # noqa: (parity)
ss_input_dicts, slat_input_dicts = [], []
for img, pm in zip(imgs, pm_tensors):
pil = Image.fromarray(img)
pmd = pipe.compute_pointmap(pil, pointmap=pm)
ss_input_dicts.append(
pipe.preprocess_image(pil, pipe.ss_preprocessor, pointmap=pmd["pointmap"]))
slat_input_dicts.append(
pipe.preprocess_image(pil, pipe.slat_preprocessor))
else:
img = imgs[0]
pmd = pipe.compute_pointmap(img, pm_tensors[0])
ss_input_dict = pipe.preprocess_image(img, pipe.ss_preprocessor,
pointmap=pmd["pointmap"])
slat_input_dict = pipe.preprocess_image(img, pipe.slat_preprocessor)
# -------- STAGE 1 : seeded SS sampler -> coords ----------------------
torch.manual_seed(seed)
ss_gen._generate_noise = ss_noise_seeded
try:
if multiview:
ss_ret = pipe.sample_sparse_structure_multi_view(
ss_input_dicts, mode="multidiffusion", ss_weighting=False)
else:
ss_ret = pipe.sample_sparse_structure(ss_input_dict)
finally:
ss_gen._generate_noise = orig_ss_noise
coords = ss_ret["coords"]
N = coords.shape[0]
# -------- build stage-2 base noise (identical for both methods) ------
bs = (slat_input_dicts[0] if multiview else slat_input_dict)["image"].shape[0]
base = orig_slat_noise((bs, N, 8), DEVICE) # draws from post-stage1 RNG
n_vis = 0
reinject_state = None
if method == "sam3d_geom_dino":
z_norm_full, visible = build_appearance_seed(M, coords, dino_pts, vinfos, force_tags)
n_vis = int(visible.sum())
vis_rows = np.nonzero(visible)[0]
vis_rows_t = torch.from_numpy(vis_rows).to(DEVICE)
zt = torch.from_numpy(z_norm_full[vis_rows]).to(DEVICE).to(base.dtype)
if schedule == "seed":
# --- DEFAULT / UNCHANGED seed-only path (byte-identical) ------
# sanity: the normalized seed should be O(1) like the noise it replaces
print(f"[af] seed std={float(zt.std()):.3f} mean={float(zt.mean()):.3f} "
f"| base-noise std={float(base.std()):.3f}", flush=True)
base[:, torch.from_numpy(vis_rows).to(DEVICE), :] = zt # SEED visible rows
elif schedule == "reinject" and n_vis > 0:
# --- HARD RE-INJECTION path (additive; base stays pure noise) -
sigma_min = float(getattr(slat_gen, "sigma_min", 0.0))
eps_fixed = base[:, vis_rows_t, :].clone() # FIXED eps (bs,Nv,8)
z_forced_b = zt.float().unsqueeze(0) # (1,Nv,8) -> broadcast bs
print(f"[af] reinject rows={n_vis}/{N} sigma_min={sigma_min:g} "
f"reinject_until={reinject_until:g} "
f"| z_forced std={float(zt.std()):.3f} mean={float(zt.mean()):.3f} "
f"| base std={float(base.std()):.3f}", flush=True)
reinject_state = dict(vr=vis_rows_t, eps=eps_fixed, z=z_forced_b,
smin=sigma_min, last=None,
until=float(reinject_until),
n_pinned=0, n_released=0, last_pinned_t=None,
first_released_t=None)
def slat_noise_fixed(x_shape, x_device):
return base.to(x_device)
# -------- validation: print the actual stage-2 t-schedule ------------
n_steps_eff = steps if steps else int(slat_gen.inference_steps)
try:
t_seq_dbg = slat_gen._prepare_t(steps if steps else None)
t_list = [float(x) for x in t_seq_dbg]
pin_flags = [("PIN" if tn <= reinject_until + 1e-6 else "free")
for tn in t_list[1:]] # decision is on each t_next
print(f"[af] stage2 steps={n_steps_eff} t_seq(len={len(t_list)}): "
f"[{t_list[0]:.3f} .. {t_list[-1]:.3f}]", flush=True)
if reinject_state is not None:
n_pin_expected = sum(1 for f in pin_flags if f == "PIN")
# first t_next that is released (>T), if any
rel = [tn for tn in t_list[1:] if tn > reinject_until + 1e-6]
stop_at = f"{rel[0]:.3f}" if rel else "never(full pin)"
print(f"[af] reinject t_next schedule={pin_flags} "
f"-> pin {n_pin_expected}/{len(pin_flags)} steps, "
f"stops reinjecting at t_next={stop_at} (until={reinject_until:g})",
flush=True)
except Exception as _e:
print(f"[af] (t-schedule debug skipped: {_e})", flush=True)
# -------- STAGE 2 : seeded SLAT sampler ------------------------------
slat_gen._generate_noise = slat_noise_fixed
orig_step = slat_gen._solver.step
if reinject_state is not None:
def patched_step(dynamics_fn, x_t, t, dt, *a, **k):
x_tp1 = orig_step(dynamics_fn, x_t, t, dt, *a, **k)
rs = reinject_state
t_next = float(t) + float(dt) # Euler advances t0 -> t0+dt = t_next
# ANNEAL: only pin while t_next <= T; release (free) afterwards.
if t_next <= rs["until"] + 1e-6:
onpath = (1.0 - (1.0 - rs["smin"]) * t_next) * rs["eps"] + t_next * rs["z"]
x_tp1[:, rs["vr"], :] = onpath.to(x_tp1.dtype)
rs["n_pinned"] += 1
rs["last_pinned_t"] = t_next
if abs(t_next - 1.0) < 1e-6: # final state pinned -> stash for check
rs["last"] = x_tp1[:, rs["vr"], :].detach().clone()
else:
rs["n_released"] += 1
if rs["first_released_t"] is None:
rs["first_released_t"] = t_next
return x_tp1
slat_gen._solver.step = patched_step
slat_steps_kw = {"inference_steps": steps} if steps else {}
try:
if multiview:
slat = pipe.sample_slat_multi_view(
slat_input_dicts, coords, mode="multidiffusion", **slat_steps_kw)
else:
slat = pipe.sample_slat(slat_input_dict, coords, **slat_steps_kw)
finally:
slat_gen._generate_noise = orig_slat_noise
if reinject_state is not None:
try:
del slat_gen._solver.step # remove instance attr -> class method
except AttributeError:
slat_gen._solver.step = orig_step
# validation: report the pin/release accounting and the final-state check
if reinject_state is not None:
rs = reinject_state
print(f"[af] reinject applied: pinned {rs['n_pinned']} steps "
f"(last pinned t_next={rs['last_pinned_t']}), released "
f"{rs['n_released']} steps (first released t_next={rs['first_released_t']}) "
f"| until={rs['until']:g}", flush=True)
if rs["last"] is not None:
# only meaningful when the final step was pinned (until>=1)
diff = float((rs["last"] - rs["z"]).abs().max())
print(f"[af] reinject check (final PINNED): "
f"max|x[vis]_final - z_forced|={diff:.3e} "
f"(expected ~sigma_min*|eps| = {rs['smin']:g}*O(1))", flush=True)
else:
print(f"[af] reinject check: final step was RELEASED (until={rs['until']:g}"
f"<1), visible rows integrated freely after t_next="
f"{rs['first_released_t']} -- annealed release confirmed", flush=True)
verts, faces, va = decode_mesh_arrays(M, slat)
print(f"[af] coords={N} visible_forced={n_vis} verts={len(verts)} "
f"faces={len(faces)}", flush=True)
return verts, faces, va, n_vis, N
# --------------------------------------------------------------------------- #
def main():
ap = argparse.ArgumentParser(description=__doc__,
formatter_class=argparse.RawDescriptionHelpFormatter)
ap.add_argument("--method", required=True,
choices=["sam3d_geom", "sam3d_geom_dino", "sam3d_naive",
"sam3d_weighted_mv"])
ap.add_argument("--schedule", default="seed", choices=["seed", "reinject"],
help="appearance-forcing schedule (only for sam3d_geom_dino): "
"seed (default, SEED-ONLY, unchanged) | reinject "
"(RePaint hard re-injection, pins visible rows every step)")
ap.add_argument("--reinject-until", type=float, default=1.0,
dest="reinject_until",
help="(reinject schedule only) ANNEAL the hard pin: only pin "
"visible rows while flow time t <= T (t:0 noise -> 1 data), "
"then STOP re-injecting for the rest of the flow, letting "
"visible rows integrate freely. Default 1.0 = full reinject "
"(byte-identical to before). 0.5 = pin the noisy first half, "
"release the second half.")
ap.add_argument("--steps", type=int, default=0,
help="override the stage-2 SLAT sampler step count (default 0 = "
"the model's own ~25). Halve = ~12. Stage-1 is unaffected.")
ap.add_argument("--selection", required=True)
ap.add_argument("--inputs", required=True)
ap.add_argument("--exp", required=True,
help="EXPDIR with renders/, inputs/, npz_1v/, npz_2v/")
ap.add_argument("--out", required=True)
ap.add_argument("--views", type=int, choices=[1, 2, 4, 8], required=True)
ap.add_argument("--seed", type=int, default=42)
ap.add_argument("--limit", type=int, default=0)
ap.add_argument("--objects", nargs="+", default=None)
ap.add_argument("--gpu", default="4")
ap.add_argument("--shard", type=int, default=0,
help="this shard index in [0, nshards) for co-located parallelism")
ap.add_argument("--nshards", type=int, default=1,
help="total number of shards; objects are split round-robin")
ap.add_argument("--async-export", dest="async_export", action="store_true",
default=True,
help="(default ON) pipeline the CPU-bound mesh finalize "
"(trimesh build + GLB write) on a background worker so "
"the GPU starts the next sample's inference immediately.")
ap.add_argument("--no-async-export", dest="async_export", action="store_false",
help="disable async export; finalize synchronously "
"(byte-identical output, for A-B).")
args = ap.parse_args()
exp = os.path.abspath(args.exp)
inputs_dir = os.path.abspath(args.inputs)
npz_dir = os.path.join(exp, {1: "npz_1v", 2: "npz_2v", 4: "npz_4v", 8: "npz_8v"}[args.views])
out_dir = os.path.abspath(args.out)
os.makedirs(out_dir, exist_ok=True)
views_tags = {1: ["front"], 2: ["front", "side"],
4: ["front", "side", "back", "oside"],
8: ["front", "side", "back", "oside", "top", "bottom", "top2", "bottom2"]}[args.views]
data = json.loads(Path(args.selection).read_text())
sels = data["selections"] if isinstance(data, dict) else data
objects = [s["object"] for s in sels]
if args.objects:
want = set(args.objects)
objects = [o for o in objects if o in want]
if args.limit:
objects = objects[: args.limit]
# co-located-process sharding: split objects round-robin so N shards run in
# parallel (2 per GPU on GPUs 4 & 5) with no overlap. Idempotent skip makes
# overlap harmless anyway, but round-robin keeps the shards balanced.
if args.nshards > 1:
objects = objects[args.shard::args.nshards]
print(f"[af] method={args.method} schedule={args.schedule} "
f"reinject_until={args.reinject_until} steps={args.steps or 'default(25)'} "
f"views={args.views} async_export={args.async_export} "
f"seed={args.seed} shard={args.shard}/{args.nshards} n_obj={len(objects)} "
f"CUDA={os.environ.get('CUDA_VISIBLE_DEVICES')}", flush=True)
t0 = time.time()
M = Models()
print(f"[af] model load {time.time() - t0:.1f}s", flush=True)
# ---- producer/consumer export pipeline ---------------------------------
# Main thread: GPU inference + mesh decode (-> CPU numpy arrays). A single
# background worker does the CPU-bound trimesh build + atomic GLB write for
# the PREVIOUS sample while the GPU runs the NEXT sample's inference.
counts = {"ok": 0, "fail": 0, "skip": 0}
executor = ThreadPoolExecutor(max_workers=1) if args.async_export else None
pending = [] # list of dict(future, i, obj, out_path, t1)
def _reap(item, wait=False):
fut = item["future"]
if not wait and not fut.done():
return False
try:
nv, nf = fut.result()
counts["ok"] += 1
sz = os.path.getsize(item["out_path"]) if os.path.isfile(item["out_path"]) else 0
print(f"[af {item['i']}/{len(objects)}] OK {item['obj']} "
f"{time.time() - item['t1']:.1f}s (async) -> {item['out_path']} "
f"({sz} bytes)", flush=True)
except Exception:
counts["fail"] += 1
traceback.print_exc()
print(f"[af {item['i']}/{len(objects)}] FAIL(export) {item['obj']} "
f"{time.time() - item['t1']:.1f}s", flush=True)
return True
for i, obj in enumerate(objects, 1):
# reap any finished exports opportunistically (non-blocking)
pending = [it for it in pending if not _reap(it, wait=False)]
out_path = os.path.join(out_dir, f"{obj}.glb")
if os.path.isfile(out_path):
counts["skip"] += 1
print(f"[af {i}/{len(objects)}] SKIP (exists) {obj}", flush=True)
continue
t1 = time.time()
try:
print(f"[af {i}/{len(objects)}] {obj}", flush=True)
if args.method == "sam3d_naive":
verts, faces, va, n_vis, N = run_object_naive(
M, obj, exp, inputs_dir, npz_dir, views_tags, args.seed)
elif args.method == "sam3d_weighted_mv":
# weighted fusion is multi-view only; at 1 view it degenerates to
# the stock single-view pipeline == naive 1v (weighting disabled).
if len(views_tags) > 1:
verts, faces, va, n_vis, N = run_object_weighted_mv(
M, obj, exp, inputs_dir, npz_dir, views_tags, args.seed)
else:
verts, faces, va, n_vis, N = run_object_naive(
M, obj, exp, inputs_dir, npz_dir, views_tags, args.seed)
else:
verts, faces, va, n_vis, N = run_object(
M, obj, args.method, exp, inputs_dir, npz_dir, views_tags,
args.seed, schedule=args.schedule,
reinject_until=args.reinject_until, steps=args.steps)
except Exception:
counts["fail"] += 1
traceback.print_exc()
print(f"[af {i}/{len(objects)}] FAIL {obj} {time.time() - t1:.1f}s",
flush=True)
torch.cuda.empty_cache()
continue
torch.cuda.empty_cache() # free GPU before the export overlaps next infer
if executor is not None:
fut = executor.submit(finalize_glb, verts, faces, va, out_path)
pending.append(dict(future=fut, i=i, obj=obj, out_path=out_path, t1=t1))
else:
try:
nv, nf = finalize_glb(verts, faces, va, out_path)
counts["ok"] += 1
print(f"[af {i}/{len(objects)}] OK {obj} {time.time() - t1:.1f}s "
f"-> {out_path} ({os.path.getsize(out_path)} bytes)", flush=True)
except Exception:
counts["fail"] += 1
traceback.print_exc()
print(f"[af {i}/{len(objects)}] FAIL(export) {obj} "
f"{time.time() - t1:.1f}s", flush=True)
# ---- drain the pipeline (no sample lost) --------------------------------
for it in pending:
_reap(it, wait=True)
if executor is not None:
executor.shutdown(wait=True)
print(f"APPFORCE SAM3D DONE method={args.method} ok={counts['ok']} "
f"fail={counts['fail']} skip={counts['skip']} total={len(objects)}",
flush=True)
if __name__ == "__main__":
main()