"""Video-stereo-matcher loader + signed inference for Phase 0. Bypasses pytorch3d/lightning/hydra: builds the core nn.Modules directly with a tiny sys.modules shim. One model per process (each repo needs cwd=its root and has a colliding `models/core` namespace). Usage: python vsm.py --model {dynamic,bida,sav} --ckpt [--selftest] python vsm.py --model ... --ckpt ... --clip --out `--clip` npz holds `video` [T,2,3,H,W] uint8 (RGB). Output npz holds `signed` [T,H,W] float32 = pre-abs signed disparity (positive = in front of screen). """ import argparse, importlib, os, sys, types, typing import numpy as np import torch TP = os.path.join(os.path.dirname(os.path.abspath(__file__)), "third_party") REPO = {"dynamic": "dynamic_stereo", "bida": "BiDAStereo", "sav": "StereoAnyVideo"} def install_pytorch3d_shim(): if "pytorch3d" in sys.modules: return def mod(n): m = types.ModuleType(n); sys.modules[n] = m; return m mod("pytorch3d"); mod("pytorch3d.common") dt = mod("pytorch3d.common.datatypes") dt.get_args, dt.get_origin = typing.get_args, typing.get_origin mod("pytorch3d.implicitron"); mod("pytorch3d.implicitron.tools") cfg = mod("pytorch3d.implicitron.tools.config") cfg.get_default_args = lambda *a, **k: {} cfg.Configurable = type("Configurable", (), {}) cfg.expand_args_fields = lambda c: c cfg.run_auto_creation = lambda s: None def _strip(sd): if isinstance(sd, dict) and "model" in sd: sd = sd["model"] if isinstance(sd, dict) and "state_dict" in sd: sd = sd["state_dict"] out = {} for k, v in sd.items(): for p in ("module.", "model."): if k.startswith(p): k = k[len(p):] out[k] = v return out def build(model, ckpt, device="cuda"): """cd into repo root, build core net, strict-load checkpoint.""" os.chdir(f"{TP}/{REPO[model]}") sys.path.insert(0, TP) install_pytorch3d_shim() if model == "dynamic": from dynamic_stereo.models.core.dynamic_stereo import DynamicStereo net = DynamicStereo(num_frames=5, attention_type="self_stereo_temporal_update_time_update_space", use_3d_update_block=True, different_update_blocks=True) elif model == "bida": # patch the vendored-RAFT wrapper: build arch only (weights come from the # full BiDAStereo checkpoint's raft.model.* keys), no .cuda()/sintel load. rm = importlib.import_module("bidastereo.models.raft_model") raft_mod = rm.raft def _raft_init(self): torch.nn.Module.__init__(self) from types import SimpleNamespace self.args = SimpleNamespace(mixed_precision=False, small=False, dropout=0.0) self.model = raft_mod.RAFT(self.args) rm.RAFTModel.__init__ = _raft_init rm.RAFTModel.__post_init__ = lambda self: None from bidastereo.models.core.bidastereo import BiDAStereo net = BiDAStereo(mixed_precision=False) elif model == "sav": from stereoanyvideo.models.core.stereoanyvideo import StereoAnyVideo net = StereoAnyVideo(mixed_precision=False) else: raise ValueError(model) sd = torch.load(ckpt, map_location="cpu", weights_only=False) net.load_state_dict(_strip(sd), strict=True) net = net.to(device).eval() return net @torch.no_grad() def infer_clip(model, net, video, iters=20, kernel_size=20, device="cuda", max_width=960): """video: torch [T,2,3,H,W] float32 in [0,255]. Returns signed disp [T,H,W] (pre-abs) in ORIGINAL-resolution pixels. Reimplements forward_batch_test's sliding window WITHOUT the final .abs() so we keep the sign (sign convention resolved later against GT). Downscales to max_width for tractable all-pairs correlation, then upscales disparity and rescales by the width factor.""" import torch.nn.functional as F from importlib import import_module pad_mod = {"dynamic": "dynamic_stereo.models.core.utils.utils", "bida": "bidastereo.models.core.utils.utils", "sav": "stereoanyvideo.models.core.utils.utils"}[model] InputPadder = import_module(pad_mod).InputPadder T = video.shape[0] H0, W0 = video.shape[-2:] disp_scale = 1.0 if max_width and W0 > max_width: s = max_width / W0 newW = max(32, int(round(W0 * s)) // 32 * 32) newH = max(32, int(round(H0 * s)) // 32 * 32) video = F.interpolate(video.reshape(T * 2, 3, H0, W0), size=(newH, newW), mode="bilinear", align_corners=True).reshape(T, 2, 3, newH, newW) disp_scale = W0 / newW def run_window(L, R): padder = InputPadder(L.shape, divis_by=32) Lp, Rp = padder.pad(L, R) disp = net(Lp[None].to(device), Rp[None].to(device), iters=iters, test_mode=True) return padder.unpad(disp[:, 0]).cpu() # [t,1,H,W] -> [t,H,W] after [:,0]? -> [t,H,W] if T <= kernel_size: out = run_window(video[:, 0], video[:, 1]) # single window covers all frames else: stride = kernel_size // 2 buf = [None] * T for i in range(0, T, stride): j = min(i + kernel_size, T) d = run_window(video[i:j, 0], video[i:j, 1]) for k in range(i, j): # last write wins (centered frames) buf[k] = d[k - i] if j == T: break out = torch.stack(buf) signed = out.squeeze(1) if out.dim() == 4 else out # -> [T,h,w] (no abs) if disp_scale != 1.0: # back to original resolution + scale signed = F.interpolate(signed[:, None], size=(H0, W0), mode="bilinear", align_corners=True)[:, 0] * disp_scale return signed.float().numpy() def main(): ap = argparse.ArgumentParser() ap.add_argument("--model", required=True, choices=list(REPO)) ap.add_argument("--ckpt", required=True) ap.add_argument("--selftest", action="store_true") ap.add_argument("--clip"); ap.add_argument("--out") ap.add_argument("--iters", type=int, default=20) a = ap.parse_args() # resolve all user paths to absolute BEFORE build() chdir's into the repo root a.ckpt = os.path.abspath(a.ckpt) if a.clip: a.clip = os.path.abspath(a.clip) if a.out: a.out = os.path.abspath(a.out) net = build(a.model, a.ckpt) n = sum(p.numel() for p in net.parameters()) print(f"[{a.model}] built + STRICT-loaded OK: {n/1e6:.1f}M params on cuda", flush=True) if a.clip: d = np.load(a.clip) video = torch.from_numpy(d["video"].astype("float32")) signed = infer_clip(a.model, net, video, iters=a.iters) np.savez_compressed(a.out, signed=signed) print(f"[{a.model}] clip {tuple(video.shape)} -> signed {signed.shape} " f"range [{signed.min():.1f},{signed.max():.1f}] neg%={100*(signed<-0.5).mean():.1f}", flush=True) if __name__ == "__main__": main()