Download scripts/evaluate_causal_single_block_fppf.py from Cccccz/Causal-Forcing: direct link, hf CLI and curl.
- Browser
- Download file 10.7 kB
-
https://huggingface.co/Cccccz/Causal-Forcing/resolve/main/scripts/evaluate_causal_single_block_fppf.py
- Command line
-
hf download hf://Cccccz/Causal-Forcing/scripts/evaluate_causal_single_block_fppf.py
-
curl -L -o evaluate_causal_single_block_fppf.py https://huggingface.co/Cccccz/Causal-Forcing/resolve/main/scripts/evaluate_causal_single_block_fppf.py
10.7 kB
| #!/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() | |
| 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() | |