File size: 53,949 Bytes
4c3d957 | 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 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 280 281 282 283 284 285 286 287 288 289 290 291 292 293 294 295 296 297 298 299 300 301 302 303 304 305 306 307 308 309 310 311 312 313 314 315 316 317 318 319 320 321 322 323 324 325 326 327 328 329 330 331 332 333 334 335 336 337 338 339 340 341 342 343 344 345 346 347 348 349 350 351 352 353 354 355 356 357 358 359 360 361 362 363 364 365 366 367 368 369 370 371 372 373 374 375 376 377 378 379 380 381 382 383 384 385 386 387 388 389 390 391 392 393 394 395 396 397 398 399 400 401 402 403 404 405 406 407 408 409 410 411 412 413 414 415 416 417 418 419 420 421 422 423 424 425 426 427 428 429 430 431 432 433 434 435 436 437 438 439 440 441 442 443 444 445 446 447 448 449 450 451 452 453 454 455 456 457 458 459 460 461 462 463 464 465 466 467 468 469 470 471 472 473 474 475 476 477 478 479 480 481 482 483 484 485 486 487 488 489 490 491 492 493 494 495 496 497 498 499 500 501 502 503 504 505 506 507 508 509 510 511 512 513 514 515 516 517 518 519 520 521 522 523 524 525 526 527 528 529 530 531 532 533 534 535 536 537 538 539 540 541 542 543 544 545 546 547 548 549 550 551 552 553 554 555 556 557 558 559 560 561 562 563 564 565 566 567 568 569 570 571 572 573 574 575 576 577 578 579 580 581 582 583 584 585 586 587 588 589 590 591 592 593 594 595 596 597 598 599 600 601 602 603 604 605 606 607 608 609 610 611 612 613 614 615 616 617 618 619 620 621 622 623 624 625 626 627 628 629 630 631 632 633 634 635 636 637 638 639 640 641 642 643 644 645 646 647 648 649 650 651 652 653 654 655 656 657 658 659 660 661 662 663 664 665 666 667 668 669 670 671 672 673 674 675 676 677 678 679 680 681 682 683 684 685 686 687 688 689 690 691 692 693 694 695 696 697 698 699 700 701 702 703 704 705 706 707 708 709 710 711 712 713 714 715 716 717 718 719 720 721 722 723 724 725 726 727 728 729 730 731 732 733 734 735 736 737 738 739 740 741 742 743 744 745 746 747 748 749 750 751 752 753 754 755 756 757 758 759 760 761 762 763 764 765 766 767 768 769 770 771 772 773 774 775 776 777 778 779 780 781 782 783 784 785 786 787 788 789 790 791 792 793 794 795 796 797 798 799 800 801 802 803 804 805 806 807 808 809 810 811 812 813 814 815 816 817 818 819 820 821 822 823 824 825 826 827 828 829 830 831 832 833 834 835 836 837 838 839 840 841 842 843 844 845 846 847 848 849 850 851 852 853 854 855 856 857 858 859 860 861 862 863 864 865 866 867 868 869 870 871 872 873 874 875 876 877 878 879 880 881 882 883 884 885 886 887 888 889 890 891 892 893 894 895 896 897 898 899 900 901 902 903 904 905 906 907 908 909 910 911 912 913 914 915 916 917 918 919 920 921 922 923 924 925 926 927 928 929 930 931 932 933 934 935 936 937 938 939 940 941 942 943 944 945 946 947 948 949 950 951 952 953 954 955 956 957 958 959 960 961 962 963 964 965 966 967 968 969 970 971 972 973 974 975 976 977 978 979 980 981 982 983 984 985 986 987 988 989 990 991 992 993 994 995 996 997 998 999 1000 1001 1002 1003 1004 1005 1006 1007 1008 1009 1010 1011 1012 1013 1014 1015 1016 1017 1018 1019 1020 1021 1022 1023 1024 1025 1026 1027 1028 1029 1030 1031 1032 1033 1034 1035 1036 1037 1038 1039 1040 1041 1042 1043 1044 1045 1046 1047 1048 1049 1050 1051 1052 1053 1054 1055 1056 1057 1058 1059 1060 1061 1062 1063 1064 1065 1066 | #!/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()
|