"""Evaluate official LIVE weights on the shared WorldMem RE10K trajectory.""" import argparse import json import os from pathlib import Path import random import sys import time import numpy as np import torch from omegaconf import OmegaConf REPO = Path(__file__).resolve().parents[1] sys.path.insert(0, str(REPO)) sys.path.insert(0, str(REPO.parent / "shared")) from tvideo.mc.models.live import LIVE class InferenceLIVE(LIVE): def create_evaluation_models(self): # Metrics are computed with the common evaluator, after generation. return torch.nn.ModuleDict(), torch.nn.ModuleDict() def init_from_ckpt(self, path, ignore_keys=(), init_from_oasis_model=False, verbose=False): checkpoint = torch.load(path, map_location="cpu", mmap=True, weights_only=False) expected = self.state_dict() weights = {} ignored = [] for name, value in checkpoint["state_dict"].items(): name = name.replace("._orig_mod.", ".") if name in expected: weights[name] = value elif ".attn_mask_" in name or name.startswith(("train_metrics.", "val_metrics.", "tokenizer.metrics", "tokenizer.perceptual_loss.")): ignored.append(name) else: raise ValueError(f"Unrecognized checkpoint tensor: {name}") self.load_state_dict(weights, strict=True) self.checkpoint_audit = { "checkpoint": str(path), "checkpoint_global_step": checkpoint["global_step"], "weight_policy": "raw state_dict; no EMA tensors in released checkpoint", "loaded_tensors": len(weights), "excluded_metric_or_deterministic_mask_tensors": len(ignored), } class LivePredictor: def __init__(self, checkpoint, vae, device="cuda"): config = OmegaConf.load(REPO / "configs/live_re10k.yaml") params = config.model.params params.ckpt_path = str(checkpoint) params.compile_model = False params.compile_tokenizer = False params.tokenizer_config.params.ckpt_path = str(vae) params.tokenizer_config.params.ckpt_path2 = str(vae) params.tokenizer_config.params.initialize_metrics = False self.model = InferenceLIVE(**params).eval().requires_grad_(False).to(device) self.device = device self.window_frames = 12 self.metadata = { **self.model.checkpoint_audit, "method": "LIVE", "official_commit": "5f0fe658330c64c585874acad3149aa15ba82d1b", "backbone_parameters": sum(p.numel() for p in self.model.model.parameters()), "backbone_parameter_tensors": len(list(self.model.model.parameters())), "vae_parameters": sum(p.numel() for p in self.model.tokenizer.parameters()), "vae_parameter_tensors": len(list(self.model.tokenizer.parameters())), "trainable_parameters": 0, "optimizer_groups": [], "adapters_lora": False, "activation_checkpointing": False, "batch_size": 1, "sampler": "official generate_v_flow_shift_dpm", "denoise_steps": 18, "flow_shift": 3.0, "window_frames": self.window_frames, "training_window_frames": 32, "context_budget": "11 recent frames + 1 target; same total slots as 8 local + 4 memory; no retrieved memory", "chunk_frames": 1, "context_noise": True, "kv_cache": False, "network_precision": "native fp16 autocast", "vae_precision": "fp32, stochastic posterior sample, scale 0.2325", "compile": False, "numerical_failure_policy": "raise instead of replacing non-finite sampler/decoder outputs with zeros", "history_frames": 100, "history_policy": "100 observations supplied; sliding window uses latest 11", "camera_policy": "DFoT w2c; native per-window first-camera Plucker normalization", "protocol_change_from_paper": "12-frame inference window instead of native 32; common stride 1, observed 100, generated 300, closed revisit instead of native stride 2/open rollout", "torch_version": torch.__version__, } @torch.inference_mode() def generate(self, observed, cameras, seed=0): random.seed(seed) np.random.seed(seed % (2 ** 32)) torch.manual_seed(seed) torch.cuda.manual_seed_all(seed) history_frames = observed.shape[0] images = observed.unsqueeze(0).to(self.device).mul(2).sub(1) # Future RGB is never passed to the encoder or sampler. latent = self.model.tokenizer_encode(images) samples = self.model.generate_v_flow_shift_dpm( model=self.model.model, vid=latent, act=cameras.unsqueeze(0).to(self.device), n_prompt_frames=history_frames, total_frames=len(cameras), max_frames=self.window_frames, chunk_size=1, denoise_steps=18, add_noise_to_ctx=True, use_kv_cache=False, use_half_precision=True, desc="WorldMem common RE10K protocol", ) generated = self.model.tokenizer_decode(samples[:, history_frames:]) return generated[0].clamp(-1, 1).add(1).div(2).float().cpu() def main(): from protocol import case_seed, load_re10k_case, save_case, select_cases parser = argparse.ArgumentParser() parser.add_argument("--manifest", required=True) parser.add_argument("--data-root", required=True) parser.add_argument("--output", required=True) parser.add_argument("--checkpoint", default=str(REPO.parent / "checkpoints/live/live-re10k.ckpt")) parser.add_argument("--vae", default=str(REPO.parent / "checkpoints/live/vae-kl16.ckpt")) parser.add_argument("--scope", choices=["monitor8", "final100"], default="final100") parser.add_argument("--seed", type=int, default=0) parser.add_argument("--limit", type=int) parser.add_argument("--max-generated", type=int, default=300) parser.add_argument("--rank", type=int, default=int(os.environ.get("SLURM_PROCID", 0))) parser.add_argument("--world-size", type=int, default=int(os.environ.get("SLURM_NTASKS", 1))) parser.add_argument("--preflight-only", action="store_true") args = parser.parse_args() if not 1 <= args.max_generated <= 300: parser.error("--max-generated must be between 1 and 300") rows = select_cases(args.manifest, args.scope, args.rank, args.world_size, args.limit) predictor = LivePredictor(args.checkpoint, args.vae, device="cpu" if args.preflight_only else "cuda") output = Path(args.output) output.mkdir(parents=True, exist_ok=True) (output / f"live_manifest_rank{args.rank}.json").write_text(json.dumps(predictor.metadata, indent=2) + "\n") print(json.dumps(predictor.metadata, indent=2), flush=True) for row in rows: case = load_re10k_case(row, args.data_root) if args.preflight_only: print(f"Preflight scene {row['scene_id']}: RGB {tuple(case['rgb'].shape)}, cameras {tuple(case['cameras'].shape)}", flush=True) continue history = case["history_frames"] torch.cuda.reset_peak_memory_stats() torch.cuda.synchronize() start = time.monotonic() seed = case_seed(row, args.seed) pred = predictor.generate(case["rgb"][:history], case["cameras"][:history + args.max_generated], seed=seed) torch.cuda.synchronize() metadata = { **predictor.metadata, "seed": seed, "generated_frames": len(pred), "generation_seconds": time.monotonic() - start, "peak_cuda_allocated_bytes": torch.cuda.max_memory_allocated(), "source_indices": case["source_indices"].tolist(), "diagnostic_only": args.max_generated != 300, } save_case(args.output, row, pred, metadata) print(f"Saved {row['scene_id']}: {len(pred)} generated frames, {metadata['generation_seconds']:.1f}s", flush=True) if __name__ == "__main__": main()