#!/usr/bin/env python3 """Evaluate a Causal-Forcing single-block Predictor against FFFF RGB videos. This reuses the established Self-Forcing FPPF rollout and pixel-metric implementation. References are reproduced with the Causal checkpoint because the Stage-1 dataset intentionally omits chunk-0 clean latents. """ from __future__ import annotations import importlib.util import json import os import sys import types from pathlib import Path from typing import Any REPO_ROOT = Path(__file__).resolve().parents[1] SELF_FORCING_ROOT = REPO_ROOT.parent / "Self-Forcing" REFERENCE_SCRIPT = SELF_FORCING_ROOT / "scripts/evaluate_single_block_fppf.py" # Keep Causal-Forcing imports authoritative while allowing the reference # script itself to be loaded. It imports this historical helper name only for # hidden_to_flow, which is identical to our Stage-1 implementation. sys.path.insert(0, str(REPO_ROOT)) if str(SELF_FORCING_ROOT) not in sys.path: sys.path.append(str(SELF_FORCING_ROOT)) from scripts.run_single_block_stage1_sweep import hidden_to_flow compat = types.ModuleType("scripts.run_single_block_init_sweep") compat.hidden_to_flow = hidden_to_flow sys.modules[compat.__name__] = compat spec = importlib.util.spec_from_file_location("self_forcing_fppf_reference", REFERENCE_SCRIPT) if spec is None or spec.loader is None: raise RuntimeError(f"Cannot load {REFERENCE_SCRIPT}") reference = importlib.util.module_from_spec(spec) spec.loader.exec_module(reference) torch = reference.torch OmegaConf = reference.OmegaConf safe_open = reference.safe_open def build_causal_pipeline( config: Any, checkpoint_path: Path, vae: torch.nn.Module, device: torch.device, ): generator = reference.WanDiffusionWrapper( **getattr(config, "model_kwargs", {}), is_causal=True ) pipeline = reference.CausalInferencePipeline( config, device=device, generator=generator, text_encoder=torch.nn.Identity(), vae=vae, ) checkpoint = torch.load( checkpoint_path, map_location="cpu", weights_only=False, mmap=True ) if "generator" not in checkpoint: raise KeyError(f"Causal checkpoint has keys {sorted(checkpoint)}") pipeline.generator.load_state_dict(checkpoint["generator"], strict=True) del checkpoint pipeline.to(dtype=torch.bfloat16) pipeline.generator.to(device=device) pipeline.eval().requires_grad_(False) return pipeline def offline_rollout_latent(dataset_root: Path, prompt_id: int) -> torch.Tensor: path = dataset_root / f"prompt_{prompt_id:04d}" / "trajectory.safetensors" with safe_open(path, framework="pt", device="cpu") as handle: return torch.cat( [ handle.get_tensor(f"chunk_{chunk:02d}_clean_latent") for chunk in range(1, reference.NUM_CHUNKS) ], dim=1, ).contiguous() @torch.inference_mode() def prepare_reference( *, pipeline, vae, dataset_root: Path, output_dir: Path, prompt_id: int, seed: int, device: torch.device, ) -> None: destination = output_dir / "ffff_reference_frames" / f"prompt_{prompt_id:04d}.safetensors" verification_path = output_dir / "ffff_verification" / f"prompt_{prompt_id:04d}.json" if destination.exists() and verification_path.exists(): return latent, counts = reference.generate_rollout( pipeline=pipeline, dataset_root=dataset_root, prompt_id=prompt_id, generation_seed=seed, device=device, predictor=None, source_layer=None, schedule="FFFF", ) expected = offline_rollout_latent(dataset_root, prompt_id).to( device=device, dtype=torch.bfloat16 ) difference = latent[:, reference.FRAMES_PER_CHUNK :].float() - expected.float() verification = { "prompt_id": prompt_id, "compared_latent_chunks": [1, 2, 3, 4, 5, 6], "max_abs_latent_error": float(difference.abs().max()), "latent_mse": float(difference.square().mean()), **counts, } if verification["max_abs_latent_error"] > 1e-3: raise RuntimeError( f"FFFF prompt {prompt_id} does not reproduce offline chunks: {verification}" ) with torch.autocast(device_type="cuda", dtype=torch.bfloat16): pixels = vae.decode_to_pixel(latent, use_cache=False) frames = reference.pixels_to_u8(pixels) reference.atomic_safetensors( destination, {"frames": frames}, { "reference": "reproduced Causal-Forcing FFFF", "prompt_id": str(prompt_id), "range": "uint8_0_255", "layout": "TCHW", }, ) reference.atomic_json(verification_path, verification) if hasattr(vae.model, "clear_cache"): vae.model.clear_cache() del latent, expected, difference, pixels, frames torch.cuda.empty_cache() def main() -> None: args = reference.parse_args() args.config_path = reference.resolve(args.config_path) args.checkpoint_path = reference.resolve(args.checkpoint_path) args.dataset_root = reference.resolve(args.dataset_root) args.sweep_dir = reference.resolve(args.sweep_dir) args.output_dir = reference.resolve(args.output_dir) args.output_dir.mkdir(parents=True, exist_ok=True) prompt_ids = sorted(set(args.prompt_ids)) if args.max_prompts is not None: prompt_ids = prompt_ids[: args.max_prompts] experiments = reference.discover_experiments( args.sweep_dir, args.experiments, args.max_experiments ) if len(experiments) != 1: raise ValueError("This evaluator expects exactly one selected experiment") experiment = experiments[0] device = torch.device("cuda") torch.set_grad_enabled(False) reference.set_seed(args.generation_seed) config = OmegaConf.merge( OmegaConf.load(REPO_ROOT / "configs/default_config.yaml"), OmegaConf.load(args.config_path), ) manifest = { "status": "running", "gpu": str(args.gpu), "config_path": str(args.config_path), "checkpoint_path": str(args.checkpoint_path), "dataset_root": str(args.dataset_root), "sweep_dir": str(args.sweep_dir), "experiment": experiment["name"], "prompt_ids": prompt_ids, "generation_seed_reset_per_prompt": args.generation_seed, "schedule": "chunk0=FFFF; chunks1-6=FPPF", "reference": "reproduced Causal-Forcing FFFF with same prompt and seed", "metrics": { "psnr": "RGB PSNR from global pixel MSE", "ssim": "11x11 Gaussian sigma=1.5 RGB SSIM, frame mean", "lpips": "AlexNet LPIPS on RGB [-1,1], frame mean", "pixel_quantization": "both inputs rounded to uint8", }, } reference.atomic_json(args.output_dir / "manifest.json", manifest) print("[setup] loading VAE, Causal generator, and LPIPS", flush=True) vae = reference.WanVAEWrapper().to( device=device, dtype=torch.bfloat16 ).eval() pipeline = build_causal_pipeline(config, args.checkpoint_path, vae, device) teacher = pipeline.generator.model lpips_model = None if not args.skip_lpips: lpips_model = reference.lpips.LPIPS(net="alex", verbose=False).to(device).eval() lpips_model.requires_grad_(False) predictor = reference.load_predictor(teacher, experiment, device) run_dir = args.output_dir / experiment["name"] run_dir.mkdir(parents=True, exist_ok=True) existing_results = { result["prompt_id"]: result for result in reference.load_completed_prompt_results(run_dir, prompt_ids) } for offset, prompt_id in enumerate(prompt_ids, start=1): if prompt_id in existing_results and ( args.skip_lpips or existing_results[prompt_id].get("lpips") is not None ): print(f"[prompt] {offset}/{len(prompt_ids)} id={prompt_id} cached", flush=True) continue print(f"[reference] {offset}/{len(prompt_ids)} id={prompt_id}", flush=True) prepare_reference( pipeline=pipeline, vae=vae, dataset_root=args.dataset_root, output_dir=args.output_dir, prompt_id=prompt_id, seed=args.generation_seed, device=device, ) started = reference.time.perf_counter() latent, counts = reference.generate_rollout( pipeline=pipeline, dataset_root=args.dataset_root, prompt_id=prompt_id, generation_seed=args.generation_seed, device=device, predictor=predictor, source_layer=experiment["source_layer"], schedule="FPPF", ) with torch.autocast(device_type="cuda", dtype=torch.bfloat16): pixels = vae.decode_to_pixel(latent, use_cache=False) prediction_u8 = reference.pixels_to_u8(pixels) reference_u8 = reference.load_reference_frames(args.output_dir, prompt_id) metrics = reference.frame_metrics( reference_u8=reference_u8, prediction_u8=prediction_u8, lpips_model=lpips_model, batch_size=args.metric_batch_size, device=device, ) result = { "prompt_id": prompt_id, "prompt": reference.load_prompt_metadata(args.dataset_root, prompt_id)["prompt"], **counts, **metrics, "total_time_s": reference.time.perf_counter() - started, } reference.atomic_json( run_dir / "per_prompt" / f"prompt_{prompt_id:04d}.json", result ) existing_results[prompt_id] = result print( f"[prompt] {offset}/{len(prompt_ids)} id={prompt_id} " f"psnr={metrics['psnr']:.4f} ssim={metrics['ssim']:.6f} " f"lpips={metrics['lpips']:.6f}", flush=True, ) if hasattr(vae.model, "clear_cache"): vae.model.clear_cache() del latent, pixels, prediction_u8, reference_u8 torch.cuda.empty_cache() prompt_results = [existing_results[prompt_id] for prompt_id in prompt_ids] aggregate = reference.aggregate_prompt_results(experiment, prompt_results) reference.atomic_json(run_dir / "metrics.json", aggregate) reference.write_summary(args.output_dir, experiments) manifest["status"] = "complete" reference.atomic_json(args.output_dir / "manifest.json", manifest) print( f"[complete] psnr={aggregate['psnr']:.6f} " f"ssim={aggregate['ssim']:.6f} lpips={aggregate['lpips']:.6f}", flush=True, ) if __name__ == "__main__": main()