Self-Forcing / scripts /evaluate_single_chunk_frrr.py
Cccccz's picture
Upload Python scripts
bc29ee3 verified
Raw
History Blame Contribute Delete
32.1 kB
#!/usr/bin/env python3
"""Measure the effect of introducing a reuse schedule in one AR chunk.
The four letters describe the four denoising steps of a chunk. ``F`` runs
the full generator. ``R`` reuses the flow prediction from the most recent
full step and only applies the timestep-dependent x0 conversion. For every
prompt this script creates an all-FFFF reference and seven interventions; in
intervention k only chunk k uses the requested schedule and all other chunks
use FFFF.
All rollouts for a prompt reset the RNG to the same seed, so the initial noise
and the three re-noising tensors per chunk are identical. Metrics are
computed on VAE-decoded, rounded uint8 RGB frames before MP4 compression.
"""
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="0")
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.torch import load_file, save_file
from torchvision.io import write_video
REPO_ROOT = Path(
os.environ.get("EVAL_REPO_ROOT", str(Path(__file__).resolve().parents[1]))
).resolve()
if str(REPO_ROOT) not in sys.path:
sys.path.insert(0, str(REPO_ROOT))
from pipeline import CausalInferencePipeline
from utils.misc import set_seed
LATENT_CHANNELS = 16
LATENT_HEIGHT = 60
LATENT_WIDTH = 104
FRAMES_PER_CHUNK = 3
NUM_CHUNKS = 7
NUM_DENOISING_STEPS = 4
TOKENS_PER_FRAME = 30 * 52
TOKENS_PER_CHUNK = FRAMES_PER_CHUNK * TOKENS_PER_FRAME
DECODED_FRAMES = 1 + 4 * (NUM_CHUNKS * FRAMES_PER_CHUNK - 1)
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_dmd.yaml")
)
parser.add_argument(
"--checkpoint_path",
type=Path,
default=Path("checkpoints/self_forcing_dmd.pt"),
)
parser.add_argument(
"--prompt_path",
type=Path,
default=Path("prompts/MovieGenVideoBench_extended.txt"),
)
parser.add_argument(
"--output_dir",
type=Path,
default=Path("outputs/single_chunk_frrr_first10"),
)
parser.add_argument("--prompt_ids", type=int, nargs="*", default=list(range(10)))
parser.add_argument("--seed", type=int, default=0)
parser.add_argument(
"--num_chunks",
type=int,
default=7,
help="Number of 3-latent-frame autoregressive chunks.",
)
parser.add_argument(
"--intervention_schedule",
choices=["FRRR", "FRRF"],
default="FRRR",
help="Four-step schedule used in the selected chunk.",
)
parser.add_argument("--metric_batch_size", type=int, default=4)
parser.add_argument("--use_ema", action=argparse.BooleanOptionalAction, default=True)
parser.add_argument("--save_videos", action=argparse.BooleanOptionalAction, default=True)
parser.add_argument(
"--low_memory",
action=argparse.BooleanOptionalAction,
default=True,
help="Keep only the currently used text encoder/generator/VAE on CUDA.",
)
parser.add_argument("--overwrite", action="store_true")
parser.add_argument("--aggregate_only", action="store_true")
parser.add_argument(
"--worker",
action="store_true",
help="Shard worker: do not rewrite shared config or aggregate files.",
)
args = parser.parse_args()
if not args.prompt_ids:
parser.error("--prompt_ids cannot be empty")
if any(index < 0 for index in args.prompt_ids):
parser.error("prompt IDs must be non-negative")
if args.metric_batch_size < 1:
parser.error("--metric_batch_size must be positive")
if args.num_chunks < 1:
parser.error("--num_chunks must be positive")
global NUM_CHUNKS, DECODED_FRAMES
NUM_CHUNKS = args.num_chunks
DECODED_FRAMES = 1 + 4 * (NUM_CHUNKS * FRAMES_PER_CHUNK - 1)
return args
def resolve(path: Path) -> Path:
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=False) + "\n",
encoding="utf-8",
)
os.replace(temporary, path)
def read_prompts(path: Path) -> list[str]:
with path.open("r", encoding="utf-8") as handle:
return [line.strip() for line in handle if line.strip()]
def build_pipeline(args: argparse.Namespace) -> CausalInferencePipeline:
config = OmegaConf.merge(
OmegaConf.load(REPO_ROOT / "configs/default_config.yaml"),
OmegaConf.load(resolve(args.config_path)),
)
pipeline = CausalInferencePipeline(config, device=torch.device("cuda"))
checkpoint = torch.load(
resolve(args.checkpoint_path), map_location="cpu", weights_only=False
)
state_key = "generator_ema" if args.use_ema else "generator"
pipeline.generator.load_state_dict(checkpoint[state_key])
del checkpoint
pipeline = pipeline.to(dtype=torch.bfloat16)
if not args.low_memory:
pipeline.text_encoder.to(device="cuda")
pipeline.generator.to(device="cuda")
pipeline.vae.to(device="cuda")
pipeline.eval()
if pipeline.num_frame_per_block != FRAMES_PER_CHUNK:
raise ValueError(
f"Expected {FRAMES_PER_CHUNK} latent frames/chunk, got "
f"{pipeline.num_frame_per_block}"
)
if len(pipeline.denoising_step_list) != NUM_DENOISING_STEPS:
raise ValueError(
f"Expected {NUM_DENOISING_STEPS} denoising steps, got "
f"{len(pipeline.denoising_step_list)}"
)
return pipeline
def reset_caches(pipeline: CausalInferencePipeline) -> None:
device = torch.device("cuda")
if pipeline.kv_cache1 is None:
pipeline._initialize_kv_cache(1, torch.bfloat16, device)
pipeline._initialize_crossattn_cache(1, torch.bfloat16, device)
required_tokens = NUM_CHUNKS * TOKENS_PER_CHUNK
for cache in pipeline.kv_cache1:
if cache["k"].shape[1] < required_tokens:
heads, head_dim = cache["k"].shape[2:]
cache["k"] = torch.zeros(
[1, required_tokens, heads, head_dim],
dtype=torch.bfloat16,
device=device,
)
cache["v"] = torch.zeros_like(cache["k"])
for cache in pipeline.kv_cache1:
cache["global_end_index"].zero_()
cache["local_end_index"].zero_()
for cache in pipeline.crossattn_cache:
cache["is_init"] = False
@torch.inference_mode()
def generate_latents(
*,
pipeline: CausalInferencePipeline,
conditional_dict: dict[str, torch.Tensor],
seed: int,
reuse_chunk: int | None,
intervention_schedule: str,
) -> tuple[torch.Tensor, dict[str, Any]]:
"""Generate 21 latents with FFFF or exactly one reuse-schedule chunk."""
reset_caches(pipeline)
set_seed(seed)
noise = torch.randn(
1,
NUM_CHUNKS * FRAMES_PER_CHUNK,
LATENT_CHANNELS,
LATENT_HEIGHT,
LATENT_WIDTH,
dtype=torch.bfloat16,
device="cuda",
)
timesteps = pipeline.denoising_step_list.to(device="cuda")
output_chunks: list[torch.Tensor] = []
full_calls = 0
reuse_calls = 0
started = time.perf_counter()
for chunk in range(NUM_CHUNKS):
noisy_input = noise[
:, chunk * FRAMES_PER_CHUNK : (chunk + 1) * FRAMES_PER_CHUNK
]
cached_flow: torch.Tensor | None = None
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="cuda"
) * current_timestep
selected_step = (
intervention_schedule[step] if reuse_chunk == chunk else "F"
)
use_reuse = selected_step == "R"
if use_reuse:
if cached_flow is None:
raise RuntimeError("Reuse requested before any full step")
flow = cached_flow
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])
reuse_calls += 1
else:
flow, 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,
)
if reuse_chunk == chunk:
cached_flow = flow.detach().clone()
full_calls += 1
if step < NUM_DENOISING_STEPS - 1:
if denoised_pred is None:
raise RuntimeError("Denoising step did not produce 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(
[FRAMES_PER_CHUNK], dtype=torch.long, device="cuda"
),
).unflatten(0, denoised_pred.shape[:2])
if denoised_pred is None or timestep is None:
raise RuntimeError("Chunk did not produce a clean latent")
output_chunks.append(denoised_pred)
# The clean context pass is always full, as in baseline Self-Forcing.
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,
)
torch.cuda.synchronize()
return torch.cat(output_chunks, dim=1), {
"generation_time_s": time.perf_counter() - started,
"full_denoising_calls": full_calls,
"reuse_denoising_calls": reuse_calls,
"full_context_calls": NUM_CHUNKS,
}
@torch.inference_mode()
def decode_u8(pipeline: CausalInferencePipeline, latents: torch.Tensor) -> torch.Tensor:
video = pipeline.vae.decode_to_pixel(latents, use_cache=False)
frames = (
((video.squeeze(0).float() + 1.0) * 127.5)
.round_()
.clamp_(0, 255)
.to(device="cpu", dtype=torch.uint8)
.contiguous()
)
pipeline.vae.model.clear_cache()
if frames.shape != (DECODED_FRAMES, 3, 480, 832):
raise ValueError(f"Unexpected decoded shape: {tuple(frames.shape)}")
return frames
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"},
)
def gaussian_kernel(device: torch.device, channels: int = 3) -> torch.Tensor:
coordinates = torch.arange(11, device=device, dtype=torch.float32) - 5
kernel_1d = torch.exp(-coordinates.square() / (2 * 1.5**2))
kernel_1d /= kernel_1d.sum()
return torch.outer(kernel_1d, kernel_1d).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, c2 = 0.01**2, 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))
def decoded_chunk_slices() -> list[slice]:
# Wan VAE maps the first latent to one pixel frame and every subsequent
# latent to four frames. The first 3-latent chunk therefore has 9 frames;
# each later chunk has 12.
result = [slice(0, 9)]
result.extend(
slice(9 + 12 * index, 9 + 12 * (index + 1))
for index in range(NUM_CHUNKS - 1)
)
if result[-1].stop != DECODED_FRAMES:
raise AssertionError(result)
return result
@torch.inference_mode()
def frame_metrics(
*,
reference_u8: torch.Tensor,
prediction_u8: torch.Tensor,
lpips_model: torch.nn.Module,
batch_size: int,
) -> dict[str, Any]:
if reference_u8.shape != prediction_u8.shape:
raise ValueError(
f"Frame shapes differ: {tuple(reference_u8.shape)} vs "
f"{tuple(prediction_u8.shape)}"
)
device = torch.device("cuda")
kernel = gaussian_kernel(device)
mse_values: list[float] = []
psnr_values: list[float] = []
ssim_values: list[float] = []
lpips_values: list[float] = []
max_abs_values: list[int] = []
for start in range(0, len(reference_u8), batch_size):
end = min(start + batch_size, len(reference_u8))
reference = reference_u8[start:end].to(device=device, dtype=torch.float32) / 255
prediction = prediction_u8[start:end].to(device=device, dtype=torch.float32) / 255
mse = (reference - prediction).square().mean(dim=(1, 2, 3))
psnr = -10 * torch.log10(mse.clamp_min(1e-12))
ssim = ssim_per_frame(reference, prediction, kernel)
distance = lpips_model(reference.mul(2).sub(1), prediction.mul(2).sub(1)).flatten()
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())
lpips_values.extend(float(value) for value in distance.cpu())
maximum = (
reference_u8[start:end].to(torch.int16)
- prediction_u8[start:end].to(torch.int16)
).abs().flatten(1).max(1).values
max_abs_values.extend(int(value) for value in maximum)
def summarize(indices: range) -> dict[str, float | int]:
selected_mse = [mse_values[index] for index in indices]
selected_ssim = [ssim_values[index] for index in indices]
selected_lpips = [lpips_values[index] for index in indices]
mse = sum(selected_mse) / len(selected_mse)
return {
"start_frame": indices.start,
"end_frame_exclusive": indices.stop,
"num_frames": len(selected_mse),
"pixel_mse": mse,
"psnr": -10 * math.log10(max(mse, 1e-12)),
"ssim": sum(selected_ssim) / len(selected_ssim),
"lpips": sum(selected_lpips) / len(selected_lpips),
"max_abs_u8": max(max_abs_values[index] for index in indices),
}
full = summarize(range(DECODED_FRAMES))
by_chunk = []
for chunk, frame_slice in enumerate(decoded_chunk_slices()):
summary = summarize(range(frame_slice.start, frame_slice.stop))
summary["chunk"] = chunk
by_chunk.append(summary)
return {
**full,
"mse_per_frame": mse_values,
"psnr_per_frame": psnr_values,
"ssim_per_frame": ssim_values,
"lpips_per_frame": lpips_values,
"max_abs_u8_per_frame": max_abs_values,
"by_output_chunk": by_chunk,
}
def save_reference_frames(path: Path, frames: torch.Tensor, prompt_id: int) -> None:
path.parent.mkdir(parents=True, exist_ok=True)
temporary = path.with_suffix(path.suffix + ".tmp")
save_file(
{"frames": frames},
str(temporary),
metadata={"prompt_id": str(prompt_id), "pixel_domain": "uint8_rgb"},
)
os.replace(temporary, path)
def load_reference_frames(path: Path) -> torch.Tensor:
return load_file(str(path), device="cpu")["frames"]
def move_module(module: torch.nn.Module, device: str) -> None:
module.to(device=device)
if device == "cpu":
torch.cuda.empty_cache()
def write_csv(path: Path, rows: list[dict[str, Any]]) -> None:
if not rows:
return
fields: list[str] = []
for row in rows:
for key in row:
if key not in fields:
fields.append(key)
with path.open("w", newline="", encoding="utf-8") as handle:
writer = csv.DictWriter(handle, fieldnames=fields)
writer.writeheader()
writer.writerows(rows)
def aggregate(
output_dir: Path, prompt_ids: list[int], intervention_schedule: str
) -> None:
prompt_rows: list[dict[str, Any]] = []
propagation_buckets: dict[tuple[int, int], list[dict[str, Any]]] = {}
intervention_metrics: dict[int, list[dict[str, Any]]] = {
chunk: [] for chunk in range(NUM_CHUNKS)
}
completed_prompt_ids: list[int] = []
for prompt_id in prompt_ids:
prompt_dir = output_dir / f"prompt_{prompt_id:04d}"
prompt_complete = True
for reuse_chunk in range(NUM_CHUNKS):
path = prompt_dir / f"reuse_chunk_{reuse_chunk}" / "metrics.json"
if not path.exists():
prompt_complete = False
continue
value = json.loads(path.read_text(encoding="utf-8"))
if value.get("status") != "complete":
prompt_complete = False
continue
prompt_rows.append(
{
"prompt_id": prompt_id,
"reuse_chunk": reuse_chunk,
"psnr": value["psnr"],
"ssim": value["ssim"],
"lpips": value["lpips"],
"pixel_mse": value["pixel_mse"],
"max_abs_u8": value["max_abs_u8"],
"generation_time_s": value["generation_time_s"],
}
)
intervention_metrics[reuse_chunk].append(value)
for chunk_value in value["by_output_chunk"]:
propagation_buckets.setdefault(
(reuse_chunk, int(chunk_value["chunk"])), []
).append(chunk_value)
if prompt_complete:
completed_prompt_ids.append(prompt_id)
summary_rows: list[dict[str, Any]] = []
for reuse_chunk in range(NUM_CHUNKS):
rows = [row for row in prompt_rows if row["reuse_chunk"] == reuse_chunk]
if not rows:
continue
global_mse = sum(float(row["pixel_mse"]) for row in rows) / len(rows)
summary_rows.append(
{
"reuse_chunk": reuse_chunk,
"num_prompts": len(rows),
"psnr_from_global_mse": -10 * math.log10(max(global_mse, 1e-12)),
"mean_prompt_psnr": sum(float(row["psnr"]) for row in rows) / len(rows),
"mean_ssim": sum(float(row["ssim"]) for row in rows) / len(rows),
"mean_lpips": sum(float(row["lpips"]) for row in rows) / len(rows),
"mean_generation_time_s": sum(
float(row["generation_time_s"]) for row in rows
) / len(rows),
}
)
propagation_rows: list[dict[str, Any]] = []
for reuse_chunk in range(NUM_CHUNKS):
for output_chunk in range(NUM_CHUNKS):
values = propagation_buckets.get((reuse_chunk, output_chunk), [])
if not values:
continue
mse = sum(float(value["pixel_mse"]) for value in values) / len(values)
propagation_rows.append(
{
"reuse_chunk": reuse_chunk,
"output_chunk": output_chunk,
"relative_chunk": output_chunk - reuse_chunk,
"num_prompts": len(values),
"psnr_from_global_mse": -10 * math.log10(max(mse, 1e-12)),
"mean_ssim": sum(float(value["ssim"]) for value in values)
/ len(values),
"mean_lpips": sum(float(value["lpips"]) for value in values)
/ len(values),
"mean_max_abs_u8": sum(float(value["max_abs_u8"]) for value in values)
/ len(values),
}
)
affected_tail_rows: list[dict[str, Any]] = []
chunk_frames = decoded_chunk_slices()
for reuse_chunk in range(NUM_CHUNKS):
values = intervention_metrics[reuse_chunk]
if not values:
continue
start_frame = int(chunk_frames[reuse_chunk].start)
mse_values = [
frame
for value in values
for frame in value["mse_per_frame"][start_frame:]
]
ssim_values = [
frame
for value in values
for frame in value["ssim_per_frame"][start_frame:]
]
lpips_values = [
frame
for value in values
for frame in value["lpips_per_frame"][start_frame:]
]
mse = sum(mse_values) / len(mse_values)
affected_tail_rows.append(
{
"reuse_chunk": reuse_chunk,
"start_frame": start_frame,
"affected_frames_per_prompt": DECODED_FRAMES - start_frame,
"num_prompts": len(values),
"psnr_from_global_mse": -10 * math.log10(max(mse, 1e-12)),
"mean_ssim": sum(ssim_values) / len(ssim_values),
"mean_lpips": sum(lpips_values) / len(lpips_values),
}
)
output_dir.mkdir(parents=True, exist_ok=True)
write_csv(output_dir / "per_prompt.csv", prompt_rows)
write_csv(output_dir / "summary_by_reuse_chunk.csv", summary_rows)
write_csv(output_dir / "error_propagation_matrix.csv", propagation_rows)
write_csv(output_dir / "summary_affected_tail.csv", affected_tail_rows)
atomic_json(
output_dir / "summary.json",
{
"status": (
"complete"
if sorted(completed_prompt_ids) == sorted(prompt_ids)
else "partial"
),
"requested_prompt_ids": prompt_ids,
"completed_prompt_ids": completed_prompt_ids,
"num_completed_prompts": len(completed_prompt_ids),
"schedule_reference": "all chunks FFFF",
"schedule_intervention": (
f"one selected chunk {intervention_schedule}; all others FFFF"
),
"reuse_definition": (
"R reuses the selected chunk's most recent F flow prediction, "
"then applies current-timestep x0 conversion"
),
"summary_by_reuse_chunk": summary_rows,
"summary_affected_tail": affected_tail_rows,
},
)
@torch.inference_mode()
def run(args: argparse.Namespace) -> None:
output_dir = resolve(args.output_dir)
prompt_path = resolve(args.prompt_path)
prompts = read_prompts(prompt_path)
if max(args.prompt_ids) >= len(prompts):
raise ValueError(
f"Prompt ID {max(args.prompt_ids)} exceeds {len(prompts)} prompts"
)
output_dir.mkdir(parents=True, exist_ok=True)
if not args.worker:
atomic_json(
output_dir / "experiment_config.json",
{
"config_path": str(resolve(args.config_path)),
"checkpoint_path": str(resolve(args.checkpoint_path)),
"prompt_path": str(prompt_path),
"prompt_ids": args.prompt_ids,
"seed": args.seed,
"physical_gpu": args.gpu,
"use_ema": args.use_ema,
"low_memory": args.low_memory,
"latent_frames": NUM_CHUNKS * FRAMES_PER_CHUNK,
"decoded_frames": DECODED_FRAMES,
"num_chunks": NUM_CHUNKS,
"denoising_steps_per_chunk": NUM_DENOISING_STEPS,
"reference_schedule": (
"FFFFFFF at chunk level; FFFF within every chunk"
),
"intervention_schedule": (
f"one {args.intervention_schedule} chunk and "
f"{NUM_CHUNKS - 1} FFFF chunks"
),
"reuse_definition": (
"R reuses the most recent F flow prediction and applies "
"the current-timestep x0 conversion"
),
"metric_domain": (
"VAE-decoded RGB rounded to uint8 before MP4 encoding"
),
},
)
if args.aggregate_only:
aggregate(output_dir, args.prompt_ids, args.intervention_schedule)
return
pipeline = build_pipeline(args)
lpips_model = lpips.LPIPS(net="alex").eval()
if not args.low_memory:
lpips_model.to("cuda")
for prompt_offset, prompt_id in enumerate(args.prompt_ids, start=1):
prompt = prompts[prompt_id]
prompt_dir = output_dir / f"prompt_{prompt_id:04d}"
prompt_dir.mkdir(parents=True, exist_ok=True)
atomic_json(
prompt_dir / "prompt.json", {"prompt_id": prompt_id, "prompt": prompt}
)
reference_frames_path = prompt_dir / "reference_ffff_frames.safetensors"
reference_video_path = prompt_dir / "reference_ffff.mp4"
reference_metrics_path = prompt_dir / "reference_ffff.json"
print(
f"[prompt] {prompt_offset}/{len(args.prompt_ids)} id={prompt_id}",
flush=True,
)
reference_pending = args.overwrite or not reference_frames_path.exists()
pending_reuse_chunks: list[int] = []
for reuse_chunk in range(NUM_CHUNKS):
variant_dir = prompt_dir / f"reuse_chunk_{reuse_chunk}"
metrics_path = variant_dir / "metrics.json"
if metrics_path.exists() and not args.overwrite:
existing = json.loads(metrics_path.read_text(encoding="utf-8"))
if existing.get("status") == "complete":
print(f"[skip] reuse_chunk={reuse_chunk}", flush=True)
continue
pending_reuse_chunks.append(reuse_chunk)
if not reference_pending and not pending_reuse_chunks:
if not args.worker:
aggregate(output_dir, args.prompt_ids, args.intervention_schedule)
continue
# Generation phase: keep only T5, then only the generator, on CUDA.
if args.low_memory:
move_module(pipeline.text_encoder, "cuda")
conditional_dict = pipeline.text_encoder(text_prompts=[prompt])
if args.low_memory:
move_module(pipeline.text_encoder, "cpu")
move_module(pipeline.generator, "cuda")
generated_latents: dict[int | None, torch.Tensor] = {}
generation_timings: dict[int | None, dict[str, Any]] = {}
if reference_pending:
print("[run] reference FFFF", flush=True)
latents, timing = generate_latents(
pipeline=pipeline,
conditional_dict=conditional_dict,
seed=args.seed,
reuse_chunk=None,
intervention_schedule=args.intervention_schedule,
)
generated_latents[None] = latents.to(device="cpu")
generation_timings[None] = timing
del latents
for reuse_chunk in pending_reuse_chunks:
print(
f"[run] reuse_chunk={reuse_chunk}: "
f"{args.intervention_schedule}",
flush=True,
)
latents, timing = generate_latents(
pipeline=pipeline,
conditional_dict=conditional_dict,
seed=args.seed,
reuse_chunk=reuse_chunk,
intervention_schedule=args.intervention_schedule,
)
generated_latents[reuse_chunk] = latents.to(device="cpu")
generation_timings[reuse_chunk] = timing
del latents
if args.low_memory:
pipeline.kv_cache1 = None
pipeline.crossattn_cache = None
move_module(pipeline.generator, "cpu")
move_module(pipeline.vae, "cuda")
move_module(lpips_model, "cuda")
# Decode/metric phase. Keeping latents on CPU makes the phase boundary
# cheap and limits peak CUDA memory while other jobs occupy the GPU.
if reference_pending:
frames = decode_u8(pipeline, generated_latents.pop(None).to("cuda"))
save_reference_frames(reference_frames_path, frames, prompt_id)
if args.save_videos:
save_mp4(frames, reference_video_path)
atomic_json(
reference_metrics_path,
{
"status": "complete",
"prompt_id": prompt_id,
"schedule": "all chunks FFFF",
**generation_timings[None],
},
)
del frames
reference_frames = load_reference_frames(reference_frames_path)
for reuse_chunk in pending_reuse_chunks:
variant_dir = prompt_dir / f"reuse_chunk_{reuse_chunk}"
metrics_path = variant_dir / "metrics.json"
frames = decode_u8(
pipeline, generated_latents.pop(reuse_chunk).to("cuda")
)
metrics = frame_metrics(
reference_u8=reference_frames,
prediction_u8=frames,
lpips_model=lpips_model,
batch_size=args.metric_batch_size,
)
variant_dir.mkdir(parents=True, exist_ok=True)
if args.save_videos:
save_mp4(frames, variant_dir / "video.mp4")
atomic_json(
metrics_path,
{
"status": "complete",
"prompt_id": prompt_id,
"reuse_chunk": reuse_chunk,
"schedule": (
f"chunk {reuse_chunk}={args.intervention_schedule}; "
"all other chunks=FFFF"
),
"reference": "matching FFFF; same prompt, seed, and RNG sequence",
**generation_timings[reuse_chunk],
**metrics,
},
)
print(
f"[metric] chunk={reuse_chunk} PSNR={metrics['psnr']:.4f} "
f"SSIM={metrics['ssim']:.6f} LPIPS={metrics['lpips']:.6f}",
flush=True,
)
del frames
torch.cuda.empty_cache()
if args.low_memory:
move_module(pipeline.vae, "cpu")
move_module(lpips_model, "cpu")
del reference_frames, conditional_dict, generated_latents
if not args.worker:
aggregate(output_dir, args.prompt_ids, args.intervention_schedule)
if not args.worker:
aggregate(output_dir, args.prompt_ids, args.intervention_schedule)
print(f"[done] {output_dir}", flush=True)
def main() -> None:
args = parse_args()
run(args)
if __name__ == "__main__":
main()