bert_simpson / forgebench /code /ours /finish_cache.py
Ronaldo-GOAT's picture
Add files using upload-large-folder tool
4c3d957 verified
Raw History Blame Contribute Delete
11.6 kB
#!/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 ``<coords-dir>/<obj>.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_<tag> (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 <pass2 out> --out <cache>/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 (<obj>.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()