Self-Forcing / scripts /evaluate_layer17_dynamic_gate.py
Cccccz's picture
Upload Python scripts
bc29ee3 verified
Raw
History Blame Contribute Delete
36.4 kB
#!/usr/bin/env python3
"""Evaluate chunk-aware dynamic gating for the frozen Layer-17 Predictor."""
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="4")
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.confidence import PredictorConfidenceHead
from scripts import evaluate_layer17_chunk_impact as impact_eval
from scripts import evaluate_single_block_fppf as base
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
BETAS = (0.0, 1.0, 1.5, 2.0)
TARGET_ACCEPTS = (4, 6, 8, 10)
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--gpu", default=PHYSICAL_GPU)
parser.add_argument("--mode", choices=("smoke", "validation", "test"), required=True)
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(
"--predictor_weights",
type=Path,
default=None,
help=(
"Optional direct Layer-17 Predictor weights. When supplied, this "
"takes precedence over teacher_layer_17 in --sweep_dir."
),
)
parser.add_argument(
"--confidence_weights", type=Path,
default=Path(
"outputs/layer17_confidence_teacher_forced_20260830/"
"confidence_best.safetensors"
),
)
parser.add_argument(
"--validation_predictions", type=Path,
default=Path(
"outputs/layer17_confidence_teacher_forced_20260830/"
"validation_predictions.csv"
),
)
parser.add_argument(
"--reference_root", type=Path, default=Path("outputs/single_block_fppf_eval")
)
parser.add_argument(
"--output_root", type=Path,
default=Path("outputs/layer17_dynamic_gate_20260830"),
)
parser.add_argument("--generation_seed", type=int, default=0)
parser.add_argument("--metric_batch_size", type=int, default=4)
parser.add_argument(
"--candidate_steps", type=int, nargs="+", choices=(1, 2, 3),
default=[1, 2],
)
parser.add_argument("--target_accepts", type=int, nargs="*", default=None)
parser.add_argument("--selected_path", type=Path, default=None)
parser.add_argument(
"--config_names", nargs="*", default=None,
help="Optional exact configuration names to run in validation/test mode.",
)
parser.add_argument(
"--prompt_ids", type=int, nargs="*", default=None,
help="Optional prompt shard; shard-level CSV/manifest files get a GPU suffix.",
)
parser.add_argument("--save_videos", action="store_true")
parser.add_argument("--overwrite", action="store_true")
parser.add_argument(
"--skip_lpips", action=argparse.BooleanOptionalAction, default=False
)
args = parser.parse_args()
for name in (
"config_path", "checkpoint_path", "dataset_root", "sweep_dir",
"predictor_weights",
"confidence_weights", "validation_predictions", "reference_root",
"output_root", "selected_path",
):
value = getattr(args, name)
if value is None:
continue
path = value.expanduser()
setattr(args, name, path.resolve() if path.is_absolute() else (REPO_ROOT / path).resolve())
args.candidate_steps = sorted(set(args.candidate_steps))
max_accepts = 6 * len(args.candidate_steps)
if args.target_accepts is None:
args.target_accepts = (
[4, 6, 8, 10] if len(args.candidate_steps) == 2
else [6, 9, 12, 15]
)
if any(value < 1 or value >= max_accepts for value in args.target_accepts):
parser.error(f"target accepts must be in [1, {max_accepts - 1}]")
return args
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) + "\n", encoding="utf-8")
os.replace(temporary, path)
def write_csv(path: Path, rows: list[dict[str, Any]], fields: list[str]) -> None:
temporary = path.with_suffix(path.suffix + ".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, path)
def quantile(values: list[float], fraction: float) -> float:
ordered = sorted(values)
position = fraction * (len(ordered) - 1)
lower = int(math.floor(position))
upper = int(math.ceil(position))
if lower == upper:
return ordered[lower]
weight = position - lower
return ordered[lower] * (1.0 - weight) + ordered[upper] * weight
def threshold_grid(
predictions_path: Path,
candidate_steps: list[int],
targets: list[int],
) -> dict[tuple[float, int], float]:
rows = list(csv.DictReader(predictions_path.open(encoding="utf-8")))
expected = 10 * 6 * len(candidate_steps)
if len(rows) != expected:
raise ValueError(f"Expected {expected} validation predictions, got {len(rows)}")
thresholds = {}
for beta in BETAS:
risks = []
for row in rows:
chunk = int(row["chunk"])
alpha = (base.NUM_CHUNKS - 1 - chunk) / (base.NUM_CHUNKS - 2)
local_error = float(row["predicted_hidden_nrmse"])
risks.append(local_error * (1.0 + beta * alpha))
max_accepts = 6 * len(candidate_steps)
for target in targets:
thresholds[(beta, target)] = quantile(risks, target / max_accepts)
return thresholds
def dynamic_configs(
predictions_path: Path,
candidate_steps: list[int],
targets: list[int],
) -> list[dict[str, Any]]:
thresholds = threshold_grid(predictions_path, candidate_steps, targets)
return [
{
"name": f"dynamic_b{str(beta).replace('.', 'p')}_k{target:02d}",
"policy": "dynamic",
"beta": beta,
"target_accepts": target,
"threshold": thresholds[(beta, target)],
}
for beta in BETAS
for target in targets
]
def static_configs(targets: list[int]) -> list[dict[str, Any]]:
return [
{
"name": f"static_late_k{target:02d}",
"policy": "static_late",
"beta": None,
"target_accepts": target,
"threshold": None,
}
for target in targets
]
def load_selected(path: Path) -> list[dict[str, Any]]:
value = json.loads(path.read_text(encoding="utf-8"))
return [
{
"name": row["config_name"],
"policy": "dynamic",
"beta": float(row["beta"]),
"target_accepts": int(row["target_accepts"]),
"threshold": float(row["threshold"]),
}
for row in value["selected_dynamic"]
]
@torch.no_grad()
def predictor_with_features(
*,
predictor: Any,
teacher: Any,
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(
[[base.FRAMES_PER_CHUNK, 30, 52]], dtype=torch.long, device="cpu"
)
history_length = int(history_cache["local_end_index"].item())
output = 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,
return_features=True,
condition_tokens=condition_tokens,
anchor_distance=anchor_distance,
)
if not isinstance(output, tuple):
raise RuntimeError("Predictor did not return internal features")
pred_hidden, transformed = output
pred_flow = hidden_to_flow(
pred_hidden, head_embedding, grid_sizes, teacher
)
return pred_hidden, pred_flow, transformed
@torch.inference_mode()
def generate(
*,
pipeline: Any,
dataset_root: Path,
prompt_id: int,
generation_seed: int,
device: torch.device,
predictor: Any,
head: PredictorConfidenceHead,
config: dict[str, Any],
candidate_steps: list[int],
) -> tuple[torch.Tensor, dict[str, Any]]:
base.reset_kv_and_load_cross_cache(pipeline, dataset_root, prompt_id, device)
set_seed(generation_seed)
noise = torch.randn(
1,
base.NUM_CHUNKS * base.FRAMES_PER_CHUNK,
base.LATENT_CHANNELS,
base.LATENT_HEIGHT,
base.LATENT_WIDTH,
dtype=torch.bfloat16,
device=device,
)
timesteps = pipeline.denoising_step_list.to(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
)
}
capture = base.FinalHiddenCapture(teacher)
output_chunks: list[torch.Tensor] = []
previous_chunk_hidden: list[torch.Tensor | None] | None = None
decisions: list[dict[str, Any]] = []
full_calls = 0
predictor_calls = 0
accepted_predictor_calls = 0
timing_events: dict[str, list[tuple[torch.cuda.Event, torch.cuda.Event]]] = {
"full_dit": [],
"predictor": [],
"confidence": [],
"context_dit": [],
}
def start_timing() -> tuple[torch.cuda.Event, torch.cuda.Event]:
start_event = torch.cuda.Event(enable_timing=True)
end_event = torch.cuda.Event(enable_timing=True)
start_event.record()
return start_event, end_event
def finish_timing(
category: str,
events: tuple[torch.cuda.Event, torch.cuda.Event],
) -> None:
events[1].record()
timing_events[category].append(events)
started = time.perf_counter()
try:
for chunk in range(base.NUM_CHUNKS):
noisy_input = noise[
:, chunk * base.FRAMES_PER_CHUNK : (chunk + 1) * base.FRAMES_PER_CHUNK
]
current_hidden: list[torch.Tensor | None] = [None] * base.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, base.FRAMES_PER_CHUNK], dtype=torch.int64, device=device
) * current_timestep
candidate = chunk > 0 and step in candidate_steps
policy = str(config["policy"])
run_predictor = False
static_accept = False
if candidate and policy == "dynamic":
run_predictor = True
elif candidate and policy == "fppf":
run_predictor = True
static_accept = True
elif candidate and policy == "static_late":
first_chunk = (
base.NUM_CHUNKS
- int(config["target_accepts"]) // len(candidate_steps)
)
static_accept = chunk >= first_chunk
run_predictor = static_accept
accepted = False
pred_hidden = None
pred_x0 = None
predicted_local_error = None
risk = None
alpha = None
if run_predictor:
if previous_chunk_hidden is None:
raise RuntimeError("Previous chunk hidden is unavailable")
anchor_hidden = current_hidden[step - 1]
previous_hidden = previous_chunk_hidden[step]
if anchor_hidden is None or previous_hidden is None:
raise RuntimeError("Predictor inputs are unavailable")
predictor_events = start_timing()
pred_hidden, pred_flow, transformed = predictor_with_features(
predictor=predictor,
teacher=teacher,
noisy_input=noisy_input,
timestep=timestep,
anchor_hidden=anchor_hidden,
previous_hidden=previous_hidden,
history_cache=pipeline.kv_cache1[17],
cross_cache=pipeline.crossattn_cache[17],
current_start=chunk * base.TOKENS_PER_CHUNK,
anchor_timestep=(
torch.ones_like(timestep) * timesteps[step - 1]
),
)
finish_timing("predictor", predictor_events)
pred_x0 = pipeline.generator._convert_flow_pred_to_x0(
flow_pred=pred_flow.flatten(0, 1),
xt=noisy_input.flatten(0, 1),
timestep=timestep.flatten(0, 1),
).unflatten(0, pred_flow.shape[:2])
predictor_calls += 1
if policy == "dynamic":
chunk_position = torch.tensor(
[(chunk - 1) / 5.0], device=device
)
step_tensor = torch.tensor([step], dtype=torch.long, device=device)
confidence_events = start_timing()
with torch.autocast(device_type="cuda", dtype=torch.bfloat16):
predicted_log = head(
transformed_hidden=transformed,
pred_hidden=pred_hidden,
anchor_hidden=anchor_hidden,
chunk_position=chunk_position,
step_id=step_tensor,
)
finish_timing("confidence", confidence_events)
predicted_local_error = float(predicted_log.exp()[0])
alpha = (base.NUM_CHUNKS - 1 - chunk) / (base.NUM_CHUNKS - 2)
risk = predicted_local_error * (
1.0 + float(config["beta"]) * alpha
)
accepted = risk <= float(config["threshold"])
else:
accepted = static_accept
if accepted:
assert pred_hidden is not None and pred_x0 is not None
current_hidden[step] = pred_hidden
denoised_pred = pred_x0
accepted_predictor_calls += 1
else:
full_events = start_timing()
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 * base.TOKENS_PER_CHUNK,
)
current_hidden[step] = capture.finish()
finish_timing("full_dit", full_events)
full_calls += 1
if candidate:
decisions.append(
{
"chunk": chunk,
"step": step,
"ran_predictor": run_predictor,
"accepted": accepted,
"predicted_local_error": predicted_local_error,
"chunk_alpha": alpha,
"impact_risk": risk,
}
)
if step < base.NUM_DENOISING_STEPS - 1:
if denoised_pred is None:
raise RuntimeError("Denoising step produced no x0")
next_timestep = timesteps[step + 1]
flat = denoised_pred.flatten(0, 1)
noisy_input = pipeline.scheduler.add_noise(
flat,
torch.randn_like(flat),
next_timestep
* torch.ones(
[base.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("Chunk produced no clean latent")
output_chunks.append(denoised_pred)
context_timestep = torch.ones_like(timestep) * pipeline.args.context_noise
context_events = start_timing()
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 * base.TOKENS_PER_CHUNK,
)
finish_timing("context_dit", context_events)
previous_chunk_hidden = current_hidden
finally:
capture.close()
torch.cuda.synchronize()
elapsed = {
category: sum(start.elapsed_time(end) for start, end in events)
for category, events in timing_events.items()
}
actual_dit_time_ms = (
elapsed["full_dit"] + elapsed["predictor"] + elapsed["context_dit"]
)
return torch.cat(output_chunks, dim=1), {
"generation_time_s": time.perf_counter() - started,
"full_calls": full_calls,
"predictor_calls": predictor_calls,
"accepted_predictor_calls": accepted_predictor_calls,
"rejected_predictor_calls": predictor_calls - accepted_predictor_calls,
"full_dit_time_ms": elapsed["full_dit"],
"predictor_time_ms": elapsed["predictor"],
"confidence_head_time_ms": elapsed["confidence"],
"context_dit_time_ms": elapsed["context_dit"],
"actual_dit_time_ms": actual_dit_time_ms,
"model_path_time_ms": actual_dit_time_ms + elapsed["confidence"],
"decisions": decisions,
}
def load_models(
args: argparse.Namespace, device: torch.device
) -> tuple[Any, Any, Any, Any, Any]:
print("[setup] loading VAE", flush=True)
vae = WanVAEWrapper().to(device=device, dtype=torch.bfloat16).eval()
config = OmegaConf.merge(
OmegaConf.load(REPO_ROOT / "configs/default_config.yaml"),
OmegaConf.load(args.config_path),
)
print("[setup] loading frozen generator and Predictor", flush=True)
pipeline = base.build_pipeline(config, args.checkpoint_path, vae, device)
if args.predictor_weights is not None:
experiment = {
"name": "direct_layer17_predictor",
"source_layer": 17,
"weights": args.predictor_weights,
"gate_mode": "baseline",
}
else:
experiment = base.discover_experiments(
args.sweep_dir, ["teacher_layer_17"], None
)[0]
predictor = base.load_predictor(pipeline.generator.model, experiment, device)
head = PredictorConfidenceHead(
num_steps=max(args.candidate_steps)
).to(device=device).eval()
head.load_state_dict(load_file(str(args.confidence_weights), device="cpu"), strict=True)
head.requires_grad_(False)
lpips_model = None
if not args.skip_lpips:
lpips_model = lpips.LPIPS(net="alex", verbose=False).to(device).eval()
lpips_model.requires_grad_(False)
return vae, pipeline, predictor, head, lpips_model
def smoke(args: argparse.Namespace, pipeline: Any, predictor: Any, head: Any, device: torch.device) -> None:
max_accepts = 6 * len(args.candidate_steps)
all_name = "fppf" if args.candidate_steps == [1, 2] else "fppp"
configurations = [
{"name": "ffff", "policy": "ffff", "beta": None, "threshold": None, "target_accepts": 0},
{"name": "dynamic_all_fallback", "policy": "dynamic", "beta": 1.5, "threshold": -math.inf, "target_accepts": 0},
{"name": all_name, "policy": "fppf", "beta": None, "threshold": None, "target_accepts": max_accepts},
{"name": "dynamic_all_accept", "policy": "dynamic", "beta": 1.5, "threshold": math.inf, "target_accepts": max_accepts},
]
latents = {}
diagnostics = {}
for config in configurations:
latent, diagnostic = generate(
pipeline=pipeline, dataset_root=args.dataset_root, prompt_id=80,
generation_seed=args.generation_seed, device=device, predictor=predictor,
head=head, config=config, candidate_steps=args.candidate_steps,
)
latents[config["name"]] = latent.cpu()
diagnostics[config["name"]] = diagnostic
print(
f"[smoke] {config['name']} full={diagnostic['full_calls']} "
f"pred={diagnostic['predictor_calls']} accept={diagnostic['accepted_predictor_calls']}",
flush=True,
)
fallback_diff = float((latents["ffff"].float() - latents["dynamic_all_fallback"].float()).abs().max())
accept_diff = float((latents[all_name].float() - latents["dynamic_all_accept"].float()).abs().max())
result = {
"status": "complete",
"prompt_id": 80,
"ffff_vs_all_fallback_max_abs": fallback_diff,
"fppf_vs_all_accept_max_abs": accept_diff,
"diagnostics": diagnostics,
}
atomic_json(args.output_root / "smoke.json", result)
if fallback_diff != 0.0 or accept_diff != 0.0:
raise RuntimeError(f"Smoke consistency failed: {result}")
print("[smoke] exact consistency passed", flush=True)
def aggregate(records: list[dict[str, Any]], output_dir: Path) -> list[dict[str, Any]]:
numeric = [
"accepted_predictor_calls", "full_calls", "predictor_calls",
"full_dit_time_ms", "predictor_time_ms", "confidence_head_time_ms",
"context_dit_time_ms", "actual_dit_time_ms", "model_path_time_ms",
"generation_time_s", "total_time_s", "latent_nrmse", "latent_tail_nrmse",
"psnr", "ssim", "lpips", "tail_psnr", "tail_ssim", "tail_lpips",
]
summary = []
for name in sorted({str(row["config_name"]) for row in records}):
selected = [row for row in records if row["config_name"] == name]
first = selected[0]
item = {
"config_name": name,
"policy": first["policy"],
"beta": first["beta"],
"target_accepts": first["target_accepts"],
"threshold": first["threshold"],
"num_prompts": len(selected),
}
for field in numeric:
item[field] = sum(float(row[field]) for row in selected) / len(selected)
summary.append(item)
fields = [
"config_name", "policy", "beta", "target_accepts", "threshold",
"num_prompts", *numeric,
]
write_csv(output_dir / "summary.csv", summary, fields)
return summary
def select_validation(
summary: list[dict[str, Any]], output_dir: Path, targets: list[int]
) -> None:
selected_dynamic = []
for target in targets:
candidates = [
row for row in summary
if row["policy"] == "dynamic" and int(row["target_accepts"]) == target
]
same_budget = [
row for row in candidates
if abs(float(row["accepted_predictor_calls"]) - target) <= 0.5 + 1e-8
]
if not same_budget:
closest = min(
abs(float(row["accepted_predictor_calls"]) - target)
for row in candidates
)
same_budget = [
row for row in candidates
if abs(abs(float(row["accepted_predictor_calls"]) - target) - closest)
<= 1e-8
]
same_budget.sort(
key=lambda row: (
float(row["tail_lpips"]),
abs(float(row["accepted_predictor_calls"]) - target),
float(row["beta"]),
)
)
selected_dynamic.append(same_budget[0])
atomic_json(
output_dir / "selected.json",
{
"selection_rule": (
"within target accepted calls +/-0.5, lowest validation tail LPIPS; "
"then budget distance and lower beta"
),
"selected_dynamic": selected_dynamic,
},
)
def formal(
args: argparse.Namespace,
vae: Any,
pipeline: Any,
predictor: Any,
head: Any,
lpips_model: Any,
device: torch.device,
) -> None:
split = args.mode
split_prompt_ids = (
list(range(80, 90)) if split == "validation" else list(range(90, 100))
)
prompt_ids = args.prompt_ids or split_prompt_ids
invalid_prompt_ids = sorted(set(prompt_ids) - set(split_prompt_ids))
if invalid_prompt_ids:
raise ValueError(
f"Prompt IDs {invalid_prompt_ids} are outside the {split} split"
)
if split == "validation":
configurations = dynamic_configs(
args.validation_predictions, args.candidate_steps, args.target_accepts
) + static_configs(args.target_accepts)
max_accepts = 6 * len(args.candidate_steps)
all_name = "fppf" if args.candidate_steps == [1, 2] else "fppp"
configurations += [
{"name": "ffff", "policy": "ffff", "beta": None, "threshold": None, "target_accepts": 0},
{"name": all_name, "policy": "fppf", "beta": None, "threshold": None, "target_accepts": max_accepts},
]
else:
selected_path = args.selected_path or (
args.output_root / "validation" / "selected.json"
)
configurations = load_selected(selected_path) + static_configs(args.target_accepts)
max_accepts = 6 * len(args.candidate_steps)
all_name = "fppf" if args.candidate_steps == [1, 2] else "fppp"
configurations += [
{"name": "ffff", "policy": "ffff", "beta": None, "threshold": None, "target_accepts": 0},
{"name": all_name, "policy": "fppf", "beta": None, "threshold": None, "target_accepts": max_accepts},
]
if args.config_names:
requested = set(args.config_names)
available = {str(config["name"]) for config in configurations}
missing = requested - available
if missing:
raise ValueError(
f"Unknown config_names {sorted(missing)}; available={sorted(available)}"
)
configurations = [
config for config in configurations if config["name"] in requested
]
output_dir = args.output_root / split
output_dir.mkdir(parents=True, exist_ok=True)
missing_references = [
prompt_id for prompt_id in prompt_ids
if not (
args.reference_root / "ffff_reference_frames"
/ f"prompt_{prompt_id:04d}.safetensors"
).exists()
]
if missing_references:
base.prepare_reference_frames(
vae=vae, dataset_root=args.dataset_root, output_dir=args.reference_root,
prompt_ids=missing_references, device=device, rebuild=False,
)
total = len(configurations) * len(prompt_ids)
records = []
completed = 0
print("[warmup] one unmeasured Full+Predictor+Head rollout", flush=True)
warmup_config = {
"name": "warmup",
"policy": "dynamic",
"beta": 1.0,
"threshold": -math.inf,
"target_accepts": 0,
}
warmup_latent, _ = generate(
pipeline=pipeline,
dataset_root=args.dataset_root,
prompt_id=prompt_ids[0],
generation_seed=args.generation_seed,
device=device,
predictor=predictor,
head=head,
config=warmup_config,
candidate_steps=args.candidate_steps,
)
del warmup_latent
torch.cuda.empty_cache()
for config in configurations:
for prompt_id in prompt_ids:
destination = output_dir / "per_run" / config["name"] / f"prompt_{prompt_id:04d}.json"
if destination.exists() and not args.overwrite:
records.append(json.loads(destination.read_text(encoding="utf-8")))
completed += 1
print(f"[cached] {completed}/{total} {config['name']} p={prompt_id}", flush=True)
continue
started = time.perf_counter()
reference_latent = base.load_ffff_latent(args.dataset_root, prompt_id).to(
device=device, dtype=torch.bfloat16
)
reference_u8 = base.load_reference_frames(args.reference_root, prompt_id)
latent, diagnostic = generate(
pipeline=pipeline, dataset_root=args.dataset_root, prompt_id=prompt_id,
generation_seed=args.generation_seed, device=device, predictor=predictor,
head=head, config=config, candidate_steps=args.candidate_steps,
)
with torch.autocast(device_type="cuda", dtype=torch.bfloat16):
pixels = vae.decode_to_pixel(latent, use_cache=False)
prediction_u8 = base.pixels_to_u8(pixels)
if args.save_videos:
base.save_mp4(
prediction_u8,
output_dir / "videos" / config["name"]
/ f"prompt_{prompt_id:04d}.mp4",
)
frame = base.frame_metrics(
reference_u8=reference_u8, prediction_u8=prediction_u8,
lpips_model=lpips_model, batch_size=args.metric_batch_size, device=device,
)
tail_start = impact_eval.chunk_frame_slice(1).start
tail = impact_eval.summarize_frame_range(frame, slice(tail_start, None))
record = {
"config_name": config["name"],
"policy": config["policy"],
"beta": config["beta"],
"target_accepts": config["target_accepts"],
"threshold": config["threshold"],
"prompt_id": prompt_id,
"latent_nrmse": impact_eval.nrmse(latent, reference_latent),
"latent_tail_nrmse": impact_eval.nrmse(
latent[:, base.FRAMES_PER_CHUNK:],
reference_latent[:, base.FRAMES_PER_CHUNK:],
),
"psnr": frame["psnr"],
"ssim": frame["ssim"],
"lpips": frame["lpips"],
"tail_psnr": tail["psnr"],
"tail_ssim": tail["ssim"],
"tail_lpips": tail["lpips"],
**diagnostic,
"total_time_s": time.perf_counter() - started,
}
atomic_json(destination, record)
records.append(record)
completed += 1
print(
f"[run] {completed}/{total} {config['name']} p={prompt_id} "
f"accept={record['accepted_predictor_calls']} "
f"tail_lpips={record['tail_lpips']:.5f} "
f"time={record['total_time_s']:.1f}s",
flush=True,
)
if hasattr(vae.model, "clear_cache"):
vae.model.clear_cache()
del reference_latent, reference_u8, latent, pixels, prediction_u8, frame
torch.cuda.empty_cache()
flat_fields = sorted({key for row in records for key in row if key != "decisions"})
shard_suffix = f"_gpu{args.gpu}" if args.prompt_ids else ""
write_csv(
output_dir / f"runs{shard_suffix}.csv",
[{key: row.get(key) for key in flat_fields} for row in records],
flat_fields,
)
summary_output_dir = output_dir
if shard_suffix:
summary_output_dir = output_dir / f".summary_shard_gpu{args.gpu}"
summary_output_dir.mkdir(parents=True, exist_ok=True)
summary = aggregate(records, summary_output_dir)
if shard_suffix:
os.replace(
summary_output_dir / "summary.csv",
output_dir / f"summary{shard_suffix}.csv",
)
summary_output_dir.rmdir()
if split == "validation" and not shard_suffix:
select_validation(summary, output_dir, args.target_accepts)
atomic_json(
output_dir / f"manifest{shard_suffix}.json",
{
"status": "complete", "split": split, "prompt_ids": prompt_ids,
"num_configs": len(configurations), "num_runs": len(records),
"configs": configurations, "candidate_steps": args.candidate_steps,
"target_accepts": args.target_accepts,
"predictor_weights": (
str(args.predictor_weights) if args.predictor_weights else None
),
},
)
print(f"[complete] {split} -> {output_dir}", flush=True)
def main() -> None:
args = parse_args()
args.output_root.mkdir(parents=True, exist_ok=True)
device = torch.device("cuda")
torch.set_grad_enabled(False)
set_seed(args.generation_seed)
vae, pipeline, predictor, head, lpips_model = load_models(args, device)
if args.mode == "smoke":
smoke(args, pipeline, predictor, head, device)
else:
formal(args, vae, pipeline, predictor, head, lpips_model, device)
if __name__ == "__main__":
main()