Causal-Forcing / scripts /evaluate_causal_single_block_fppf.py
Cccccz's picture
Upload code and configuration only
ae8ade0 verified
Raw History Blame Contribute Delete
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()
@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()