Download LIVE/scripts/evaluate_worldmem.py from BonanDing/worldmem-baseline-evals: direct link, hf CLI and curl.
- Browser
- Download file 8.19 kB
-
https://huggingface.co/BonanDing/worldmem-baseline-evals/resolve/main/LIVE/scripts/evaluate_worldmem.py
- Command line
-
hf download hf://BonanDing/worldmem-baseline-evals/LIVE/scripts/evaluate_worldmem.py
-
curl -L -o evaluate_worldmem.py https://huggingface.co/BonanDing/worldmem-baseline-evals/resolve/main/LIVE/scripts/evaluate_worldmem.py
8.19 kB
| """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__, | |
| } | |
| 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() | |