Self-Forcing / scripts /evaluate_single_block_fppf.py
Cccccz's picture
Upload Python scripts
bc29ee3 verified
Raw
History Blame Contribute Delete
42.4 kB
#!/usr/bin/env python3
"""Evaluate trained one-block Predictors with an FPPF rollout against FFFF.
The evaluation uses the held-out offline prompt shards. FFFF clean latents
are decoded once and cached as 8-bit RGB reference frames. For every trained
Predictor, chunk 0 is generated with FFFF (there is no previous chunk), while
chunks 1..6 use Full-Predictor-Predictor-Full. PSNR, Gaussian SSIM, and
AlexNet LPIPS are computed frame by frame against the matching FFFF video.
"""
from __future__ import annotations
import argparse
import csv
import json
import math
import os
import sys
import time
from pathlib import Path
from typing import Any
def _preparse_gpu() -> str:
parser = argparse.ArgumentParser(add_help=False)
parser.add_argument("--gpu", default="2")
args, _ = parser.parse_known_args()
os.environ["CUDA_VISIBLE_DEVICES"] = str(args.gpu)
return str(args.gpu)
PHYSICAL_GPU = _preparse_gpu()
import lpips
import torch
import torch.nn.functional as F
from omegaconf import OmegaConf
from safetensors import safe_open
from safetensors.torch import load_file, save_file
from torchvision.io import write_video
REPO_ROOT = Path(__file__).resolve().parents[1]
if str(REPO_ROOT) not in sys.path:
sys.path.insert(0, str(REPO_ROOT))
from pipeline import CausalInferencePipeline
from predictor_training.offline_data import TOKENS_PER_CHUNK
from predictor_training.single_block import (
SingleBlockPredictor,
initialize_predictor_block,
)
from scripts.run_single_block_init_sweep import hidden_to_flow
from utils.misc import set_seed
from utils.wan_wrapper import WanDiffusionWrapper, WanVAEWrapper
from wan.modules.model import sinusoidal_embedding_1d
LATENT_CHANNELS = 16
LATENT_HEIGHT = 60
LATENT_WIDTH = 104
FRAMES_PER_CHUNK = 3
NUM_CHUNKS = 7
NUM_DENOISING_STEPS = 4
PIXEL_FRAMES_FIRST_CHUNK = 1 + 4 * (FRAMES_PER_CHUNK - 1)
DEFAULT_PROMPT_IDS = list(range(80, 100))
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--gpu", default=PHYSICAL_GPU)
parser.add_argument(
"--config_path",
type=Path,
default=Path("configs/self_forcing_sid.yaml"),
)
parser.add_argument(
"--checkpoint_path",
type=Path,
default=Path("checkpoints/self_forcing_dmd.pt"),
)
parser.add_argument(
"--dataset_root",
type=Path,
default=Path("outputs/predictor_offline_100_all_blocks"),
)
parser.add_argument(
"--sweep_dir",
type=Path,
default=Path("outputs/single_block_init_sweep"),
)
parser.add_argument(
"--output_dir",
type=Path,
default=Path("outputs/single_block_fppf_eval"),
)
parser.add_argument(
"--prompt_ids", type=int, nargs="*", default=DEFAULT_PROMPT_IDS
)
parser.add_argument(
"--experiments",
nargs="*",
default=None,
help="Experiment directory names. Omit to evaluate all summary rows.",
)
parser.add_argument("--max_prompts", type=int, default=None)
parser.add_argument("--max_experiments", type=int, default=None)
parser.add_argument("--metric_batch_size", type=int, default=4)
parser.add_argument("--generation_seed", type=int, default=0)
parser.add_argument(
"--verify_ffff",
action=argparse.BooleanOptionalAction,
default=True,
help="Re-run FFFF once and compare its latent exactly to offline data.",
)
parser.add_argument(
"--skip_lpips",
action=argparse.BooleanOptionalAction,
default=False,
help="Only for quick diagnostics; formal evaluation should keep LPIPS.",
)
parser.add_argument(
"--rebuild_references",
action=argparse.BooleanOptionalAction,
default=False,
)
parser.add_argument(
"--save_videos",
action=argparse.BooleanOptionalAction,
default=False,
help="Save each FPPF prediction as a 16-fps H.264 MP4 for VBench.",
)
args = parser.parse_args()
if args.metric_batch_size < 1:
parser.error("--metric_batch_size must be positive")
if not args.prompt_ids:
parser.error("At least one prompt ID is required")
if any(value < 0 or value >= 100 for value in args.prompt_ids):
parser.error("Prompt IDs must be in [0, 99]")
return args
def resolve(path: Path) -> Path:
path = path.expanduser()
return path.resolve() if path.is_absolute() else (REPO_ROOT / path).resolve()
def atomic_json(path: Path, value: Any) -> None:
path.parent.mkdir(parents=True, exist_ok=True)
temporary = path.with_suffix(path.suffix + ".tmp")
temporary.write_text(
json.dumps(value, indent=2, ensure_ascii=False, allow_nan=True) + "\n",
encoding="utf-8",
)
os.replace(temporary, path)
def atomic_safetensors(
path: Path, tensors: dict[str, torch.Tensor], metadata: dict[str, str]
) -> None:
path.parent.mkdir(parents=True, exist_ok=True)
temporary = path.with_suffix(path.suffix + ".tmp")
save_file(tensors, temporary, metadata=metadata)
os.replace(temporary, path)
def load_prompt_metadata(dataset_root: Path, prompt_id: int) -> dict[str, Any]:
path = dataset_root / f"prompt_{prompt_id:04d}" / "metadata.json"
return json.loads(path.read_text(encoding="utf-8"))
def load_ffff_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:
chunks = [
handle.get_tensor(f"chunk_{chunk:02d}_clean_latent")
for chunk in range(NUM_CHUNKS)
]
return torch.cat(chunks, dim=1).contiguous()
def pixels_to_u8(video: torch.Tensor) -> torch.Tensor:
"""Convert [1,T,3,H,W] pixels in [-1,1] to CPU uint8 frames."""
return (
((video.squeeze(0).float() + 1.0) * 127.5)
.round_()
.clamp_(0, 255)
.to(device="cpu", dtype=torch.uint8)
.contiguous()
)
def save_mp4(frames: torch.Tensor, path: Path) -> None:
path.parent.mkdir(parents=True, exist_ok=True)
write_video(
str(path), frames.permute(0, 2, 3, 1), fps=16,
video_codec="libx264", options={"crf": "18"},
)
@torch.inference_mode()
def prepare_reference_frames(
*,
vae: WanVAEWrapper,
dataset_root: Path,
output_dir: Path,
prompt_ids: list[int],
device: torch.device,
rebuild: bool,
) -> None:
reference_dir = output_dir / "ffff_reference_frames"
reference_dir.mkdir(parents=True, exist_ok=True)
for offset, prompt_id in enumerate(prompt_ids, start=1):
destination = reference_dir / f"prompt_{prompt_id:04d}.safetensors"
if destination.exists() and not rebuild:
print(
f"[reference] {offset}/{len(prompt_ids)} prompt={prompt_id} cached",
flush=True,
)
continue
latent = load_ffff_latent(dataset_root, prompt_id).to(
device=device, dtype=torch.bfloat16
)
started = time.perf_counter()
with torch.autocast(device_type="cuda", dtype=torch.bfloat16):
pixels = vae.decode_to_pixel(latent, use_cache=False)
frames = pixels_to_u8(pixels)
atomic_safetensors(
destination,
{"frames": frames},
{
"reference": "FFFF",
"prompt_id": str(prompt_id),
"range": "uint8_0_255",
"layout": "TCHW",
},
)
if hasattr(vae.model, "clear_cache"):
vae.model.clear_cache()
del latent, pixels, frames
torch.cuda.empty_cache()
print(
f"[reference] {offset}/{len(prompt_ids)} prompt={prompt_id} "
f"decoded={time.perf_counter() - started:.1f}s",
flush=True,
)
def load_reference_frames(output_dir: Path, prompt_id: int) -> torch.Tensor:
path = output_dir / "ffff_reference_frames" / f"prompt_{prompt_id:04d}.safetensors"
with safe_open(path, framework="pt", device="cpu") as handle:
return handle.get_tensor("frames")
def discover_experiments(
sweep_dir: Path,
requested: list[str] | None,
max_experiments: int | None,
) -> list[dict[str, Any]]:
summary_path = sweep_dir / "summary.csv"
with summary_path.open("r", encoding="utf-8", newline="") as handle:
rows = list(csv.DictReader(handle))
by_name = {row["name"]: row for row in rows}
names = list(by_name) if requested is None else requested
unknown = [name for name in names if name not in by_name]
if unknown:
raise KeyError(f"Unknown sweep experiments: {unknown}")
if max_experiments is not None:
names = names[:max_experiments]
experiments = []
for name in names:
run_dir = sweep_dir / name
config = json.loads((run_dir / "config.json").read_text(encoding="utf-8"))
weights = run_dir / "predictor_final.safetensors"
if not weights.exists():
raise FileNotFoundError(weights)
row = by_name[name]
experiment = {
"name": name,
"initialization_method": config["initialization_method"],
"source_layer": int(config["source_layer"]),
"weights": weights,
"gate_mode": config.get("gate_mode", "baseline"),
"gate_hidden_dim": int(config.get("gate_hidden_dim", 128)),
"gate_initial_bias": float(
config.get("gate_initial_bias", 4.6)
),
"gate_floor": float(config.get("gate_floor", 0.0)),
"constant_gate": float(config.get("constant_gate", 1.0)),
"gate_override": config.get("gate_override"),
"offline_final_val_flow_mse": float(row["final_val_flow_mse"]),
"offline_final_val_hidden_mse": float(row["final_val_hidden_mse"]),
}
experiments.append(experiment)
if experiment["gate_mode"] == "learned":
training_metrics = json.loads(
(run_dir / "metrics.json").read_text(encoding="utf-8")
)
gate_mean = float(training_metrics["evaluations"][-1]["gate_mean"])
experiments.append(
{
**experiment,
"name": f"{name}_constant_mean",
"gate_override": gate_mean,
"constant_gate": gate_mean,
}
)
return experiments
def build_pipeline(
config: Any,
checkpoint_path: Path,
vae: WanVAEWrapper,
device: torch.device,
) -> CausalInferencePipeline:
generator = WanDiffusionWrapper(
**getattr(config, "model_kwargs", {}), is_causal=True
)
pipeline = 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 set(checkpoint) != {"generator_ema"}:
raise KeyError(f"Unexpected Teacher checkpoint keys: {sorted(checkpoint)}")
pipeline.generator.load_state_dict(checkpoint["generator_ema"], strict=True)
del checkpoint
pipeline.to(dtype=torch.bfloat16)
pipeline.generator.to(device=device)
pipeline.eval()
pipeline.generator.requires_grad_(False)
return pipeline
def reset_kv_and_load_cross_cache(
pipeline: CausalInferencePipeline,
dataset_root: Path,
prompt_id: int,
device: torch.device,
) -> None:
if pipeline.kv_cache1 is None:
pipeline._initialize_kv_cache(1, torch.bfloat16, device)
pipeline._initialize_crossattn_cache(1, torch.bfloat16, device)
for cache in pipeline.kv_cache1:
cache["global_end_index"].zero_()
cache["local_end_index"].zero_()
cross_path = (
dataset_root / f"prompt_{prompt_id:04d}" / "cross_attention.safetensors"
)
with safe_open(cross_path, framework="pt", device="cpu") as handle:
for layer, cache in enumerate(pipeline.crossattn_cache):
cache["k"] = handle.get_tensor(f"block_{layer:02d}_k").to(
device=device, dtype=torch.bfloat16
)
cache["v"] = handle.get_tensor(f"block_{layer:02d}_v").to(
device=device, dtype=torch.bfloat16
)
cache["is_init"] = True
class FinalHiddenCapture:
def __init__(self, teacher: torch.nn.Module) -> None:
self.enabled = False
self.value: torch.Tensor | None = None
self.handle = teacher.head.register_forward_pre_hook(self._hook)
def close(self) -> None:
self.handle.remove()
def _hook(
self, _module: torch.nn.Module, inputs: tuple[torch.Tensor, ...]
) -> None:
if self.enabled:
if self.value is not None:
raise RuntimeError("Teacher head was called twice in one Full step")
self.value = inputs[0].detach()
def start(self) -> None:
self.value = None
self.enabled = True
def finish(self) -> torch.Tensor:
self.enabled = False
if self.value is None:
raise RuntimeError("Teacher final hidden was not captured")
value = self.value
self.value = None
return value
def load_predictor(
teacher: torch.nn.Module,
experiment: dict[str, Any],
device: torch.device,
) -> SingleBlockPredictor:
source_layer = int(experiment["source_layer"])
block = initialize_predictor_block(
teacher.blocks[source_layer], "teacher_full"
)
predictor = SingleBlockPredictor(
block=block,
dim=teacher.dim,
gradient_checkpointing=False,
input_variant=experiment.get("input_variant", "self_forcing"),
gate_mode=experiment.get("gate_mode", "baseline"),
gate_hidden_dim=int(experiment.get("gate_hidden_dim", 128)),
gate_initial_bias=float(experiment.get("gate_initial_bias", 4.6)),
gate_floor=float(experiment.get("gate_floor", 0.0)),
constant_gate=float(experiment.get("constant_gate", 1.0)),
atc_previous_scope=experiment.get("atc_previous_scope", "chunk"),
atc_freq_dim=int(experiment.get("atc_freq_dim", 256)),
atc_mlp_hidden_dim=int(experiment.get("atc_mlp_hidden_dim", 3072)),
atc_gate_hidden_dim=int(experiment.get("atc_gate_hidden_dim", 512)),
atc_transport_residual_scale=float(
experiment.get("atc_transport_residual_scale", 0.1)
),
atc_gate_initial_probability=float(
experiment.get("atc_gate_initial_probability", 0.3)
),
atc_collect_diagnostics=bool(
experiment.get("atc_collect_diagnostics", False)
),
)
state = load_file(str(experiment["weights"]), device="cpu")
predictor.load_state_dict(state, strict=True)
if experiment.get("gate_override") is not None:
predictor.fusion.gate_override = float(experiment["gate_override"])
predictor.to(device=device)
predictor.eval().requires_grad_(False)
return predictor
@torch.inference_mode()
def predictor_step(
*,
predictor: SingleBlockPredictor,
teacher: torch.nn.Module,
noisy_input: torch.Tensor,
timestep: torch.Tensor,
anchor_hidden: torch.Tensor,
previous_hidden: torch.Tensor,
history_cache: dict[str, torch.Tensor],
cross_cache: dict[str, torch.Tensor],
current_start: int,
anchor_timestep: torch.Tensor | None = None,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
with torch.autocast(device_type="cuda", dtype=torch.bfloat16):
current_tokens = teacher.patch_embedding(
noisy_input.permute(0, 2, 1, 3, 4)
).flatten(2).transpose(1, 2)
time_embedding = teacher.time_embedding(
sinusoidal_embedding_1d(
teacher.freq_dim, timestep.flatten()
).type_as(current_tokens)
)
timestep_modulation = teacher.time_projection(
time_embedding
).unflatten(1, (6, teacher.dim)).unflatten(
dim=0, sizes=timestep.shape
)
head_embedding = time_embedding.unflatten(
dim=0, sizes=timestep.shape
).unsqueeze(2)
condition_per_frame = time_embedding.unflatten(
dim=0, sizes=timestep.shape
)
condition_tokens = (
condition_per_frame[:, :, None, :]
.expand(
timestep.shape[0],
timestep.shape[1],
30 * 52,
teacher.dim,
)
.reshape(timestep.shape[0], -1, teacher.dim)
)
anchor_distance = None
if predictor.input_variant == "atc":
if anchor_timestep is None:
raise ValueError("ATC inference requires anchor_timestep")
anchor_distance = (
timestep.float() - anchor_timestep.float()
).abs().mean(dim=1)
grid_sizes = torch.tensor(
[[FRAMES_PER_CHUNK, 30, 52]], dtype=torch.long, device="cpu"
)
history_length = int(history_cache["local_end_index"].item())
pred_hidden = predictor(
current_tokens=current_tokens,
anchor_hidden=anchor_hidden,
previous_hidden=previous_hidden,
timestep_modulation=timestep_modulation,
grid_sizes=grid_sizes,
freqs=teacher.freqs,
history_k=history_cache["k"][:, :history_length],
history_v=history_cache["v"][:, :history_length],
cross_k=cross_cache["k"],
cross_v=cross_cache["v"],
current_start=current_start,
condition_tokens=condition_tokens,
anchor_distance=anchor_distance,
)
pred_flow = hidden_to_flow(
pred_hidden, head_embedding, grid_sizes, teacher
)
return pred_hidden, pred_flow, current_tokens
@torch.inference_mode()
def generate_rollout(
*,
pipeline: CausalInferencePipeline,
dataset_root: Path,
prompt_id: int,
generation_seed: int,
device: torch.device,
predictor: SingleBlockPredictor | None,
source_layer: int | None,
schedule: str,
) -> tuple[torch.Tensor, dict[str, float | int]]:
if schedule not in {"FFFF", "FPPF"}:
raise ValueError(schedule)
if schedule == "FPPF" and (predictor is None or source_layer is None):
raise ValueError("FPPF requires a Predictor and source layer")
reset_kv_and_load_cross_cache(pipeline, dataset_root, prompt_id, device)
set_seed(generation_seed)
noise = torch.randn(
1,
NUM_CHUNKS * FRAMES_PER_CHUNK,
LATENT_CHANNELS,
LATENT_HEIGHT,
LATENT_WIDTH,
dtype=torch.bfloat16,
device=device,
)
teacher = pipeline.generator.model
text_dim = int(teacher.text_embedding[0].in_features)
conditional_dict = {
"prompt_embeds": torch.zeros(
1, 1, text_dim, dtype=torch.bfloat16, device=device
)
}
timesteps = pipeline.denoising_step_list.to(device=device)
output_chunks: list[torch.Tensor] = []
previous_chunk_hidden: list[torch.Tensor | None] | None = None
capture = FinalHiddenCapture(teacher)
full_calls = 0
predictor_calls = 0
started = time.perf_counter()
try:
for chunk in range(NUM_CHUNKS):
noisy_input = noise[
:, chunk * FRAMES_PER_CHUNK : (chunk + 1) * FRAMES_PER_CHUNK
]
current_hidden: list[torch.Tensor | None] = [None] * NUM_DENOISING_STEPS
denoised_pred: torch.Tensor | None = None
timestep: torch.Tensor | None = None
for step, current_timestep in enumerate(timesteps):
timestep = torch.ones(
[1, FRAMES_PER_CHUNK], dtype=torch.int64, device=device
) * current_timestep
use_predictor = schedule == "FPPF" and chunk > 0 and step in {1, 2}
if use_predictor:
anchor_hidden = current_hidden[step - 1]
assert anchor_hidden is not None
assert previous_chunk_hidden is not None
previous_hidden = previous_chunk_hidden[step]
assert previous_hidden is not None
history = pipeline.kv_cache1[int(source_layer)]
cross = pipeline.crossattn_cache[int(source_layer)]
pred_hidden, flow, _ = predictor_step(
predictor=predictor,
teacher=teacher,
noisy_input=noisy_input,
timestep=timestep,
anchor_hidden=anchor_hidden,
previous_hidden=previous_hidden,
history_cache=history,
cross_cache=cross,
current_start=chunk * TOKENS_PER_CHUNK,
anchor_timestep=(
torch.ones_like(timestep) * timesteps[step - 1]
),
)
denoised_pred = pipeline.generator._convert_flow_pred_to_x0(
flow_pred=flow.flatten(0, 1),
xt=noisy_input.flatten(0, 1),
timestep=timestep.flatten(0, 1),
).unflatten(0, flow.shape[:2])
current_hidden[step] = pred_hidden
predictor_calls += 1
else:
capture.start()
_, denoised_pred = pipeline.generator(
noisy_image_or_video=noisy_input,
conditional_dict=conditional_dict,
timestep=timestep,
kv_cache=pipeline.kv_cache1,
crossattn_cache=pipeline.crossattn_cache,
current_start=chunk * TOKENS_PER_CHUNK,
)
current_hidden[step] = capture.finish()
full_calls += 1
if step < NUM_DENOISING_STEPS - 1:
next_timestep = timesteps[step + 1]
denoised_flat = denoised_pred.flatten(0, 1)
noisy_input = pipeline.scheduler.add_noise(
denoised_flat,
torch.randn_like(denoised_flat),
next_timestep
* torch.ones(
[FRAMES_PER_CHUNK], dtype=torch.long, device=device
),
).unflatten(0, denoised_pred.shape[:2])
if denoised_pred is None or timestep is None:
raise RuntimeError("Denoising loop produced no clean latent")
output_chunks.append(denoised_pred)
context_timestep = torch.ones_like(timestep) * pipeline.args.context_noise
pipeline.generator(
noisy_image_or_video=denoised_pred,
conditional_dict=conditional_dict,
timestep=context_timestep,
kv_cache=pipeline.kv_cache1,
crossattn_cache=pipeline.crossattn_cache,
current_start=chunk * TOKENS_PER_CHUNK,
)
previous_chunk_hidden = current_hidden
finally:
capture.close()
torch.cuda.synchronize()
return torch.cat(output_chunks, dim=1), {
"generation_time_s": time.perf_counter() - started,
"full_calls": full_calls,
"predictor_calls": predictor_calls,
}
def gaussian_kernel(
device: torch.device, dtype: torch.dtype, channels: int = 3
) -> torch.Tensor:
coordinates = torch.arange(11, device=device, dtype=dtype) - 5
kernel_1d = torch.exp(-(coordinates.square()) / (2 * 1.5**2))
kernel_1d /= kernel_1d.sum()
kernel_2d = torch.outer(kernel_1d, kernel_1d)
return kernel_2d.expand(channels, 1, 11, 11).contiguous()
def ssim_per_frame(
reference: torch.Tensor, prediction: torch.Tensor, kernel: torch.Tensor
) -> torch.Tensor:
channels = reference.shape[1]
mu_x = F.conv2d(reference, kernel, groups=channels)
mu_y = F.conv2d(prediction, kernel, groups=channels)
mu_x2 = mu_x.square()
mu_y2 = mu_y.square()
mu_xy = mu_x * mu_y
sigma_x2 = F.conv2d(reference.square(), kernel, groups=channels) - mu_x2
sigma_y2 = F.conv2d(prediction.square(), kernel, groups=channels) - mu_y2
sigma_xy = F.conv2d(reference * prediction, kernel, groups=channels) - mu_xy
c1 = 0.01**2
c2 = 0.03**2
score = ((2 * mu_xy + c1) * (2 * sigma_xy + c2)) / (
(mu_x2 + mu_y2 + c1) * (sigma_x2 + sigma_y2 + c2)
)
return score.mean(dim=(1, 2, 3))
@torch.inference_mode()
def frame_metrics(
*,
reference_u8: torch.Tensor,
prediction_u8: torch.Tensor,
lpips_model: torch.nn.Module | None,
batch_size: int,
device: torch.device,
) -> dict[str, Any]:
if reference_u8.shape != prediction_u8.shape:
raise ValueError(
f"Reference/prediction shapes differ: {reference_u8.shape}, "
f"{prediction_u8.shape}"
)
kernel = gaussian_kernel(device, torch.float32)
psnr_values: list[float] = []
mse_values: list[float] = []
ssim_values: list[float] = []
lpips_values: list[float] = []
for start in range(0, reference_u8.shape[0], batch_size):
end = min(start + batch_size, reference_u8.shape[0])
reference = reference_u8[start:end].to(
device=device, dtype=torch.float32
) / 255.0
prediction = prediction_u8[start:end].to(
device=device, dtype=torch.float32
) / 255.0
mse = (reference - prediction).square().mean(dim=(1, 2, 3))
psnr = -10.0 * torch.log10(mse.clamp_min(1e-12))
ssim = ssim_per_frame(reference, prediction, kernel)
mse_values.extend(float(value) for value in mse.cpu())
psnr_values.extend(float(value) for value in psnr.cpu())
ssim_values.extend(float(value) for value in ssim.cpu())
if lpips_model is not None:
distance = lpips_model(
reference.mul(2).sub(1), prediction.mul(2).sub(1)
).flatten()
lpips_values.extend(float(value) for value in distance.cpu())
del reference, prediction, mse, psnr, ssim
global_mse = sum(mse_values) / len(mse_values)
rollout_mse = sum(mse_values[PIXEL_FRAMES_FIRST_CHUNK:]) / len(
mse_values[PIXEL_FRAMES_FIRST_CHUNK:]
)
return {
"mse_per_frame": mse_values,
"psnr_per_frame": psnr_values,
"ssim_per_frame": ssim_values,
"lpips_per_frame": lpips_values,
"pixel_mse": global_mse,
"psnr": -10.0 * math.log10(max(global_mse, 1e-12)),
"psnr_frame_mean": sum(psnr_values) / len(psnr_values),
"ssim": sum(ssim_values) / len(ssim_values),
"lpips": (
sum(lpips_values) / len(lpips_values)
if lpips_values
else None
),
"rollout_start_frame": PIXEL_FRAMES_FIRST_CHUNK,
"rollout_pixel_mse": rollout_mse,
"rollout_psnr": -10.0 * math.log10(max(rollout_mse, 1e-12)),
"rollout_ssim": sum(ssim_values[PIXEL_FRAMES_FIRST_CHUNK:])
/ len(ssim_values[PIXEL_FRAMES_FIRST_CHUNK:]),
"rollout_lpips": (
sum(lpips_values[PIXEL_FRAMES_FIRST_CHUNK:])
/ len(lpips_values[PIXEL_FRAMES_FIRST_CHUNK:])
if lpips_values
else None
),
"num_frames": len(psnr_values),
}
def mean_std(values: list[float]) -> tuple[float, float]:
mean = sum(values) / len(values)
variance = sum((value - mean) ** 2 for value in values) / len(values)
return mean, math.sqrt(variance)
def aggregate_prompt_results(
experiment: dict[str, Any], prompt_results: list[dict[str, Any]]
) -> dict[str, Any]:
mse_frames = [
value
for result in prompt_results
for value in result["mse_per_frame"]
]
psnr_frames = [
value
for result in prompt_results
for value in result["psnr_per_frame"]
]
ssim_frames = [
value
for result in prompt_results
for value in result["ssim_per_frame"]
]
lpips_frames = [
value
for result in prompt_results
for value in result["lpips_per_frame"]
]
rollout_mse_frames = [
value
for result in prompt_results
for value in result["mse_per_frame"][PIXEL_FRAMES_FIRST_CHUNK:]
]
rollout_ssim_frames = [
value
for result in prompt_results
for value in result["ssim_per_frame"][PIXEL_FRAMES_FIRST_CHUNK:]
]
rollout_lpips_frames = [
value
for result in prompt_results
for value in result["lpips_per_frame"][PIXEL_FRAMES_FIRST_CHUNK:]
]
pixel_mse = sum(mse_frames) / len(mse_frames)
psnr_frame_mean, psnr_std = mean_std(psnr_frames)
ssim, ssim_std = mean_std(ssim_frames)
if lpips_frames:
lpips_mean, lpips_std = mean_std(lpips_frames)
else:
lpips_mean, lpips_std = None, None
rollout_pixel_mse = sum(rollout_mse_frames) / len(rollout_mse_frames)
rollout_ssim = sum(rollout_ssim_frames) / len(rollout_ssim_frames)
rollout_lpips = (
sum(rollout_lpips_frames) / len(rollout_lpips_frames)
if rollout_lpips_frames
else None
)
return {
"status": "complete",
"name": experiment["name"],
"initialization_method": experiment["initialization_method"],
"source_layer": experiment["source_layer"],
"offline_final_val_flow_mse": experiment["offline_final_val_flow_mse"],
"offline_final_val_hidden_mse": experiment[
"offline_final_val_hidden_mse"
],
"schedule": "chunk0=FFFF; chunks1-6=FPPF",
"reference": "matching FFFF, same prompt and seed",
"pixel_domain": "VAE-decoded RGB, rounded to uint8",
"aggregation": (
"PSNR from global pixel MSE; SSIM/LPIPS mean over decoded frames"
),
"num_prompts": len(prompt_results),
"num_frames": len(psnr_frames),
"pixel_mse": pixel_mse,
"psnr": -10.0 * math.log10(max(pixel_mse, 1e-12)),
"psnr_frame_mean": psnr_frame_mean,
"psnr_frame_std": psnr_std,
"ssim": ssim,
"ssim_frame_std": ssim_std,
"lpips": lpips_mean,
"lpips_frame_std": lpips_std,
"rollout_start_frame": PIXEL_FRAMES_FIRST_CHUNK,
"rollout_num_frames": len(rollout_mse_frames),
"rollout_pixel_mse": rollout_pixel_mse,
"rollout_psnr": -10.0
* math.log10(max(rollout_pixel_mse, 1e-12)),
"rollout_ssim": rollout_ssim,
"rollout_lpips": rollout_lpips,
"mean_generation_time_s": sum(
result["generation_time_s"] for result in prompt_results
)
/ len(prompt_results),
"full_calls_per_prompt": prompt_results[0]["full_calls"],
"predictor_calls_per_prompt": prompt_results[0]["predictor_calls"],
"prompt_ids": [result["prompt_id"] for result in prompt_results],
}
def write_summary(
output_dir: Path, experiments: list[dict[str, Any]]
) -> None:
rows: list[dict[str, Any]] = []
for experiment in experiments:
path = output_dir / experiment["name"] / "metrics.json"
if not path.exists():
continue
metrics = json.loads(path.read_text(encoding="utf-8"))
if metrics.get("status") != "complete":
continue
rows.append(
{
"name": metrics["name"],
"initialization_method": metrics["initialization_method"],
"source_layer": metrics["source_layer"],
"num_prompts": metrics["num_prompts"],
"num_frames": metrics["num_frames"],
"psnr": metrics["psnr"],
"ssim": metrics["ssim"],
"lpips": metrics["lpips"],
"rollout_psnr": metrics["rollout_psnr"],
"rollout_ssim": metrics["rollout_ssim"],
"rollout_lpips": metrics["rollout_lpips"],
"offline_final_val_flow_mse": metrics[
"offline_final_val_flow_mse"
],
"mean_generation_time_s": metrics["mean_generation_time_s"],
}
)
rows.sort(key=lambda row: float(row["lpips"] or math.inf))
fields = [
"name",
"initialization_method",
"source_layer",
"num_prompts",
"num_frames",
"psnr",
"ssim",
"lpips",
"rollout_psnr",
"rollout_ssim",
"rollout_lpips",
"offline_final_val_flow_mse",
"mean_generation_time_s",
]
destination = output_dir / "summary.csv"
temporary = destination.with_suffix(".csv.tmp")
with temporary.open("w", encoding="utf-8", newline="") as handle:
writer = csv.DictWriter(handle, fieldnames=fields)
writer.writeheader()
writer.writerows(rows)
os.replace(temporary, destination)
def load_completed_prompt_results(
run_dir: Path, prompt_ids: list[int]
) -> list[dict[str, Any]]:
results = []
for prompt_id in prompt_ids:
path = run_dir / "per_prompt" / f"prompt_{prompt_id:04d}.json"
if path.exists():
results.append(json.loads(path.read_text(encoding="utf-8")))
return results
def main() -> None:
args = parse_args()
args.config_path = resolve(args.config_path)
args.checkpoint_path = resolve(args.checkpoint_path)
args.dataset_root = resolve(args.dataset_root)
args.sweep_dir = resolve(args.sweep_dir)
args.output_dir = 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 = discover_experiments(
args.sweep_dir, args.experiments, args.max_experiments
)
device = torch.device("cuda")
torch.set_grad_enabled(False)
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),
"prompt_ids": prompt_ids,
"generation_seed_reset_per_prompt": args.generation_seed,
"experiments": [item["name"] for item in experiments],
"fppf_definition": "chunk0=FFFF; chunks1-6=FPPF",
"reference": "offline FFFF clean latents from the same prompt/seed",
"metrics": {
"psnr": "RGB PSNR from global pixel MSE",
"ssim": "11x11 Gaussian sigma=1.5 RGB SSIM, then frame mean",
"lpips": "AlexNet LPIPS on RGB [-1,1], then frame mean",
"pixel_quantization": "both inputs rounded to uint8",
"rollout_only": (
"also reported for decoded frames 9..80 after the FFFF-only "
"first chunk"
),
},
}
atomic_json(args.output_dir / "manifest.json", manifest)
print("[setup] loading VAE and preparing FFFF reference frames", flush=True)
vae = WanVAEWrapper().to(device=device, dtype=torch.bfloat16).eval()
prepare_reference_frames(
vae=vae,
dataset_root=args.dataset_root,
output_dir=args.output_dir,
prompt_ids=prompt_ids,
device=device,
rebuild=args.rebuild_references,
)
print("[setup] loading frozen generator_ema", flush=True)
pipeline = build_pipeline(config, args.checkpoint_path, vae, device)
teacher = pipeline.generator.model
lpips_model = None
if not args.skip_lpips:
print("[setup] loading AlexNet LPIPS", flush=True)
lpips_model = lpips.LPIPS(net="alex", verbose=False).to(device).eval()
lpips_model.requires_grad_(False)
if args.verify_ffff:
prompt_id = prompt_ids[0]
print(f"[verify] reproducing FFFF prompt={prompt_id}", flush=True)
reproduced, counts = generate_rollout(
pipeline=pipeline,
dataset_root=args.dataset_root,
prompt_id=prompt_id,
generation_seed=args.generation_seed,
device=device,
predictor=None,
source_layer=None,
schedule="FFFF",
)
expected = load_ffff_latent(args.dataset_root, prompt_id).to(
device=device, dtype=torch.bfloat16
)
difference = reproduced.float() - expected.float()
verification = {
"prompt_id": prompt_id,
"max_abs_latent_error": float(difference.abs().max()),
"latent_mse": float(difference.square().mean()),
**counts,
}
atomic_json(args.output_dir / "ffff_reproduction.json", verification)
print(f"[verify] {verification}", flush=True)
if verification["max_abs_latent_error"] > 1e-3:
raise RuntimeError(
"FFFF reproduction differs from offline reference; refusing "
"to evaluate FPPF with unmatched randomness/caches"
)
del reproduced, expected, difference
torch.cuda.empty_cache()
for experiment_index, experiment in enumerate(experiments, start=1):
run_dir = args.output_dir / experiment["name"]
run_dir.mkdir(parents=True, exist_ok=True)
metrics_path = run_dir / "metrics.json"
if metrics_path.exists():
existing = json.loads(metrics_path.read_text(encoding="utf-8"))
if (
existing.get("status") == "complete"
and existing.get("prompt_ids") == prompt_ids
and (args.skip_lpips or existing.get("lpips") is not None)
):
print(
f"[run] {experiment_index}/{len(experiments)} "
f"skip complete {experiment['name']}",
flush=True,
)
continue
print(
f"[run] {experiment_index}/{len(experiments)} "
f"{experiment['name']} source={experiment['source_layer']}",
flush=True,
)
predictor = load_predictor(teacher, experiment, device)
existing_results = {
result["prompt_id"]: result
for result in load_completed_prompt_results(run_dir, prompt_ids)
}
for prompt_index, prompt_id in enumerate(prompt_ids, start=1):
video_path = run_dir / "videos" / f"prompt_{prompt_id:04d}.mp4"
if prompt_id in existing_results and (
not args.save_videos or video_path.exists()
):
print(
f"[prompt] {experiment['name']} {prompt_index}/{len(prompt_ids)} "
f"id={prompt_id} cached",
flush=True,
)
continue
started = time.perf_counter()
latent, counts = 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 = pixels_to_u8(pixels)
if args.save_videos:
save_mp4(prediction_u8, video_path)
reference_u8 = load_reference_frames(args.output_dir, prompt_id)
metrics = frame_metrics(
reference_u8=reference_u8,
prediction_u8=prediction_u8,
lpips_model=lpips_model,
batch_size=args.metric_batch_size,
device=device,
)
prompt_result = {
"prompt_id": prompt_id,
"prompt": load_prompt_metadata(args.dataset_root, prompt_id)[
"prompt"
],
**counts,
**metrics,
"total_time_s": time.perf_counter() - started,
}
atomic_json(
run_dir / "per_prompt" / f"prompt_{prompt_id:04d}.json",
prompt_result,
)
existing_results[prompt_id] = prompt_result
print(
f"[prompt] {experiment['name']} {prompt_index}/{len(prompt_ids)} "
f"id={prompt_id} psnr={metrics['psnr']:.4f} "
f"ssim={metrics['ssim']:.6f} "
f"lpips={metrics['lpips']} "
f"time={prompt_result['total_time_s']:.1f}s",
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 = aggregate_prompt_results(experiment, prompt_results)
atomic_json(metrics_path, aggregate)
write_summary(args.output_dir, experiments)
print(
f"[result] {experiment['name']} psnr={aggregate['psnr']:.4f} "
f"ssim={aggregate['ssim']:.6f} lpips={aggregate['lpips']}",
flush=True,
)
del predictor
torch.cuda.empty_cache()
manifest["status"] = "complete"
atomic_json(args.output_dir / "manifest.json", manifest)
write_summary(args.output_dir, experiments)
print(
f"[complete] {len(experiments)} experiments -> "
f"{args.output_dir / 'summary.csv'}",
flush=True,
)
if __name__ == "__main__":
main()