Self-Forcing / scripts /evaluate_two_block_fppf.py
Cccccz's picture
Upload Python scripts
bc29ee3 verified
Raw
History Blame Contribute Delete
24.4 kB
#!/usr/bin/env python3
"""Evaluate trained multi-block Predictors against matching FFFF rollouts."""
from __future__ import annotations
import argparse
import csv
import json
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
from omegaconf import OmegaConf
from safetensors.torch import load_file
REPO_ROOT = Path(__file__).resolve().parents[1]
if str(REPO_ROOT) not in sys.path:
sys.path.insert(0, str(REPO_ROOT))
from predictor_training.offline_data import TOKENS_PER_CHUNK
from predictor_training.three_block import ThreeBlockPredictor
from predictor_training.two_block import TwoBlockPredictor
from scripts.evaluate_single_block_fppf import (
DEFAULT_PROMPT_IDS,
FRAMES_PER_CHUNK,
LATENT_CHANNELS,
LATENT_HEIGHT,
LATENT_WIDTH,
NUM_CHUNKS,
NUM_DENOISING_STEPS,
FinalHiddenCapture,
aggregate_prompt_results,
atomic_json,
build_pipeline,
frame_metrics,
load_ffff_latent,
load_prompt_metadata,
load_reference_frames,
pixels_to_u8,
prepare_reference_frames,
reset_kv_and_load_cross_cache,
)
from scripts.run_single_block_init_sweep import hidden_to_flow
from utils.misc import set_seed
from utils.wan_wrapper import WanVAEWrapper
from wan.modules.model import sinusoidal_embedding_1d
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--gpu", default=PHYSICAL_GPU)
parser.add_argument(
"--architecture",
choices=("two_block", "three_block"),
default="two_block",
help="Predictor architecture represented by the sweep directory.",
)
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/two_block_pair_sweep")
)
parser.add_argument(
"--output_dir", type=Path, default=Path("outputs/two_block_pair_fppf_eval")
)
parser.add_argument(
"--schedule",
choices=("FPPF", "FPPP"),
default="FPPF",
help=(
"Denoising schedule for chunks 1-6; chunk 0 always uses FFFF. "
"FPPF predicts steps 1-2, while FPPP predicts steps 1-3."
),
)
parser.add_argument(
"--reference_root",
type=Path,
default=Path("outputs/single_block_fppf_eval"),
help="Directory containing reusable ffff_reference_frames/.",
)
parser.add_argument(
"--prompt_ids", type=int, nargs="*", default=DEFAULT_PROMPT_IDS
)
parser.add_argument("--experiments", nargs="*", default=None)
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
)
parser.add_argument(
"--skip_lpips", action=argparse.BooleanOptionalAction, default=False
)
args = parser.parse_args()
if not args.prompt_ids:
parser.error("At least one prompt ID is required")
if args.metric_batch_size < 1:
parser.error("--metric_batch_size must be positive")
return args
def resolve(path: Path) -> Path:
path = path.expanduser()
return path.resolve() if path.is_absolute() else (REPO_ROOT / path).resolve()
def discover_experiments(
sweep_dir: Path,
requested: list[str] | None,
max_experiments: int | None,
architecture: str = "two_block",
) -> list[dict[str, Any]]:
with (sweep_dir / "summary.csv").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 experiments: {unknown}")
if max_experiments is not None:
names = names[:max_experiments]
output = []
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]
kind_key = "pair_kind" if architecture == "two_block" else "triple_kind"
output.append(
{
"name": name,
"source_layers": [int(value) for value in config["source_layers"]],
"experiment_kind": config[kind_key],
"initialization_method": "teacher_full",
"weights": weights,
"offline_final_val_flow_mse": float(row["final_val_flow_mse"]),
"offline_final_val_hidden_mse": float(
row["final_val_hidden_mse"]
),
}
)
return output
def load_predictor(
teacher: torch.nn.Module,
experiment: dict[str, Any],
device: torch.device,
) -> TwoBlockPredictor | ThreeBlockPredictor:
source_layers = experiment["source_layers"]
predictor_class = (
TwoBlockPredictor if len(source_layers) == 2 else ThreeBlockPredictor
)
predictor = predictor_class(
[teacher.blocks[layer] for layer in source_layers],
dim=teacher.dim,
gradient_checkpointing=False,
)
predictor.load_state_dict(
load_file(str(experiment["weights"]), device="cpu"), strict=True
)
predictor.to(device=device).eval().requires_grad_(False)
return predictor
@torch.inference_mode()
def predictor_step(
*,
predictor: TwoBlockPredictor | ThreeBlockPredictor,
teacher: torch.nn.Module,
noisy_input: torch.Tensor,
timestep: torch.Tensor,
anchor_hidden: torch.Tensor,
previous_hidden: torch.Tensor,
history_caches: list[dict[str, torch.Tensor]],
cross_caches: list[dict[str, torch.Tensor]],
current_start: int,
) -> tuple[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)
grid_sizes = torch.tensor(
[[FRAMES_PER_CHUNK, 30, 52]], dtype=torch.long, device="cpu"
)
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_ks=[
cache["k"][:, :current_start] for cache in history_caches
],
history_vs=[
cache["v"][:, :current_start] for cache in history_caches
],
cross_ks=[cache["k"] for cache in cross_caches],
cross_vs=[cache["v"] for cache in cross_caches],
current_start=current_start,
)
pred_flow = hidden_to_flow(
pred_hidden, head_embedding, grid_sizes, teacher
)
return pred_hidden, pred_flow
@torch.inference_mode()
def generate_rollout(
*,
pipeline,
dataset_root: Path,
prompt_id: int,
generation_seed: int,
device: torch.device,
predictor: TwoBlockPredictor | ThreeBlockPredictor | None,
source_layers: list[int] | None,
schedule: str,
) -> tuple[torch.Tensor, dict[str, float | int]]:
if schedule not in {"FFFF", "FPPF", "FPPP"}:
raise ValueError(schedule)
if schedule != "FFFF" and (predictor is None or source_layers is None):
raise ValueError(f"{schedule} requires a Predictor and source layers")
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 = None
timestep = None
for step, current_timestep in enumerate(timesteps):
timestep = torch.ones(
[1, FRAMES_PER_CHUNK], dtype=torch.int64, device=device
) * current_timestep
predictor_steps = {1, 2} if schedule == "FPPF" else {1, 2, 3}
use_predictor = (
schedule != "FFFF" and chunk > 0 and step in predictor_steps
)
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
pred_hidden, flow = predictor_step(
predictor=predictor,
teacher=teacher,
noisy_input=noisy_input,
timestep=timestep,
anchor_hidden=anchor_hidden,
previous_hidden=previous_hidden,
history_caches=[
pipeline.kv_cache1[layer] for layer in source_layers
],
cross_caches=[
pipeline.crossattn_cache[layer]
for layer in source_layers
],
current_start=chunk * TOKENS_PER_CHUNK,
)
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)
pipeline.generator(
noisy_image_or_video=denoised_pred,
conditional_dict=conditional_dict,
timestep=torch.ones_like(timestep) * pipeline.args.context_noise,
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 write_summary(output_dir: Path, experiments: list[dict[str, Any]]) -> None:
rows = []
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
source_layers = metrics["source_layers"]
kind_field = "pair_kind" if len(source_layers) == 2 else "triple_kind"
row = {
"name": metrics["name"],
**{
f"source_layer_{index + 1}": layer
for index, layer in enumerate(source_layers)
},
kind_field: metrics.get(
kind_field, metrics.get("experiment_kind")
),
"schedule": metrics["schedule"],
"num_prompts": metrics["num_prompts"],
"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.append(row)
rows.sort(key=lambda row: float(row["lpips"]))
if not rows:
return
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=list(rows[0]))
writer.writeheader()
writer.writerows(rows)
os.replace(temporary, destination)
atomic_json(output_dir / "summary.json", rows)
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.reference_root = resolve(args.reference_root)
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, args.architecture
)
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),
)
schedule_description = f"chunk0=FFFF; chunks1-6={args.schedule}"
manifest = {
"status": "running",
"architecture": f"{args.architecture}_predictor",
"prompt_ids": prompt_ids,
"experiments": [item["name"] for item in experiments],
"rollout_schedule": args.schedule,
"rollout_definition": schedule_description,
"reference_root": str(args.reference_root),
"generation_seed_reset_per_prompt": args.generation_seed,
}
atomic_json(args.output_dir / "manifest.json", manifest)
print("[setup] loading VAE and checking 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.reference_root,
prompt_ids=prompt_ids,
device=device,
rebuild=False,
)
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:
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]
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_layers=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 does not match offline reference")
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 existing.get("schedule") == schedule_description
and (args.skip_lpips or existing.get("lpips") is not None)
):
print(f"[run] skip complete {experiment['name']}", flush=True)
continue
print(
f"[run] {experiment_index}/{len(experiments)} {experiment['name']}",
flush=True,
)
predictor = load_predictor(teacher, experiment, device)
existing_results = {}
for prompt_id in prompt_ids:
path = run_dir / "per_prompt" / f"prompt_{prompt_id:04d}.json"
if path.exists():
cached = json.loads(path.read_text(encoding="utf-8"))
if cached.get("rollout_schedule", "FPPF") == args.schedule:
existing_results[prompt_id] = cached
for prompt_index, prompt_id in enumerate(prompt_ids, start=1):
if prompt_id in existing_results:
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_layers=experiment["source_layers"],
schedule=args.schedule,
)
with torch.autocast(device_type="cuda", dtype=torch.bfloat16):
pixels = vae.decode_to_pixel(latent, use_cache=False)
prediction_u8 = pixels_to_u8(pixels)
reference_u8 = load_reference_frames(args.reference_root, 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"
],
"rollout_schedule": args.schedule,
**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} 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()
base_experiment = {
**experiment,
"source_layer": experiment["source_layers"],
}
aggregate = aggregate_prompt_results(
base_experiment,
[existing_results[prompt_id] for prompt_id in prompt_ids],
)
aggregate["source_layers"] = experiment["source_layers"]
kind_field = (
"pair_kind"
if len(experiment["source_layers"]) == 2
else "triple_kind"
)
aggregate[kind_field] = experiment["experiment_kind"]
aggregate["schedule"] = schedule_description
aggregate.pop("source_layer", None)
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 -> {args.output_dir / 'summary.csv'}",
flush=True,
)
if __name__ == "__main__":
main()