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()