worldmem-baseline-evals / LIVE /scripts /evaluate_worldmem.py
BonanDing's picture
Match LIVE to 12 frame slots and gate full evaluation with target GPU preflight
c1ce191 verified
Raw History Blame Contribute Delete
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__,
}
@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()