#!/usr/bin/env python """PASS 3 (REFERENCE ENV) — stage-2 base noise + appearance seed on GIVEN coords. Three-pass SSFlow eval driver, pass 3 of 3 (see dump_seed.py for the design). A near-copy of ``tools/val_stage1_ref.py --mode stage1``. Steps 1-4 of ``stage1_and_seed`` run verbatim (only to reproduce ``bs`` and to build the DINO features / view infos); step 5 (the pretrained SS sampling) is REPLACED by loading ``coords`` from ``/.npz`` (written by pass 2 from the trained SSFlowModel); steps 6 and 7 are then verbatim: base = slat_gen._generate_noise((bs, N, 8), DEVICE) z_norm_full, visible = baf.build_appearance_seed(M, coords, dino_pts, vinfos, views_tags) and the npz carries the SAME keys as the reference cache: coords (N,4) int32 | base (1,N,8) f32 | z_norm (N,8) f32 | visible (N,) bool + dino_ (1,1024,37,37) f32 per view tag. ====================================================================== CAVEAT — THE STAGE-2 ``base`` NOISE IS A DIFFERENT DRAW ====================================================================== In the reference, ``torch.manual_seed(seed)`` is called, the SS pipeline then samples (consuming the torch RNG stream), and ``base`` is drawn from the POST-STAGE-1 RNG state. Here the SS sampling does not happen in this process at all, so that RNG state cannot be reproduced. We therefore call ``torch.manual_seed(seed)`` and draw ``base`` IMMEDIATELY: a FRESH draw at the same seed, NOT the reference run's post-SS draw. Consequences, stated plainly for the results table: * ``base`` here is NOT bit-comparable to the reference ``val_appforce/stage1`` dumps, so a parity check against them WILL fail on ``base`` (and on the SLAT that follows from it) even when everything else is correct. * It IS deterministic and identical across objects-with-the-same-N and across reruns, so trained-vs-trained comparisons over this cache are fair. * For an apples-to-apples ours-vs-reference stage-2 comparison, the reference arm must be regenerated with the SAME fresh-draw convention (i.e. this script pointed at the reference coords), not with the existing dumps. Everything preceding ``base`` (the preprocessing block, which runs BEFORE ``manual_seed`` in the reference too) is unchanged, so ``bs`` is derived exactly as the reference derives it. Launch (re-execs itself into the reference env; caller sets CUDA_VISIBLE_DEVICES to ONE gpu): CUDA_VISIBLE_DEVICES=7 python finish_cache.py \ --exp .../toys4k100_tex --views 2 --coords-dir --out /2v """ import os import sys ENV_PY = "/lp-dev/jonghoon/mv-mesh/envs/mv-sam3d/bin/python" ENV = "/lp-dev/jonghoon/mv-mesh/envs/mv-sam3d" MVMESH = "/lp-dev/jonghoon/mv-mesh" METRICS_DIR = os.path.join(MVMESH, "metrics") # --- re-exec preamble: VERBATIM from tools/val_stage1_ref.py ----------------- if not os.environ.get("_VAL_STAGE1_INENV"): env = dict(os.environ) env["_VAL_STAGE1_INENV"] = "1" env["_APPFORCE_SAM3D_INENV"] = "1" # batch_appforce must not re-exec env.pop("PYTHONPATH", None) # never the training repo's sam3d_objects env["PATH"] = os.path.join(ENV, "bin") + os.pathsep + env.get("PATH", "") assert env.get("CUDA_VISIBLE_DEVICES"), "set CUDA_VISIBLE_DEVICES to one GPU" os.execve(ENV_PY, [ENV_PY, os.path.abspath(__file__)] + sys.argv[1:], env) os.environ["_APPFORCE_SAM3D_INENV"] = "1" sys.path.insert(0, METRICS_DIR) import batch_appforce_sam3d as baf # noqa: E402 (chdir REPO, sys.path, env defaults) import argparse # noqa: E402 import json # noqa: E402 import time # noqa: E402 import traceback # noqa: E402 from pathlib import Path # noqa: E402 import numpy as np # noqa: E402 import torch # noqa: E402 from PIL import Image # noqa: E402 DEVICE = baf.DEVICE # --- VIEW TAG TABLE: verbatim from tools/val_stage1_ref.py ------------------ VIEW_TAGS = ["front", "side", "back", "oside", "top", "bottom", "top2", "bottom2"] def views_tags_for(n: int): assert 1 <= n <= len(VIEW_TAGS), f"--views {n} unsupported (max {len(VIEW_TAGS)})" return VIEW_TAGS[:n] def npz_dir_for(exp: str, n: int) -> str: return os.path.join(exp, f"npz_{n}v") def _objects(args): """Verbatim from tools/val_stage1_ref.py::_objects.""" if args.objects_file: objs = [l.strip() for l in open(args.objects_file) if l.strip()] elif args.objects: objs = list(args.objects) else: data = json.loads(Path(os.path.join(args.exp, "selection.json")).read_text()) sels = data["selections"] if isinstance(data, dict) else data objs = [s["object"] for s in sels] if args.limit: objs = objs[: args.limit] if args.nshards > 1: objs = objs[args.shard::args.nshards] return objs # --------------------------------------------------------------------------- # def seed_from_given_coords(M, obj, exp, inputs_dir, npz_dir, views_tags, seed, coords_np): """``val_stage1_ref.stage1_and_seed`` with step 5 replaced by ``coords_np``. Every line below is the reference block unchanged EXCEPT: (a) the SS generator is never built/monkeypatched/sampled, (b) ``coords`` comes from the pass-2 dump, (c) ``torch.manual_seed(seed)`` is immediately followed by the ``base`` draw (see the module docstring's CAVEAT). """ 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))] dino_pts = {t: baf.dino_features(M.dino, M.dino_norm, os.path.join(inputs_dir, f"{obj}_{t}.png")) for t in views_tags} vinfos = {t: baf.load_view(exp, obj, t) for t in views_tags} slat_gen = pipe.models["slat_generator"] orig_slat_noise = slat_gen._generate_noise with pipe.device: # ---- preprocessing: VERBATIM (runs BEFORE manual_seed in the reference # too), needed here only to derive ``bs`` exactly as the reference does. 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) # ---- step 5 REPLACED: coords come from the trained SSFlowModel ------- coords = torch.from_numpy(coords_np.astype(np.int32)).to(DEVICE) assert coords.dim() == 2 and coords.shape[1] == 4, tuple(coords.shape) N = coords.shape[0] if N == 0: raise ValueError(f"{obj}: empty coords from the trained stage-1") # ---- step 6 (verbatim, modulo the documented RNG caveat) ------------ bs = (slat_input_dicts[0] if multiview else slat_input_dict)["image"].shape[0] torch.manual_seed(seed) # CAVEAT: fresh draw base = orig_slat_noise((bs, N, 8), DEVICE) z_norm_full, visible = baf.build_appearance_seed(M, coords, dino_pts, vinfos, views_tags) # ---- step 7 (verbatim) -------------------------------------------------- out = dict(coords=coords.detach().cpu().numpy().astype(np.int32), base=base.detach().float().cpu().numpy(), z_norm=z_norm_full.astype(np.float32), visible=visible.astype(bool)) for t in views_tags: # RAW DINO patch tokens (1,1024,37,37), parity aid out[f"dino_{t}"] = dino_pts[t].detach().float().cpu().numpy() return out def main(): ap = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter) ap.add_argument("--exp", required=True, help="EXPDIR (renders/, inputs/, npz_{n}v/, selection.json)") ap.add_argument("--views", type=int, default=2) ap.add_argument("--coords-dir", dest="coords_dir", required=True, help="pass-2 dump dir (.npz with coords)") ap.add_argument("--out", required=True, help="stage-1 cache dir") ap.add_argument("--seed", type=int, default=42) ap.add_argument("--objects", nargs="+", default=None) ap.add_argument("--objects-file", default=None) ap.add_argument("--limit", type=int, default=0) ap.add_argument("--shard", type=int, default=0) ap.add_argument("--nshards", type=int, default=1) ap.add_argument("--force", action="store_true") args = ap.parse_args() exp = os.path.abspath(args.exp) inputs_dir = os.path.join(exp, "inputs") npz_dir = npz_dir_for(exp, args.views) views_tags = views_tags_for(args.views) coords_dir = Path(args.coords_dir) out = Path(args.out) out.mkdir(parents=True, exist_ok=True) objs = _objects(args) print(f"[fin] exp={exp} views={args.views} n_obj={len(objs)} " f"shard={args.shard}/{args.nshards} seed={args.seed} " f"coords_dir={coords_dir} CUDA={os.environ.get('CUDA_VISIBLE_DEVICES')}", flush=True) print("[fin] NOTE: stage-2 `base` is a FRESH torch.manual_seed(seed) draw, " "NOT the reference's post-SS-sampling draw (see the docstring).", flush=True) t0 = time.time() M = baf.Models() print(f"[fin] models ready {time.time() - t0:.1f}s", flush=True) n_ok = n_fail = n_skip = 0 for i, obj in enumerate(objs, 1): f = out / f"{obj}.npz" if f.is_file() and not args.force: n_skip += 1 continue t1 = time.time() try: cz = np.load(coords_dir / f"{obj}.npz") d = seed_from_given_coords(M, obj, exp, inputs_dir, npz_dir, views_tags, args.seed, cz["coords"]) tmp = str(f) + ".tmp.npz" np.savez(tmp, **d, views=np.array(views_tags), seed=args.seed, base_noise_is_fresh_draw=np.array(True)) os.replace(tmp, f) n_ok += 1 print(f"[fin {i}/{len(objs)}] OK {obj} N={len(d['coords'])} " f"vis={int(d['visible'].sum())} {time.time() - t1:.1f}s", flush=True) except Exception: n_fail += 1 traceback.print_exc() print(f"[fin {i}/{len(objs)}] FAIL {obj}", flush=True) torch.cuda.empty_cache() print(f"FINISH_CACHE DONE ok={n_ok} fail={n_fail} skip={n_skip} " f"total={len(objs)}", flush=True) if __name__ == "__main__": main()