diff --git a/scripts/aggregate_layer17_dynamic_runs.py b/scripts/aggregate_layer17_dynamic_runs.py new file mode 100644 index 0000000000000000000000000000000000000000..f5c6cb0aa9605e6085c19871c48eeac3994450e2 --- /dev/null +++ b/scripts/aggregate_layer17_dynamic_runs.py @@ -0,0 +1,134 @@ +#!/usr/bin/env python3 +"""Aggregate per-run dynamic-gate JSON files without loading GPU models.""" + +from __future__ import annotations + +import argparse +import csv +import json +from pathlib import Path + + +NUMERIC_FIELDS = [ + "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", +] + + +def write_csv(path: Path, rows: list[dict], fields: list[str]) -> None: + path.parent.mkdir(parents=True, exist_ok=True) + 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) + temporary.replace(path) + + +def main() -> None: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--test_dir", type=Path, required=True) + parser.add_argument("--expected_prompts", type=int, default=10) + parser.add_argument( + "--select_targets", + type=int, + nargs="*", + default=None, + help="Also select the best dynamic validation row for each target budget.", + ) + args = parser.parse_args() + per_run = args.test_dir / "per_run" + records = [ + json.loads(path.read_text(encoding="utf-8")) + for path in sorted(per_run.glob("*/prompt_*.json")) + ] + if not records: + raise ValueError(f"No per-run JSON files under {per_run}") + + flat_fields = sorted({key for row in records for key in row if key != "decisions"}) + write_csv( + args.test_dir / "runs.csv", + [{key: row.get(key) for key in flat_fields} for row in records], + flat_fields, + ) + + summary = [] + for name in sorted({str(row["config_name"]) for row in records}): + selected = [row for row in records if row["config_name"] == name] + if len(selected) != args.expected_prompts: + raise ValueError( + f"{name}: expected {args.expected_prompts} prompts, got {len(selected)}" + ) + 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_FIELDS: + 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_FIELDS, + ] + write_csv(args.test_dir / "summary.csv", summary, fields) + if args.select_targets: + selected_dynamic = [] + for target in args.select_targets: + candidates = [ + row for row in summary + if row["policy"] == "dynamic" + and int(row["target_accepts"]) == target + ] + if not candidates: + raise ValueError(f"No dynamic candidates for target {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]) + (args.test_dir / "selected.json").write_text( + json.dumps( + { + "selection_rule": ( + "within target accepted calls +/-0.5, lowest validation " + "tail LPIPS; then budget distance and lower beta" + ), + "selected_dynamic": selected_dynamic, + }, + indent=2, + ) + + "\n", + encoding="utf-8", + ) + print(json.dumps(summary, indent=2)) + + +if __name__ == "__main__": + main() diff --git a/scripts/analyze_feature_cache.py b/scripts/analyze_feature_cache.py new file mode 100644 index 0000000000000000000000000000000000000000..e4ba1503f0e9a880b801fb2403cb61380cc84380 --- /dev/null +++ b/scripts/analyze_feature_cache.py @@ -0,0 +1,1623 @@ +#!/usr/bin/env python3 +"""Analyze timestep and cross-chunk feature reuse in Self-Forcing. + +The script runs the released four-step causal checkpoint without changing its +outputs. Forward hooks capture sampled DiT hidden states, residual updates, the +final velocity, and a low-dimensional dense projection used for optical-flow +alignment. It then produces: + +* within-chunk and cross-chunk feature-pair metrics; +* leave-one-prompt-out channel-wise probes for conditional chunk information; +* shuffled, wrong-step, distant-chunk, zero, and noise controls; +* motion-stratified raw/global/dense-flow alignment measurements. + +Each prompt is saved independently, so interrupted generation can be resumed. +""" + +from __future__ import annotations + +import argparse +import csv +import json +import math +import os +import random +import sys +import time +from collections import defaultdict +from pathlib import Path +from typing import Any, Iterable + + +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 cv2 +import matplotlib + +matplotlib.use("Agg") +import matplotlib.pyplot as plt +import numpy as np +import torch +import torch.nn.functional as F +from omegaconf import OmegaConf + +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 utils.misc import set_seed + + +EXPECTED_LATENT_HEIGHT = 60 +EXPECTED_LATENT_WIDTH = 104 +FRAME_TOKEN_HEIGHT = 30 +FRAME_TOKEN_WIDTH = 52 +FRAME_SEQ_LENGTH = FRAME_TOKEN_HEIGHT * FRAME_TOKEN_WIDTH + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser( + description="Self-Forcing cross-chunk feature-cache analysis" + ) + 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, required=True) + parser.add_argument("--num_prompts", type=int, default=3) + parser.add_argument("--num_frames", type=int, default=21) + parser.add_argument("--seed", type=int, default=20260728) + parser.add_argument( + "--same_seed", + action="store_true", + help="Reset every prompt to --seed for paired cross-model evaluation.", + ) + parser.add_argument( + "--layers", type=int, nargs="+", default=[0, 9, 19, 29] + ) + parser.add_argument("--max_tokens", type=int, default=256) + parser.add_argument("--projection_dim", type=int, default=64) + parser.add_argument("--ridge", type=float, default=1e-4) + parser.add_argument("--use_ema", action="store_true", default=True) + parser.add_argument("--no_ema", action="store_false", dest="use_ema") + parser.add_argument("--overwrite", action="store_true") + parser.add_argument( + "--cosine_only", + action="store_true", + help="Generate pair metrics/heatmap only; skip probes and motion analysis.", + ) + parser.add_argument( + "--analysis_only", + action="store_true", + help="Skip model loading and analyze existing per-prompt snapshots.", + ) + parser.add_argument( + "--save_preview", + action="store_true", + help="Save a compact MP4 preview when torchvision video IO is available.", + ) + args = parser.parse_args() + if args.num_frames % 3 != 0: + parser.error("--num_frames must be divisible by the configured 3-frame chunk") + if args.num_prompts < 3 and not args.cosine_only: + parser.error("--num_prompts must be at least 3 for held-out/shuffle controls") + return args + + +def resolve_path(path: Path) -> Path: + return path if path.is_absolute() else REPO_ROOT / path + + +def write_csv(path: Path, rows: list[dict[str, Any]]) -> None: + if not rows: + return + path.parent.mkdir(parents=True, exist_ok=True) + 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, extrasaction="ignore") + writer.writeheader() + writer.writerows(rows) + + +def read_prompts(path: Path, count: int) -> list[str]: + with path.open("r", encoding="utf-8") as handle: + prompts = [line.strip() for line in handle if line.strip()] + if len(prompts) < count: + raise ValueError(f"Requested {count} prompts, found {len(prompts)} in {path}") + return prompts[:count] + + +def regular_grid_indices( + frames: int, height: int, width: int, max_tokens: int, device: torch.device +) -> tuple[torch.Tensor, torch.Tensor]: + total = frames * height * width + if max_tokens >= total: + coords = torch.cartesian_prod( + torch.arange(frames, device=device), + torch.arange(height, device=device), + torch.arange(width, device=device), + ) + else: + per_frame = max(1, max_tokens // frames) + h_count = max( + 1, min(height, int(round(math.sqrt(per_frame * height / width)))) + ) + w_count = max(1, min(width, per_frame // h_count)) + while frames * h_count * w_count > max_tokens and w_count > 1: + w_count -= 1 + while frames * h_count * w_count > max_tokens and h_count > 1: + h_count -= 1 + hs = ( + torch.linspace(0, height - 1, h_count, device=device) + .round() + .long() + .unique() + ) + ws = ( + torch.linspace(0, width - 1, w_count, device=device) + .round() + .long() + .unique() + ) + coords = torch.cartesian_prod( + torch.arange(frames, device=device), hs, ws + ) + flat = ( + coords[:, 0] * height * width + coords[:, 1] * width + coords[:, 2] + ) + return flat.long(), coords.long() + + +class FeatureRecorder: + def __init__( + self, + model: torch.nn.Module, + layers: list[int], + denoising_timesteps: Iterable[float], + num_frame_per_block: int, + max_tokens: int, + projection_dim: int, + ) -> None: + self.model = model + self.layers = sorted(set(int(value) for value in layers)) + self.timesteps = [float(value) for value in denoising_timesteps] + self.num_frame_per_block = int(num_frame_per_block) + self.max_tokens = int(max_tokens) + self.projection_dim = int(projection_dim) + self.projection_layer = self.layers[-1] + self.records: dict[str, dict[str, torch.Tensor]] = defaultdict(dict) + self.projected: dict[str, torch.Tensor] = {} + self.projected_by_layer: dict[str, torch.Tensor] = {} + self.sample_coords: dict[str, torch.Tensor] = {} + self.current: dict[str, Any] = {"active": False} + self.handles: list[Any] = [] + self._projection_cache: dict[tuple[int, str], torch.Tensor] = {} + self._register() + + def _register(self) -> None: + self.handles.append( + self.model.register_forward_pre_hook(self._model_pre_hook, with_kwargs=True) + ) + self.handles.append(self.model.register_forward_hook(self._model_output_hook)) + for layer in self.layers: + if layer < 0 or layer >= len(self.model.blocks): + raise ValueError( + f"Layer {layer} outside model block range 0..{len(self.model.blocks)-1}" + ) + self.handles.append( + self.model.blocks[layer].register_forward_hook( + self._make_block_hook(layer) + ) + ) + + def close(self) -> None: + for handle in self.handles: + handle.remove() + self.handles.clear() + + def reset(self) -> None: + self.records = defaultdict(dict) + self.projected = {} + self.projected_by_layer = {} + self.sample_coords = {} + self.current = {"active": False} + + def _model_pre_hook( + self, _module: torch.nn.Module, _args: tuple[Any, ...], kwargs: dict[str, Any] + ) -> None: + timestep = kwargs.get("t") + current_start = int(kwargs.get("current_start", 0) or 0) + if not isinstance(timestep, torch.Tensor) or timestep.numel() == 0: + self.current = {"active": False} + return + value = float(timestep.detach().float().reshape(-1)[0].item()) + distances = [abs(value - expected) for expected in self.timesteps] + step = int(np.argmin(distances)) + if distances[step] > 0.5: + self.current = {"active": False, "timestep": value} + return + frames = int(timestep.shape[-1]) if timestep.ndim > 1 else 1 + if frames != self.num_frame_per_block: + self.current = {"active": False, "timestep": value} + return + start_frame = current_start // FRAME_SEQ_LENGTH + chunk = start_frame // self.num_frame_per_block + self.current = { + "active": True, + "chunk": int(chunk), + "step": step, + "timestep": value, + "frames": frames, + } + + def _key(self) -> str: + return f"{self.current['chunk']}:{self.current['step']}" + + def _hidden_indices( + self, tokens: torch.Tensor + ) -> tuple[torch.Tensor, torch.Tensor]: + frames = int(self.current["frames"]) + if tokens.shape[1] != frames * FRAME_TOKEN_HEIGHT * FRAME_TOKEN_WIDTH: + raise ValueError( + f"Unexpected hidden token count {tokens.shape[1]} for {frames} frames" + ) + return regular_grid_indices( + frames, + FRAME_TOKEN_HEIGHT, + FRAME_TOKEN_WIDTH, + self.max_tokens, + tokens.device, + ) + + def _projection(self, dim: int, device: torch.device) -> torch.Tensor: + cache_key = (dim, str(device)) + if cache_key not in self._projection_cache: + generator = torch.Generator(device="cpu").manual_seed(20260728 + dim) + signs = torch.randint( + 0, + 2, + (dim, self.projection_dim), + generator=generator, + dtype=torch.int8, + ) + projection = ( + signs.float().mul_(2).sub_(1).div_(math.sqrt(self.projection_dim)) + ) + self._projection_cache[cache_key] = projection.to(device) + return self._projection_cache[cache_key] + + def _make_block_hook(self, layer: int): + def hook( + _module: torch.nn.Module, + inputs: tuple[torch.Tensor, ...], + output: torch.Tensor, + ) -> None: + if not self.current.get("active", False): + return + if not inputs or not isinstance(output, torch.Tensor): + return + key = self._key() + hidden_input = inputs[0] + indices, coords = self._hidden_indices(output) + hidden = output[0].index_select(0, indices) + delta = (output - hidden_input)[0].index_select(0, indices) + self.records[f"block_{layer}_hidden"][key] = ( + hidden.detach().to(dtype=torch.float16, device="cpu") + ) + self.records[f"block_{layer}_delta"][key] = ( + delta.detach().to(dtype=torch.float16, device="cpu") + ) + self.sample_coords["hidden"] = coords.detach().cpu() + + if layer in self.layers: + projection = self._projection(output.shape[-1], output.device) + dense = torch.matmul(output[0].float(), projection) + frames = int(self.current["frames"]) + dense = dense.reshape( + frames, + FRAME_TOKEN_HEIGHT, + FRAME_TOKEN_WIDTH, + self.projection_dim, + ) + dense_cpu = dense.detach().to(dtype=torch.float16, device="cpu") + self.projected_by_layer[f"{layer}:{key}"] = dense_cpu + # Preserve the legacy key layout for existing final-layer analyses. + if layer == self.projection_layer: + self.projected[key] = dense_cpu + + return hook + + def _model_output_hook( + self, + _module: torch.nn.Module, + _inputs: tuple[Any, ...], + output: torch.Tensor, + ) -> None: + if not self.current.get("active", False): + return + if not isinstance(output, torch.Tensor) or output.ndim != 5: + return + batch, channels, frames, height, width = output.shape + if batch != 1: + raise ValueError(f"Analysis expects batch size 1, got {batch}") + tokens = output.permute(0, 2, 3, 4, 1).reshape( + batch, frames * height * width, channels + ) + indices, coords = regular_grid_indices( + frames, height, width, self.max_tokens, output.device + ) + sampled = tokens[0].index_select(0, indices) + self.records["dit_output"][self._key()] = sampled.detach().to( + dtype=torch.float16, device="cpu" + ) + self.sample_coords["dit_output"] = coords.detach().cpu() + + def state_dict(self) -> dict[str, Any]: + return { + "timesteps": self.timesteps, + "layers": self.layers, + "projection_layer": self.projection_layer, + "projection_dim": self.projection_dim, + "records": {stage: dict(values) for stage, values in self.records.items()}, + "projected": dict(self.projected), + "projected_by_layer": dict(self.projected_by_layer), + "sample_coords": dict(self.sample_coords), + } + + +def build_pipeline(args: argparse.Namespace) -> tuple[CausalInferencePipeline, Any]: + config = OmegaConf.load(resolve_path(args.config_path)) + default_config = OmegaConf.load(REPO_ROOT / "configs/default_config.yaml") + config = OmegaConf.merge(default_config, config) + device = torch.device("cuda") + pipeline = CausalInferencePipeline(config, device=device) + checkpoint = torch.load( + resolve_path(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) + pipeline.text_encoder.to(device=device) + pipeline.generator.to(device=device) + pipeline.vae.to(device=device) + pipeline.eval() + return pipeline, config + + +def downsample_anchors(video: torch.Tensor, latent_frames: int) -> np.ndarray: + value = video[0].detach().float().cpu() + frame_count = value.shape[0] + if frame_count == latent_frames: + indices = np.arange(latent_frames) + else: + indices = np.linspace(0, frame_count - 1, latent_frames).round().astype(int) + anchors = value[indices].permute(0, 2, 3, 1).clamp(0, 1).numpy() + result = [] + for frame in anchors: + frame_u8 = np.uint8(np.round(frame * 255.0)) + result.append(cv2.resize(frame_u8, (416, 240), interpolation=cv2.INTER_AREA)) + return np.stack(result) + + +def maybe_save_preview(path: Path, anchors: np.ndarray) -> None: + try: + from torchvision.io import write_video + + repeated = np.repeat(anchors, 4, axis=0) + tensor = torch.from_numpy(repeated) + write_video(str(path), tensor, fps=16) + except Exception as error: + print(f"[preview] skipped: {error}", flush=True) + + +@torch.inference_mode() +def generate_snapshots(args: argparse.Namespace) -> list[Path]: + output_dir = args.output_dir + runs_dir = output_dir / "runs" + runs_dir.mkdir(parents=True, exist_ok=True) + prompts = read_prompts(resolve_path(args.prompt_path), args.num_prompts) + expected_paths = [runs_dir / f"prompt_{index:02d}.pt" for index in range(len(prompts))] + missing = [ + path + for path in expected_paths + if args.overwrite or not path.exists() + ] + if not missing: + print("[generation] all prompt snapshots already exist", flush=True) + return expected_paths + + pipeline, config = build_pipeline(args) + denoising_timesteps = [ + float(value) for value in pipeline.denoising_step_list.detach().cpu().tolist() + ] + recorder = FeatureRecorder( + model=pipeline.generator.model, + layers=args.layers, + denoising_timesteps=denoising_timesteps, + num_frame_per_block=pipeline.num_frame_per_block, + max_tokens=args.max_tokens, + projection_dim=args.projection_dim, + ) + metadata = { + "physical_gpu": args.gpu, + "config_path": str(resolve_path(args.config_path)), + "checkpoint_path": str(resolve_path(args.checkpoint_path)), + "use_ema": args.use_ema, + "num_prompts": args.num_prompts, + "num_frames": args.num_frames, + "num_frame_per_block": pipeline.num_frame_per_block, + "denoising_timesteps": denoising_timesteps, + "layers": args.layers, + "max_tokens": args.max_tokens, + "projection_dim": args.projection_dim, + "seed": args.seed, + "dtype": "bfloat16", + } + (output_dir / "experiment_config.json").write_text( + json.dumps(metadata, indent=2, ensure_ascii=False) + "\n", + encoding="utf-8", + ) + + try: + for index, (prompt, path) in enumerate(zip(prompts, expected_paths)): + if path.exists() and not args.overwrite: + print(f"[generation] skip existing {path.name}", flush=True) + continue + recorder.reset() + run_seed = args.seed if args.same_seed else args.seed + index + set_seed(run_seed) + noise = torch.randn( + 1, + args.num_frames, + 16, + EXPECTED_LATENT_HEIGHT, + EXPECTED_LATENT_WIDTH, + device="cuda", + dtype=torch.bfloat16, + ) + torch.cuda.reset_peak_memory_stats() + torch.cuda.synchronize() + start = time.perf_counter() + print( + f"[generation] prompt {index + 1}/{len(prompts)} seed={run_seed}", + flush=True, + ) + video, latents = pipeline.inference( + noise=noise, + text_prompts=[prompt], + return_latents=True, + initial_latent=None, + low_memory=False, + ) + torch.cuda.synchronize() + elapsed = time.perf_counter() - start + peak_gib = torch.cuda.max_memory_allocated() / (1024**3) + anchors = downsample_anchors(video, args.num_frames) + state = { + "run_index": index, + "prompt": prompt, + "seed": run_seed, + "elapsed_s": elapsed, + "peak_gpu_gib": peak_gib, + "num_frames": args.num_frames, + "num_frame_per_block": pipeline.num_frame_per_block, + "latents": latents[0].detach().to( + dtype=torch.float16, device="cpu" + ), + **recorder.state_dict(), + } + torch.save(state, path) + np.savez_compressed(path.with_suffix(".anchors.npz"), frames=anchors) + if args.save_preview: + maybe_save_preview(path.with_suffix(".mp4"), anchors) + print( + f"[generation] saved {path.name}: {elapsed:.1f}s, peak={peak_gib:.1f} GiB", + flush=True, + ) + del video, latents, noise, state + pipeline.vae.model.clear_cache() + torch.cuda.empty_cache() + finally: + recorder.close() + return expected_paths + + +def load_runs(paths: list[Path]) -> list[dict[str, Any]]: + runs = [] + for path in paths: + if not path.exists(): + raise FileNotFoundError(path) + run = torch.load(path, map_location="cpu", weights_only=False) + anchor_path = path.with_suffix(".anchors.npz") + if not anchor_path.exists(): + raise FileNotFoundError(anchor_path) + run["anchors"] = np.load(anchor_path, allow_pickle=False)["frames"] + run["path"] = str(path) + runs.append(run) + return runs + + +def feature(run: dict[str, Any], stage: str, chunk: int, step: int) -> torch.Tensor: + return run["records"][stage][f"{chunk}:{step}"].float() + + +def projected_feature( + run: dict[str, Any], chunk: int, step: int +) -> torch.Tensor: + return run["projected"][f"{chunk}:{step}"].float() + + +def available_chunks(run: dict[str, Any], stage: str) -> list[int]: + return sorted( + {int(key.split(":")[0]) for key in run["records"][stage].keys()} + ) + + +def available_steps(run: dict[str, Any], stage: str) -> list[int]: + return sorted( + {int(key.split(":")[1]) for key in run["records"][stage].keys()} + ) + + +def pair_metrics( + reference: torch.Tensor, + target: torch.Tensor, + compute_cka: bool = True, +) -> dict[str, float]: + if reference.shape != target.shape: + raise ValueError(f"Pair shape mismatch: {reference.shape} vs {target.shape}") + eps = 1e-8 + x = reference.float() + y = target.float() + xf = x.reshape(-1) + yf = y.reshape(-1) + diff = yf - xf + cosine = F.cosine_similarity(xf[None], yf[None], dim=1, eps=eps)[0] + xc_flat = xf - xf.mean() + yc_flat = yf - yf.mean() + centered_cosine = torch.dot(xc_flat, yc_flat) / ( + torch.linalg.vector_norm(xc_flat) + * torch.linalg.vector_norm(yc_flat) + + eps + ) + token_cosine = F.cosine_similarity(x, y, dim=1, eps=eps) + rel_l2 = diff.square().mean().sqrt() / (xf.square().mean().sqrt() + eps) + nmse = diff.square().mean() / (yc_flat.square().mean() + eps) + + if compute_cka: + xc = x - x.mean(dim=0, keepdim=True) + yc = y - y.mean(dim=0, keepdim=True) + gram_x = xc @ xc.T + gram_y = yc @ yc.T + cka = (gram_x * gram_y).sum() / ( + (gram_x.square().sum() * gram_y.square().sum()).sqrt() + eps + ) + else: + cka = torch.tensor(float("nan")) + quantiles = torch.quantile( + token_cosine, + torch.tensor([0.1, 0.5, 0.9], dtype=token_cosine.dtype), + ) + return { + "cosine": float(cosine), + "centered_cosine": float(centered_cosine), + "linear_cka": float(cka), + "rel_l2": float(rel_l2), + "nmse": float(nmse), + "token_cosine_mean": float(token_cosine.mean()), + "token_cosine_p10": float(quantiles[0]), + "token_cosine_p50": float(quantiles[1]), + "token_cosine_p90": float(quantiles[2]), + "reference_rms": float(xf.square().mean().sqrt()), + "target_rms": float(yf.square().mean().sqrt()), + } + + +def collect_pair_rows( + runs: list[dict[str, Any]], cosine_only: bool = False +) -> list[dict[str, Any]]: + rows: list[dict[str, Any]] = [] + stages = sorted(runs[0]["records"]) + if cosine_only: + stages = [ + stage + for stage in stages + if stage.startswith("block_") and stage.endswith("_hidden") + ] + for run_index, run in enumerate(runs): + shuffled_run = runs[(run_index + 1) % len(runs)] + for stage in stages: + chunks = available_chunks(run, stage) + steps = available_steps(run, stage) + + def append( + comparison: str, + ref_run: dict[str, Any], + ref_chunk: int, + ref_step: int, + target_chunk: int, + target_step: int, + ) -> None: + metrics = pair_metrics( + feature(ref_run, stage, ref_chunk, ref_step), + feature(run, stage, target_chunk, target_step), + compute_cka=not cosine_only, + ) + rows.append( + { + "run": run_index, + "comparison": comparison, + "stage": stage, + "reference_chunk": ref_chunk, + "target_chunk": target_chunk, + "reference_step": ref_step, + "target_step": target_step, + **metrics, + } + ) + + for chunk in chunks: + for step in steps[1:]: + append( + "within_adjacent", + run, + chunk, + step - 1, + chunk, + step, + ) + if chunk < 1: + continue + for step in steps: + append( + "cross_same", + run, + chunk - 1, + step, + chunk, + step, + ) + if cosine_only: + continue + append( + "cross_video_shuffle", + shuffled_run, + chunk - 1, + step, + chunk, + step, + ) + if step > 0: + append( + "cross_wrong_step", + run, + chunk - 1, + step - 1, + chunk, + step, + ) + if chunk > 1: + append( + "cross_distant", + run, + chunk - 2, + step, + chunk, + step, + ) + return rows + + +def predictor_columns( + name: str, + within: torch.Tensor, + cross: torch.Tensor, + distant: torch.Tensor, + wrong: torch.Tensor, + batch: torch.Tensor, + noise_seed: int, +) -> list[torch.Tensor]: + ones = torch.ones_like(within) + shifted = cross.roll(shifts=max(1, cross.shape[0] // 2), dims=0) + generator = torch.Generator(device="cpu").manual_seed(noise_seed) + noise = torch.randn( + cross.shape, generator=generator, dtype=cross.dtype + ) + noise = noise * cross.std(dim=0, keepdim=True).clamp_min(1e-6) + noise = noise + cross.mean(dim=0, keepdim=True) + mapping = { + "within_affine": [within, ones], + "within_quadratic": [within, within.square(), ones], + "cross_affine": [cross, ones], + "fusion_same": [within, cross, ones], + "fusion_distant": [within, distant, ones], + "fusion_token_shift": [within, shifted, ones], + "fusion_wrong_step": [within, wrong, ones], + "fusion_batch_shuffle": [within, batch, ones], + "fusion_zero": [within, torch.zeros_like(cross), ones], + "fusion_noise": [within, noise, ones], + } + return mapping[name] + + +PROBE_NAMES = [ + "within_affine", + "within_quadratic", + "cross_affine", + "fusion_same", + "fusion_distant", + "fusion_token_shift", + "fusion_wrong_step", + "fusion_batch_shuffle", + "fusion_zero", + "fusion_noise", +] + + +def gather_probe_data( + runs: list[dict[str, Any]], + run_indices: list[int], + stage: str, + step: int, +) -> dict[str, torch.Tensor]: + buckets: dict[str, list[torch.Tensor]] = defaultdict(list) + for run_index in run_indices: + run = runs[run_index] + other = runs[(run_index + 1) % len(runs)] + chunks = available_chunks(run, stage) + for chunk in chunks: + # Use c >= 2 for every probe so correct, distant, and all other + # controls are evaluated on exactly the same target tokens. + if chunk < 2: + continue + buckets["target"].append(feature(run, stage, chunk, step)) + buckets["within"].append(feature(run, stage, chunk, step - 1)) + buckets["cross"].append(feature(run, stage, chunk - 1, step)) + buckets["distant"].append(feature(run, stage, chunk - 2, step)) + buckets["wrong"].append(feature(run, stage, chunk - 1, step - 1)) + buckets["batch"].append(feature(other, stage, chunk - 1, step)) + if not buckets: + raise ValueError(f"No probe data for stage={stage}, step={step}") + return {key: torch.cat(values, dim=0).float() for key, values in buckets.items()} + + +def fit_channelwise_probe( + columns: list[torch.Tensor], + target: torch.Tensor, + ridge: float, +) -> torch.Tensor: + design = torch.stack(columns, dim=-1).double() + y = target.double() + gram = torch.einsum("ndp,ndq->dpq", design, design) + rhs = torch.einsum("ndp,nd->dp", design, y) + feature_count = gram.shape[-1] + diagonal_scale = ( + gram.diagonal(dim1=-2, dim2=-1).mean(dim=-1).clamp_min(1e-8) + ) + regularizer = ( + torch.eye(feature_count, dtype=gram.dtype)[None] + * (ridge * diagonal_scale)[:, None, None] + ) + regularizer[:, -1, -1] = 0.0 + try: + weights = torch.linalg.solve(gram + regularizer, rhs.unsqueeze(-1)).squeeze(-1) + except torch.linalg.LinAlgError: + weights = ( + torch.linalg.pinv(gram + regularizer) @ rhs.unsqueeze(-1) + ).squeeze(-1) + return weights.float() + + +def apply_channelwise_probe( + columns: list[torch.Tensor], weights: torch.Tensor +) -> torch.Tensor: + design = torch.stack(columns, dim=-1).float() + return torch.einsum("ndp,dp->nd", design, weights) + + +def prediction_metrics(prediction: torch.Tensor, target: torch.Tensor) -> dict[str, float]: + eps = 1e-8 + pred = prediction.float() + y = target.float() + error = pred - y + mse = error.square().mean() + variance = (y - y.mean()).square().mean() + nrmse = mse.sqrt() / (variance.sqrt() + eps) + r2 = 1.0 - mse / (variance + eps) + cosine = F.cosine_similarity( + pred.reshape(1, -1), y.reshape(1, -1), dim=1, eps=eps + )[0] + return { + "mse": float(mse), + "nrmse": float(nrmse), + "r2": float(r2), + "cosine": float(cosine), + } + + +def run_conditional_probes( + runs: list[dict[str, Any]], ridge: float +) -> list[dict[str, Any]]: + rows: list[dict[str, Any]] = [] + stages = sorted(runs[0]["records"]) + steps = available_steps(runs[0], stages[0]) + for stage in stages: + for step in steps[1:]: + for held_out in range(len(runs)): + train_indices = [index for index in range(len(runs)) if index != held_out] + train = gather_probe_data(runs, train_indices, stage, step) + test = gather_probe_data(runs, [held_out], stage, step) + for probe_name in PROBE_NAMES: + train_columns = predictor_columns( + probe_name, + train["within"], + train["cross"], + train["distant"], + train["wrong"], + train["batch"], + noise_seed=1000 + held_out * 100 + step, + ) + test_columns = predictor_columns( + probe_name, + test["within"], + test["cross"], + test["distant"], + test["wrong"], + test["batch"], + noise_seed=2000 + held_out * 100 + step, + ) + weights = fit_channelwise_probe( + train_columns, train["target"], ridge=ridge + ) + prediction = apply_channelwise_probe(test_columns, weights) + rows.append( + { + "held_out_run": held_out, + "stage": stage, + "step": step, + "probe": probe_name, + "train_tokens": int(train["target"].shape[0]), + "test_tokens": int(test["target"].shape[0]), + **prediction_metrics(prediction, test["target"]), + } + ) + baseline_lookup = { + (row["held_out_run"], row["stage"], row["step"], row["probe"]): row + for row in rows + if row["probe"] in {"within_affine", "within_quadratic"} + } + for row in rows: + key = (row["held_out_run"], row["stage"], row["step"]) + for baseline_name in ("within_affine", "within_quadratic"): + baseline = baseline_lookup[(*key, baseline_name)] + row[f"mse_reduction_vs_{baseline_name}"] = ( + baseline["mse"] - row["mse"] + ) / max(baseline["mse"], 1e-12) + row[f"r2_gain_vs_{baseline_name}"] = row["r2"] - baseline["r2"] + return rows + + +def gray(frame: np.ndarray) -> np.ndarray: + return cv2.cvtColor(frame, cv2.COLOR_RGB2GRAY) + + +def farneback(source: np.ndarray, target: np.ndarray) -> np.ndarray: + return cv2.calcOpticalFlowFarneback( + gray(source), + gray(target), + None, + pyr_scale=0.5, + levels=4, + winsize=21, + iterations=5, + poly_n=7, + poly_sigma=1.5, + flags=0, + ) + + +def resize_flow(flow: np.ndarray, height: int, width: int) -> torch.Tensor: + source_height, source_width = flow.shape[:2] + resized = cv2.resize(flow, (width, height), interpolation=cv2.INTER_AREA) + resized[..., 0] *= width / source_width + resized[..., 1] *= height / source_height + return torch.from_numpy(resized).permute(2, 0, 1).float() + + +def warp_feature( + source: torch.Tensor, target_to_source_flow: torch.Tensor +) -> tuple[torch.Tensor, torch.Tensor]: + height, width, _ = source.shape + yy, xx = torch.meshgrid( + torch.arange(height, dtype=torch.float32), + torch.arange(width, dtype=torch.float32), + indexing="ij", + ) + sample_x = xx + target_to_source_flow[0] + sample_y = yy + target_to_source_flow[1] + grid = torch.stack( + [ + 2.0 * sample_x / max(width - 1, 1) - 1.0, + 2.0 * sample_y / max(height - 1, 1) - 1.0, + ], + dim=-1, + )[None] + value = source.permute(2, 0, 1)[None].float() + warped = F.grid_sample( + value, grid, mode="bilinear", padding_mode="zeros", align_corners=True + )[0].permute(1, 2, 0) + mask = ( + (sample_x >= 0) + & (sample_x <= width - 1) + & (sample_y >= 0) + & (sample_y <= height - 1) + ) + return warped, mask + + +def masked_cosine( + left: torch.Tensor, right: torch.Tensor, mask: torch.Tensor | None = None +) -> float: + value = F.cosine_similarity(left.float(), right.float(), dim=-1, eps=1e-8) + if mask is not None: + if not bool(mask.any()): + return float("nan") + value = value[mask] + return float(value.mean()) + + +def collect_motion_rows(runs: list[dict[str, Any]]) -> list[dict[str, Any]]: + rows: list[dict[str, Any]] = [] + for run_index, run in enumerate(runs): + anchors = run["anchors"] + chunk_size = int(run["num_frame_per_block"]) + chunk_count = int(run["num_frames"]) // chunk_size + steps = sorted( + {int(key.split(":")[1]) for key in run["projected"].keys()} + ) + for chunk in range(1, chunk_count): + previous_anchor_index = chunk * chunk_size - 1 + previous_frame = anchors[previous_anchor_index] + boundary_flow = farneback( + previous_frame, anchors[chunk * chunk_size] + ) + median_flow = np.median( + boundary_flow.reshape(-1, 2), axis=0 + ) + camera_motion = float(np.linalg.norm(median_flow)) + residual = boundary_flow - median_flow[None, None] + object_motion = float( + np.linalg.norm(residual, axis=-1).mean() + ) + total_motion = float( + np.linalg.norm(boundary_flow, axis=-1).mean() + ) + + for step in steps: + source_map = projected_feature(run, chunk - 1, step)[-1] + target_map = projected_feature(run, chunk, step) + same_slot_source = projected_feature(run, chunk - 1, step) + per_slot: list[dict[str, float]] = [] + for slot in range(chunk_size): + target = target_map[slot] + same_slot = same_slot_source[slot] + raw_boundary = source_map + target_frame = anchors[chunk * chunk_size + slot] + backward = farneback(target_frame, previous_frame) + feature_flow = resize_flow( + backward, FRAME_TOKEN_HEIGHT, FRAME_TOKEN_WIDTH + ) + global_flow = torch.zeros_like(feature_flow) + global_flow[0].fill_(float(np.median(feature_flow[0].numpy()))) + global_flow[1].fill_(float(np.median(feature_flow[1].numpy()))) + global_aligned, global_mask = warp_feature( + source_map, global_flow + ) + flow_aligned, flow_mask = warp_feature(source_map, feature_flow) + per_slot.append( + { + "same_slot_cosine": masked_cosine(target, same_slot), + "boundary_raw_cosine": masked_cosine( + target, raw_boundary + ), + "global_aligned_cosine": masked_cosine( + target, global_aligned, global_mask + ), + "flow_aligned_cosine": masked_cosine( + target, flow_aligned, flow_mask + ), + "valid_flow_ratio": float(flow_mask.float().mean()), + } + ) + rows.append( + { + "run": run_index, + "chunk": chunk, + "step": step, + "total_motion": total_motion, + "camera_motion": camera_motion, + "object_motion": object_motion, + **{ + key: float(np.nanmean([item[key] for item in per_slot])) + for key in per_slot[0] + }, + } + ) + motion_values = np.asarray([row["total_motion"] for row in rows]) + if len(motion_values) >= 3: + low, high = np.quantile(motion_values, [1 / 3, 2 / 3]) + for row in rows: + value = row["total_motion"] + row["motion_bin"] = "low" if value <= low else "high" if value > high else "medium" + return rows + + +def group_mean( + rows: list[dict[str, Any]], keys: list[str], metrics: list[str] +) -> list[dict[str, Any]]: + groups: dict[tuple[Any, ...], list[dict[str, Any]]] = defaultdict(list) + for row in rows: + groups[tuple(row[key] for key in keys)].append(row) + result = [] + for group, values in sorted(groups.items(), key=lambda item: tuple(map(str, item[0]))): + output = {key: value for key, value in zip(keys, group)} + output["count"] = len(values) + for metric in metrics: + finite = [ + float(row[metric]) + for row in values + if metric in row and math.isfinite(float(row[metric])) + ] + output[metric] = float(np.mean(finite)) if finite else float("nan") + result.append(output) + return result + + +def paired_probe_reduction( + rows: list[dict[str, Any]], + stage: str, + reference_probe: str, + candidate_probe: str = "fusion_same", +) -> dict[str, Any]: + lookup = { + (int(row["held_out_run"]), int(row["step"]), row["probe"]): row + for row in rows + if row["stage"] == stage + and row["probe"] in {reference_probe, candidate_probe} + } + pairs = sorted( + { + (held_out, step) + for held_out, step, probe in lookup + if probe == candidate_probe + and (held_out, step, reference_probe) in lookup + } + ) + reductions = [] + for held_out, step in pairs: + reference = float(lookup[(held_out, step, reference_probe)]["mse"]) + candidate = float(lookup[(held_out, step, candidate_probe)]["mse"]) + reductions.append((reference - candidate) / max(reference, 1e-12)) + return { + "reference": reference_probe, + "paired_count": len(reductions), + "mean_mse_reduction": ( + float(np.mean(reductions)) if reductions else float("nan") + ), + "median_mse_reduction": ( + float(np.median(reductions)) if reductions else float("nan") + ), + "wins": int(sum(value > 0 for value in reductions)), + } + + +def plot_similarity(pair_rows: list[dict[str, Any]], output_dir: Path) -> None: + stages = [ + stage + for stage in sorted({row["stage"] for row in pair_rows}) + if stage.endswith("_hidden") + ] + comparisons = [ + "within_adjacent", + "cross_same", + "cross_wrong_step", + "cross_distant", + "cross_video_shuffle", + ] + steps = sorted({int(row["target_step"]) for row in pair_rows}) + fig, axes = plt.subplots( + len(stages), 2, figsize=(12, max(3.2, 2.8 * len(stages))), squeeze=False + ) + for stage_index, stage in enumerate(stages): + for metric_index, metric in enumerate(["cosine", "nmse"]): + matrix = np.full((len(comparisons), len(steps)), np.nan) + for row_index, comparison in enumerate(comparisons): + for col_index, step in enumerate(steps): + values = [ + float(row[metric]) + for row in pair_rows + if row["stage"] == stage + and row["comparison"] == comparison + and int(row["target_step"]) == step + ] + if values: + matrix[row_index, col_index] = np.mean(values) + ax = axes[stage_index, metric_index] + image = ax.imshow( + matrix, + aspect="auto", + cmap="viridis_r" if metric == "nmse" else "viridis", + ) + ax.set_title(f"{stage}: {metric}") + ax.set_xticks(range(len(steps)), labels=steps) + ax.set_yticks(range(len(comparisons)), labels=comparisons) + ax.set_xlabel("target denoising step") + fig.colorbar(image, ax=ax, fraction=0.03) + fig.tight_layout() + fig.savefig(output_dir / "feature_redundancy_heatmap.png", dpi=180) + plt.close(fig) + + +def plot_probe_gain(probe_rows: list[dict[str, Any]], output_dir: Path) -> None: + aggregate = group_mean( + probe_rows, + ["stage", "step", "probe"], + ["mse", "nrmse", "r2", "mse_reduction_vs_within_quadratic"], + ) + stages = [ + stage + for stage in sorted({row["stage"] for row in aggregate}) + if stage.endswith("_hidden") + ] + probes = [ + "fusion_same", + "fusion_distant", + "fusion_token_shift", + "fusion_wrong_step", + "fusion_batch_shuffle", + ] + steps = sorted({int(row["step"]) for row in aggregate}) + fig, axes = plt.subplots( + len(stages), 1, figsize=(9, max(3.4, 3.0 * len(stages))), squeeze=False + ) + for stage_index, stage in enumerate(stages): + ax = axes[stage_index, 0] + for probe in probes: + values = [] + for step in steps: + match = [ + row + for row in aggregate + if row["stage"] == stage + and int(row["step"]) == step + and row["probe"] == probe + ] + values.append( + match[0]["mse_reduction_vs_within_quadratic"] + if match + else np.nan + ) + ax.plot(steps, values, marker="o", label=probe) + ax.axhline(0, color="black", linewidth=0.8) + ax.set_title(stage) + ax.set_ylabel("MSE reduction vs within quadratic") + ax.set_xlabel("target denoising step") + ax.legend(fontsize=8) + ax.grid(alpha=0.25) + fig.tight_layout() + fig.savefig(output_dir / "conditional_chunk_gain.png", dpi=180) + plt.close(fig) + + +def plot_motion(motion_rows: list[dict[str, Any]], output_dir: Path) -> None: + fig, axes = plt.subplots(1, 2, figsize=(12, 4.5)) + axes[0].scatter( + [row["total_motion"] for row in motion_rows], + [row["boundary_raw_cosine"] for row in motion_rows], + s=18, + alpha=0.65, + label="boundary raw", + ) + axes[0].scatter( + [row["total_motion"] for row in motion_rows], + [row["flow_aligned_cosine"] for row in motion_rows], + s=18, + alpha=0.65, + label="dense-flow aligned", + ) + axes[0].set_xlabel("optical-flow magnitude") + axes[0].set_ylabel("feature cosine") + axes[0].legend() + axes[0].grid(alpha=0.25) + + aggregate = group_mean( + motion_rows, + ["motion_bin"], + [ + "same_slot_cosine", + "boundary_raw_cosine", + "global_aligned_cosine", + "flow_aligned_cosine", + ], + ) + bins = ["low", "medium", "high"] + metrics = [ + "same_slot_cosine", + "boundary_raw_cosine", + "global_aligned_cosine", + "flow_aligned_cosine", + ] + width = 0.18 + x = np.arange(len(bins)) + for metric_index, metric in enumerate(metrics): + values = [] + for name in bins: + match = [row for row in aggregate if row["motion_bin"] == name] + values.append(match[0][metric] if match else np.nan) + axes[1].bar( + x + (metric_index - 1.5) * width, + values, + width=width, + label=metric.replace("_cosine", ""), + ) + axes[1].set_xticks(x, bins) + axes[1].set_ylabel("mean feature cosine") + axes[1].set_xlabel("motion bin") + axes[1].legend(fontsize=8) + axes[1].grid(axis="y", alpha=0.25) + fig.tight_layout() + fig.savefig(output_dir / "motion_alignment_analysis.png", dpi=180) + plt.close(fig) + + +def markdown_table( + rows: list[dict[str, Any]], columns: list[str], digits: int = 4 +) -> str: + lines = [ + "| " + " | ".join(columns) + " |", + "|" + "|".join(["---"] * len(columns)) + "|", + ] + for row in rows: + values = [] + for column in columns: + value = row.get(column, "") + if isinstance(value, float): + values.append(f"{value:.{digits}f}") + else: + values.append(str(value)) + lines.append("| " + " | ".join(values) + " |") + return "\n".join(lines) + + +def build_report( + runs: list[dict[str, Any]], + pair_rows: list[dict[str, Any]], + probe_rows: list[dict[str, Any]], + motion_rows: list[dict[str, Any]], + output_dir: Path, +) -> None: + final_stage = f"block_{runs[0]['projection_layer']}_hidden" + pair_summary = group_mean( + [ + row + for row in pair_rows + if row["stage"] == final_stage + ], + ["comparison"], + ["cosine", "linear_cka", "rel_l2", "nmse", "token_cosine_p10"], + ) + probe_summary = group_mean( + [ + row + for row in probe_rows + if row["stage"] == final_stage + and row["probe"] in set(PROBE_NAMES) + ], + ["probe"], + [ + "nrmse", + "r2", + "mse_reduction_vs_within_affine", + "mse_reduction_vs_within_quadratic", + ], + ) + motion_summary = group_mean( + motion_rows, + ["motion_bin"], + [ + "total_motion", + "same_slot_cosine", + "boundary_raw_cosine", + "global_aligned_cosine", + "flow_aligned_cosine", + ], + ) + control_order = [ + "within_affine", + "within_quadratic", + "fusion_distant", + "fusion_wrong_step", + "fusion_token_shift", + "fusion_batch_shuffle", + "fusion_zero", + "fusion_noise", + ] + control_summary = [ + paired_probe_reduction( + probe_rows, + stage=final_stage, + reference_probe=control, + ) + for control in control_order + ] + control_lookup = {row["reference"]: row for row in control_summary} + affine = control_lookup["within_affine"] + shuffled = control_lookup["fusion_batch_shuffle"] + conditional_statement = ( + f"在主层 `{final_stage}` 上,加入正确的前一 chunk 同 timestep 特征," + f"相对 within-only affine probe 平均降低 held-out MSE " + f"{100*affine['mean_mse_reduction']:.2f}%,并在 " + f"{affine['wins']}/{affine['paired_count']} 个 " + f"prompt–timestep 配对中取得改善。相对参数量一致的跨视频 " + f"shuffle 对照,MSE 平均降低 " + f"{100*shuffled['mean_mse_reduction']:.2f}%。" + ) + + pair_lookup = {row["comparison"]: row for row in pair_summary} + same_pair = pair_lookup.get("cross_same") + random_pair = pair_lookup.get("cross_video_shuffle") + pair_statement = "" + if same_pair is not None and random_pair is not None: + pair_statement = ( + f"正确相邻 chunk 的平均 cosine/CKA 为 " + f"{same_pair['cosine']:.4f}/{same_pair['linear_cka']:.4f}," + f"跨视频 shuffle 为 " + f"{random_pair['cosine']:.4f}/{random_pair['linear_cka']:.4f}。" + ) + + total_motion = np.asarray( + [float(row["total_motion"]) for row in motion_rows], dtype=np.float64 + ) + raw_cosine = np.asarray( + [float(row["boundary_raw_cosine"]) for row in motion_rows], + dtype=np.float64, + ) + flow_cosine = np.asarray( + [float(row["flow_aligned_cosine"]) for row in motion_rows], + dtype=np.float64, + ) + motion_correlation = ( + float(np.corrcoef(total_motion, raw_cosine)[0, 1]) + if len(motion_rows) > 1 + else float("nan") + ) + flow_gain = float(np.mean(flow_cosine - raw_cosine)) + flow_wins = int(np.sum(flow_cosine > raw_cosine)) + motion_statement = ( + f"运动强度与未对齐跨 chunk cosine 的 Pearson 相关系数为 " + f"{motion_correlation:.3f};dense-flow 对齐平均恢复 " + f"{flow_gain:.4f} cosine,并在 {flow_wins}/{len(motion_rows)} " + f"个 chunk–timestep 样本上改善。" + ) + + report = f"""# Self-Forcing Feature Cache 分析结果 + +## 实验配置 + +- Prompts:{len(runs)} +- 每条视频 latent frames:{runs[0]['num_frames']} +- 每个 AR chunk latent frames:{runs[0]['num_frame_per_block']} +- Denoising timesteps:{runs[0]['timesteps']} +- Hook layers:{runs[0]['layers']} +- 主分析 stage:`{final_stage}` + +## 核心观察 + +{conditional_statement} + +{pair_statement} + +{motion_statement} + +以下结果是 3 条 prompt 的 pilot 分析。条件 probe 按 prompt 留一测试, +统计单位是 held-out prompt 与 timestep;不能将 token 数量解释为独立视频样本数, +也不能据此声称数据集级统计显著性。 + +## 主层特征配对 + +{markdown_table(pair_summary, ['comparison', 'count', 'cosine', 'linear_cka', 'rel_l2', 'nmse', 'token_cosine_p10'])} + +## 条件 Probe + +Probe 仅使用 chunk index `c ≥ 2` 的目标 chunk,使 `correct`、`c-2 distant` +及其他控制组在完全相同的 token 上比较。 + +{markdown_table(probe_summary, ['probe', 'count', 'nrmse', 'r2', 'mse_reduction_vs_within_affine', 'mse_reduction_vs_within_quadratic'])} + +### `fusion_same` 的成对控制实验 + +正值表示正确前一 chunk 同 timestep 输入的 MSE 更低。 + +{markdown_table(control_summary, ['reference', 'paired_count', 'mean_mse_reduction', 'median_mse_reduction', 'wins'])} + +## 运动与对齐 + +{markdown_table(motion_summary, ['motion_bin', 'count', 'total_motion', 'same_slot_cosine', 'boundary_raw_cosine', 'global_aligned_cosine', 'flow_aligned_cosine'])} + +这里的 `global_aligned` 是由光流中位数估计的全局平移对齐, +`flow_aligned` 是 dense optical-flow oracle;本轮未实现 homography。 + +## 本轮结论边界 + +- 已完成:四步 DMD 模型的 hidden/residual-delta 特征采集、相似度与 + nMSE/CKA、linear/ridge 条件 probe、负对照、运动分桶和光流对齐。 +- 尚未完成:接入用户现有的小型非线性预测网络、真实 cache 替换干预、 + 最终视频质量与端到端加速评估、规模化多视频置信区间。 +- 因而当前结果支持“前一 chunk 提供额外且具有空间对应性的条件信息”, + 但还不能单独证明最终生成质量或真实加速收益。 + +## 输出文件 + +- `feature_pair_metrics.csv` +- `feature_pair_summary.csv` +- `conditional_probe_folds.csv` +- `conditional_probe_summary.csv` +- `motion_alignment_metrics.csv` +- `motion_alignment_summary.csv` +- `feature_redundancy_heatmap.png` +- `conditional_chunk_gain.png` +- `motion_alignment_analysis.png` +""" + (output_dir / "REPORT.md").write_text(report, encoding="utf-8") + + +def analyze(args: argparse.Namespace, paths: list[Path]) -> None: + print("[analysis] loading snapshots", flush=True) + runs = load_runs(paths) + pair_rows = collect_pair_rows(runs, cosine_only=args.cosine_only) + if args.cosine_only: + pair_summary = group_mean( + pair_rows, + ["comparison", "stage", "target_step"], + [ + "cosine", + "centered_cosine", + "linear_cka", + "rel_l2", + "nmse", + "token_cosine_mean", + "token_cosine_p10", + "token_cosine_p50", + "token_cosine_p90", + ], + ) + write_csv(args.output_dir / "feature_pair_metrics.csv", pair_rows) + write_csv(args.output_dir / "feature_pair_summary.csv", pair_summary) + plot_similarity(pair_rows, args.output_dir) + summary = { + "runs": len(runs), + "pair_rows": len(pair_rows), + "cosine_only": True, + } + (args.output_dir / "summary.json").write_text( + json.dumps(summary, indent=2) + "\n", encoding="utf-8" + ) + print(f"[analysis] cosine-only complete: {args.output_dir}", flush=True) + return + probe_rows = run_conditional_probes(runs, ridge=args.ridge) + motion_rows = collect_motion_rows(runs) + + pair_summary = group_mean( + pair_rows, + ["comparison", "stage", "target_step"], + [ + "cosine", + "centered_cosine", + "linear_cka", + "rel_l2", + "nmse", + "token_cosine_mean", + "token_cosine_p10", + "token_cosine_p50", + "token_cosine_p90", + ], + ) + probe_summary = group_mean( + probe_rows, + ["stage", "step", "probe"], + [ + "mse", + "nrmse", + "r2", + "cosine", + "mse_reduction_vs_within_affine", + "mse_reduction_vs_within_quadratic", + "r2_gain_vs_within_affine", + "r2_gain_vs_within_quadratic", + ], + ) + motion_summary = group_mean( + motion_rows, + ["motion_bin", "step"], + [ + "total_motion", + "camera_motion", + "object_motion", + "same_slot_cosine", + "boundary_raw_cosine", + "global_aligned_cosine", + "flow_aligned_cosine", + "valid_flow_ratio", + ], + ) + + write_csv(args.output_dir / "feature_pair_metrics.csv", pair_rows) + write_csv(args.output_dir / "feature_pair_summary.csv", pair_summary) + write_csv(args.output_dir / "conditional_probe_folds.csv", probe_rows) + write_csv(args.output_dir / "conditional_probe_summary.csv", probe_summary) + write_csv(args.output_dir / "motion_alignment_metrics.csv", motion_rows) + write_csv(args.output_dir / "motion_alignment_summary.csv", motion_summary) + plot_similarity(pair_rows, args.output_dir) + plot_probe_gain(probe_rows, args.output_dir) + plot_motion(motion_rows, args.output_dir) + build_report(runs, pair_rows, probe_rows, motion_rows, args.output_dir) + summary = { + "runs": len(runs), + "pair_rows": len(pair_rows), + "probe_rows": len(probe_rows), + "motion_rows": len(motion_rows), + "report": str(args.output_dir / "REPORT.md"), + } + (args.output_dir / "summary.json").write_text( + json.dumps(summary, indent=2) + "\n", encoding="utf-8" + ) + print(f"[analysis] complete: {args.output_dir / 'REPORT.md'}", flush=True) + + +def main() -> None: + args = parse_args() + args.config_path = resolve_path(args.config_path) + args.checkpoint_path = resolve_path(args.checkpoint_path) + args.prompt_path = resolve_path(args.prompt_path) + args.output_dir = resolve_path(args.output_dir) + args.output_dir.mkdir(parents=True, exist_ok=True) + random.seed(args.seed) + np.random.seed(args.seed) + torch.manual_seed(args.seed) + torch.set_grad_enabled(False) + + run_paths = [ + args.output_dir / "runs" / f"prompt_{index:02d}.pt" + for index in range(args.num_prompts) + ] + if not args.analysis_only: + run_paths = generate_snapshots(args) + analyze(args, run_paths) + + +if __name__ == "__main__": + main() diff --git a/scripts/analyze_fullgrid_bilinear_3models.py b/scripts/analyze_fullgrid_bilinear_3models.py new file mode 100644 index 0000000000000000000000000000000000000000..5d4681ed66b01ed4c4ab410df738d52412fbd2e7 --- /dev/null +++ b/scripts/analyze_fullgrid_bilinear_3models.py @@ -0,0 +1,556 @@ +#!/usr/bin/env python3 +"""Full-grid bilinear motion alignment for the three four-step backbones. + +This intentionally follows the original Self-Forcing pilot: compare every +current-chunk temporal slot with the previous chunk's boundary feature map and +warp the source map with continuous target-to-source flow using bilinear +``grid_sample``. Correct, global, negated, and spatially shuffled flow fields +are evaluated with identical interpolation and in-bounds masking. +""" + +from __future__ import annotations + +import argparse +import csv +import json +import math +from collections import defaultdict +from pathlib import Path +from typing import Any, Iterable + +import cv2 +import matplotlib + +matplotlib.use("Agg") +import matplotlib.pyplot as plt +import numpy as np +import torch +import torch.nn.functional as F + + +GRID_H, GRID_W = 30, 52 +PROJECTION_SEED_BASE = 20260728 + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser() + parser.add_argument("--self_root", type=Path, required=True) + parser.add_argument("--causal_root", type=Path, required=True) + parser.add_argument("--hy_root", type=Path, required=True) + parser.add_argument("--hy_right_root", type=Path, required=True) + parser.add_argument("--output_root", type=Path, required=True) + parser.add_argument("--projection_dim", type=int, default=64) + parser.add_argument("--projection_device", default="cpu") + parser.add_argument("--overwrite_projection_cache", action="store_true") + return parser.parse_args() + + +def mean(values: Iterable[float]) -> float: + values = [float(value) for value in values if np.isfinite(value)] + return float(np.mean(values)) if values else float("nan") + + +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) + path.parent.mkdir(parents=True, exist_ok=True) + with path.open("w", newline="", encoding="utf-8") as handle: + writer = csv.DictWriter(handle, fieldnames=fields) + writer.writeheader() + writer.writerows(rows) + + +def projection_matrix(dim: int, output_dim: int, device: torch.device) -> torch.Tensor: + generator = torch.Generator(device="cpu").manual_seed(PROJECTION_SEED_BASE + dim) + signs = torch.randint(0, 2, (dim, output_dim), generator=generator, dtype=torch.int8) + return signs.float().mul_(2).sub_(1).div_(math.sqrt(output_dim)).to(device) + + +def farneback(source: np.ndarray, target: np.ndarray) -> np.ndarray: + def gray(frame: np.ndarray) -> np.ndarray: + if frame.dtype != np.uint8: + frame = np.uint8(np.clip(np.round(frame * 255.0), 0, 255)) + return cv2.cvtColor(frame, cv2.COLOR_RGB2GRAY) + + return cv2.calcOpticalFlowFarneback( + gray(source), + gray(target), + None, + pyr_scale=0.5, + levels=4, + winsize=21, + iterations=5, + poly_n=7, + poly_sigma=1.5, + flags=0, + ) + + +def resize_flow(flow: np.ndarray) -> torch.Tensor: + source_h, source_w = flow.shape[:2] + resized = cv2.resize(flow, (GRID_W, GRID_H), interpolation=cv2.INTER_AREA) + resized[..., 0] *= GRID_W / source_w + resized[..., 1] *= GRID_H / source_h + return torch.from_numpy(resized).permute(2, 0, 1).float() + + +def warp(source: torch.Tensor, flow: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: + yy, xx = torch.meshgrid( + torch.arange(GRID_H, dtype=torch.float32), + torch.arange(GRID_W, dtype=torch.float32), + indexing="ij", + ) + sample_x = xx + flow[0] + sample_y = yy + flow[1] + grid = torch.stack( + [ + 2.0 * sample_x / (GRID_W - 1) - 1.0, + 2.0 * sample_y / (GRID_H - 1) - 1.0, + ], + dim=-1, + )[None] + value = source.permute(2, 0, 1)[None].float() + warped = F.grid_sample( + value, + grid, + mode="bilinear", + padding_mode="zeros", + align_corners=True, + )[0].permute(1, 2, 0) + mask = ( + (sample_x >= 0) + & (sample_x <= GRID_W - 1) + & (sample_y >= 0) + & (sample_y <= GRID_H - 1) + ) + return warped, mask + + +def cosine(target: torch.Tensor, source: torch.Tensor, mask: torch.Tensor | None = None) -> float: + values = F.cosine_similarity(target.float(), source.float(), dim=-1, eps=1e-8) + if mask is not None: + values = values[mask] + return float(values.mean()) if values.numel() else float("nan") + + +class GridRun: + def __init__( + self, + model: str, + action: str, + prompt_id: int, + anchors: np.ndarray, + chunk_size: int, + features: dict[tuple[int, int], torch.Tensor], + source: Path, + ): + self.model = model + self.action = action + self.prompt_id = int(prompt_id) + self.anchors = anchors.astype(np.uint8) + self.chunk_size = int(chunk_size) + self.features = features + self.source = source + self.chunks = max(chunk for chunk, _ in features) + 1 + self.steps = sorted({step for _, step in features}) + + +def load_self_runs(root: Path) -> list[GridRun]: + runs: list[GridRun] = [] + for path in sorted((root / "runs").glob("prompt_*.pt")): + state = torch.load(path, map_location="cpu", weights_only=False) + features = { + tuple(int(value) for value in key.split(":")): tensor.float() + for key, tensor in state["projected"].items() + } + anchors = np.load(path.with_suffix(".anchors.npz"), allow_pickle=False)["frames"] + runs.append( + GridRun( + "self_forcing", + "none", + int(state.get("run_index", len(runs))), + anchors, + int(state["num_frame_per_block"]), + features, + path, + ) + ) + del state + return runs + + +def load_causal_runs(root: Path) -> list[GridRun]: + runs: list[GridRun] = [] + for run_dir in sorted((root / "runs").glob("prompt_*")): + path = run_dir / "feature_snapshots.pt" + anchor_path = run_dir / "rgb_anchor_frames.npz" + if not path.exists() or not anchor_path.exists(): + continue + state = torch.load(path, map_location="cpu", weights_only=False) + projected = state.get("projected", {}) + if not projected: + raise ValueError(f"No full-grid projected features in {path}") + features: dict[tuple[int, int], torch.Tensor] = {} + for key, tensor in projected.items(): + layer, chunk, step = (int(value) for value in key.split(":")) + if layer == max(state["layers"]): + features[(chunk, step)] = tensor.float() + anchors = np.load(anchor_path, allow_pickle=False)["frames"] + runs.append( + GridRun( + "causal_forcing", + "none", + int(state["prompt_id"]), + anchors, + 3, + features, + path, + ) + ) + del state + return runs + + +def project_hy_run( + snapshot_path: Path, + cache_path: Path, + projection_dim: int, + device: torch.device, + overwrite: bool, +) -> dict[tuple[int, int], torch.Tensor]: + if cache_path.exists() and not overwrite: + state = torch.load(cache_path, map_location="cpu", weights_only=False) + return { + tuple(int(value) for value in key.split(":")): tensor.float() + for key, tensor in state["features"].items() + } + + data = np.load(snapshot_path, allow_pickle=False) + stages = data["stages"].astype(str) + chunks = data["chunks"].astype(int) + steps = data["steps"].astype(int) + coords = data["coords"].astype(int) + expected_coords = np.stack( + np.meshgrid(np.arange(4), np.arange(GRID_H), np.arange(GRID_W), indexing="ij"), + axis=-1, + ).reshape(-1, 3) + if coords.shape != expected_coords.shape or not np.array_equal(coords, expected_coords): + raise ValueError(f"Unexpected HY coordinate order in {snapshot_path}") + feature_array = data["features"] + projection = projection_matrix(int(feature_array.shape[-1]), projection_dim, device) + features: dict[tuple[int, int], torch.Tensor] = {} + selected = np.flatnonzero(stages == "block_53") + for position, index in enumerate(selected): + value = torch.from_numpy(np.asarray(feature_array[index])).to(device=device, dtype=torch.float32) + value = torch.matmul(value, projection).reshape(4, GRID_H, GRID_W, projection_dim) + features[(int(chunks[index]), int(steps[index]))] = value.to("cpu", torch.float16) + if position % 4 == 3: + print(f"[HY projection] {snapshot_path.parent.name}: {position + 1}/{len(selected)}", flush=True) + data.close() + cache_path.parent.mkdir(parents=True, exist_ok=True) + torch.save( + { + "source": str(snapshot_path), + "projection_dim": projection_dim, + "features": {f"{chunk}:{step}": value for (chunk, step), value in features.items()}, + }, + cache_path, + ) + return {key: value.float() for key, value in features.items()} + + +def load_hy_runs( + root: Path, + action: str, + cache_root: Path, + projection_dim: int, + device: torch.device, + overwrite: bool, +) -> list[GridRun]: + runs: list[GridRun] = [] + for case_dir in sorted((root / "runs").glob("prompt_*")): + run_dir = case_dir / action + snapshot = run_dir / "dense_selected_snapshots.npz" + anchor_path = run_dir / "rgb_anchor_frames.npz" + if not snapshot.exists() or not anchor_path.exists(): + continue + prompt_id = int(case_dir.name.split("_")[-1]) + cache_path = cache_root / action / f"prompt_{prompt_id:04d}.pt" + features = project_hy_run(snapshot, cache_path, projection_dim, device, overwrite) + anchors = np.load(anchor_path, allow_pickle=False)["frames"] + runs.append( + GridRun( + "hy_worldplay", + action, + prompt_id, + anchors, + 4, + features, + run_dir, + ) + ) + return runs + + +def shuffled_flow(flow: torch.Tensor, seed: int) -> torch.Tensor: + generator = torch.Generator(device="cpu").manual_seed(seed) + permutation = torch.randperm(GRID_H * GRID_W, generator=generator) + return flow.reshape(2, -1)[:, permutation].reshape_as(flow) + + +def collect_rows(runs: list[GridRun]) -> list[dict[str, Any]]: + rows: list[dict[str, Any]] = [] + for run_index, run in enumerate(runs): + for chunk in range(1, run.chunks): + source_frame_index = chunk * run.chunk_size - 1 + source_frame = run.anchors[source_frame_index] + boundary_flow = farneback(source_frame, run.anchors[chunk * run.chunk_size]) + median = np.median(boundary_flow.reshape(-1, 2), axis=0) + residual = boundary_flow - median[None, None] + motion = { + "total_motion": float(np.linalg.norm(boundary_flow, axis=-1).mean()), + "camera_motion": float(np.linalg.norm(median)), + "object_motion": float(np.linalg.norm(residual, axis=-1).mean()), + } + for step in run.steps: + source_map = run.features[(chunk - 1, step)][-1].float() + target_maps = run.features[(chunk, step)].float() + slot_rows: list[dict[str, float]] = [] + for slot in range(run.chunk_size): + target = target_maps[slot] + target_frame = run.anchors[chunk * run.chunk_size + slot] + flow = resize_flow(farneback(target_frame, source_frame)) + global_flow = torch.zeros_like(flow) + global_flow[0].fill_(float(torch.median(flow[0]))) + global_flow[1].fill_(float(torch.median(flow[1]))) + controls = { + "global": global_flow, + "flow": flow, + "negated": -flow, + "shuffled": shuffled_flow( + flow, + seed=(run.prompt_id + 1) * 100000 + chunk * 1000 + slot * 10 + step, + ), + } + values = {"raw_cosine": cosine(target, source_map)} + valid_ratios = [] + for name, control_flow in controls.items(): + aligned, mask = warp(source_map, control_flow) + values[f"{name}_aligned_cosine"] = cosine(target, aligned, mask) + valid_ratios.append(float(mask.float().mean())) + values["valid_flow_ratio"] = mean(valid_ratios) + slot_rows.append(values) + row = { + "model": run.model, + "action": run.action, + "prompt_id": run.prompt_id, + "chunk": chunk, + "step": step, + "source": str(run.source), + **motion, + } + for key in slot_rows[0]: + row[key] = mean(item[key] for item in slot_rows) + for name in ("global", "flow", "negated", "shuffled"): + row[f"{name}_gain"] = row[f"{name}_aligned_cosine"] - row["raw_cosine"] + row["flow_over_global"] = row["flow_aligned_cosine"] - row["global_aligned_cosine"] + row["flow_over_shuffled"] = row["flow_aligned_cosine"] - row["shuffled_aligned_cosine"] + rows.append(row) + print(f"[analysis] {run.model}/{run.action} prompt {run.prompt_id}: {run_index + 1}/{len(runs)}", flush=True) + return rows + + +def add_motion_bins(rows: list[dict[str, Any]]) -> None: + groups: dict[tuple[str, str], list[dict[str, Any]]] = defaultdict(list) + for row in rows: + groups[(row["model"], row["action"])].append(row) + for values in groups.values(): + low, high = np.quantile([row["total_motion"] for row in values], [1 / 3, 2 / 3]) + for row in values: + row["motion_bin"] = ( + "low" if row["total_motion"] <= low else "high" if row["total_motion"] > high else "medium" + ) + + +METRICS = [ + "total_motion", + "camera_motion", + "object_motion", + "raw_cosine", + "global_aligned_cosine", + "flow_aligned_cosine", + "negated_aligned_cosine", + "shuffled_aligned_cosine", + "global_gain", + "flow_gain", + "negated_gain", + "shuffled_gain", + "flow_over_global", + "flow_over_shuffled", + "valid_flow_ratio", +] + + +def summarize(rows: list[dict[str, Any]], keys: list[str]) -> list[dict[str, Any]]: + groups: dict[tuple[Any, ...], list[dict[str, Any]]] = defaultdict(list) + for row in rows: + groups[tuple(row[key] for key in keys)].append(row) + output = [] + for group, values in sorted(groups.items(), key=lambda item: tuple(map(str, item[0]))): + item = {key: value for key, value in zip(keys, group)} + item["count"] = len(values) + for metric in METRICS: + item[metric] = mean(row[metric] for row in values) + item["flow_win_fraction"] = mean(row["flow_gain"] > 0 for row in values) + item["flow_beats_shuffled_fraction"] = mean( + row["flow_aligned_cosine"] > row["shuffled_aligned_cosine"] for row in values + ) + motion = np.asarray([row["total_motion"] for row in values], dtype=np.float64) + raw = np.asarray([row["raw_cosine"] for row in values], dtype=np.float64) + item["motion_raw_pearson"] = ( + float(np.corrcoef(motion, raw)[0, 1]) if len(values) >= 3 and np.std(motion) > 0 else float("nan") + ) + output.append(item) + return output + + +def plot(rows: list[dict[str, Any]], output: Path) -> None: + groups = sorted({(row["model"], row["action"]) for row in rows}) + names = [f"{model}\n{action}" for model, action in groups] + methods = ["raw_cosine", "global_aligned_cosine", "flow_aligned_cosine", "negated_aligned_cosine", "shuffled_aligned_cosine"] + labels = ["raw", "global", "correct flow", "negated", "shuffled"] + fig, axes = plt.subplots(1, 2, figsize=(15, 5.5)) + x = np.arange(len(groups)) + width = 0.15 + for index, (metric, label) in enumerate(zip(methods, labels)): + values = [mean(row[metric] for row in rows if (row["model"], row["action"]) == group) for group in groups] + axes[0].bar(x + (index - 2) * width, values, width=width, label=label) + axes[0].set_xticks(x, names) + axes[0].set_ylabel("cosine") + axes[0].set_title("Full-grid bilinear alignment") + axes[0].legend(fontsize=8) + for group in groups: + selected = [row for row in rows if (row["model"], row["action"]) == group] + axes[1].scatter( + [row["total_motion"] for row in selected], + [row["flow_gain"] for row in selected], + s=14, + alpha=0.5, + label="/".join(group), + ) + axes[1].axhline(0, color="black", linewidth=1) + axes[1].set_xlabel("motion magnitude") + axes[1].set_ylabel("correct-flow cosine gain") + axes[1].set_title("Alignment gain vs motion") + axes[1].legend(fontsize=7) + fig.tight_layout() + fig.savefig(output / "fullgrid_bilinear_alignment.png", dpi=180) + plt.close(fig) + + +def markdown_table(rows: list[dict[str, Any]]) -> str: + columns = [ + "model", + "action", + "motion_bin", + "count", + "raw_cosine", + "global_aligned_cosine", + "flow_aligned_cosine", + "negated_aligned_cosine", + "shuffled_aligned_cosine", + "flow_gain", + "flow_over_shuffled", + "flow_win_fraction", + ] + lines = ["| " + " | ".join(columns) + " |", "|" + "|".join("---" for _ in columns) + "|"] + for row in rows: + cells = [] + for column in columns: + value = row.get(column, "") + cells.append(f"{value:.4f}" if isinstance(value, float) and np.isfinite(value) else str(value)) + lines.append("| " + " | ".join(cells) + " |") + return "\n".join(lines) + + +def main() -> None: + args = parse_args() + output = args.output_root.resolve() + output.mkdir(parents=True, exist_ok=True) + device = torch.device(args.projection_device) + rows: list[dict[str, Any]] = [] + + self_runs = load_self_runs(args.self_root.resolve()) + print(f"[load] Self runs: {len(self_runs)}", flush=True) + rows.extend(collect_rows(self_runs)) + del self_runs + + causal_runs = load_causal_runs(args.causal_root.resolve()) + print(f"[load] Causal runs: {len(causal_runs)}", flush=True) + rows.extend(collect_rows(causal_runs)) + del causal_runs + + cache_root = output / "projected_cache" / "hy_worldplay" + for action, root in ( + ("static", args.hy_root.resolve()), + ("forward", args.hy_root.resolve()), + ("right", args.hy_right_root.resolve()), + ): + runs = load_hy_runs( + root, + action, + cache_root, + args.projection_dim, + device, + args.overwrite_projection_cache, + ) + print(f"[load] HY {action} runs: {len(runs)}", flush=True) + rows.extend(collect_rows(runs)) + del runs + if device.type == "cuda": + torch.cuda.empty_cache() + + add_motion_bins(rows) + overall = summarize(rows, ["model", "action"]) + bins = summarize(rows, ["model", "action", "motion_bin"]) + chunks = summarize(rows, ["model", "action", "chunk"]) + write_csv(output / "fullgrid_bilinear_metrics.csv", rows) + write_csv(output / "fullgrid_bilinear_summary.csv", overall) + write_csv(output / "fullgrid_bilinear_motion_bins.csv", bins) + write_csv(output / "fullgrid_bilinear_chunks.csv", chunks) + plot(rows, output) + metadata = { + "projection_dim": args.projection_dim, + "projection_seed_base": PROJECTION_SEED_BASE, + "grid": [GRID_H, GRID_W], + "flow": "Farneback target-to-source, resized with vector scaling", + "warp": "bilinear grid_sample, align_corners=True, in-bounds mask", + "motion_bins": "tertiles computed independently inside each model/action", + "run_count": len({(row["model"], row["action"], row["prompt_id"]) for row in rows}), + "row_count": len(rows), + } + (output / "analysis_config.json").write_text(json.dumps(metadata, indent=2) + "\n", encoding="utf-8") + report = [ + "# Three-backbone full-grid bilinear alignment", + "", + "Primary comparison follows the original Self-Forcing pilot and keeps a complete 30x52 feature grid. Correct flow is evaluated against global, negated, and spatially shuffled flow fields under the same bilinear interpolation.", + "", + markdown_table(bins), + "", + "A correct-flow gain alone can include interpolation effects. `flow_over_shuffled` and `flow_beats_shuffled_fraction` test whether spatially correct displacement adds value beyond a magnitude-matched interpolating control.", + "", + "In the current results, Self-Forcing and Causal-Forcing retain material correct-flow advantages over shuffled flow (0.0040 and 0.0075 cosine overall). HY's static/forward/right advantages are only about 0.0002: its raw-to-warp gain is therefore dominated by bilinear smoothing rather than verified optical-flow correspondence.", + "", + "HY has only two target chunk boundaries. Motion buckets are strongly confounded with chunk position and must not be read as a clean low/medium/high causal trend; use the chunk-stratified CSV for diagnosis.", + ] + (output / "REPORT.md").write_text("\n".join(report) + "\n", encoding="utf-8") + print(f"[complete] {output}: {len(rows)} rows", flush=True) + + +if __name__ == "__main__": + main() diff --git a/scripts/analyze_motion_alignment_3models.py b/scripts/analyze_motion_alignment_3models.py new file mode 100644 index 0000000000000000000000000000000000000000..f396bd51ed1524570eb99c5337a84879ab988fbf --- /dev/null +++ b/scripts/analyze_motion_alignment_3models.py @@ -0,0 +1,859 @@ +#!/usr/bin/env python3 +"""Motion-stratified and correspondence-aware cache analysis for the three AR4 models. + +The generation jobs deliberately remain project-native. This file is an offline +consumer of their saved feature snapshots and RGB anchor frames. It keeps the +feature comparison at token locations (rather than comparing only a pooled +vector), estimates RAFT flow with a forward/backward consistency mask, and +reports raw, global-transform, homography, dense-flow, and (for WorldPlay) +camera-action rotation alignment. +""" + +from __future__ import annotations + +import argparse +import csv +import gc +import json +import math +import os +from collections import defaultdict +from pathlib import Path +from typing import Any, Iterable + +import cv2 +import matplotlib + +matplotlib.use("Agg") +import matplotlib.pyplot as plt +import numpy as np +import torch +import torch.nn.functional as F + + +IMG_W, IMG_H = 416, 240 +ROLE_ORDER = ("early", "middle", "late", "final") + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser() + parser.add_argument("--self_root", type=Path, required=True) + parser.add_argument("--causal_root", type=Path, required=True) + parser.add_argument("--hy_root", type=Path, required=True) + parser.add_argument("--hy_right_root", type=Path, default=None) + parser.add_argument("--output_root", type=Path, required=True) + parser.add_argument("--flow_device", default="auto") + parser.add_argument("--flow_backend", choices=("raft", "farneback"), default="raft") + parser.add_argument("--corr_max_tokens", type=int, default=512) + parser.add_argument("--overwrite", action="store_true") + return parser.parse_args() + + +def finite(value: Any) -> bool: + try: + return math.isfinite(float(value)) + except (TypeError, ValueError): + return False + + +def mean_or_nan(values: Iterable[float]) -> float: + values = [float(v) for v in values if finite(v)] + return float(np.mean(values)) if values else float("nan") + + +def write_csv(path: Path, rows: list[dict[str, Any]]) -> None: + if not rows: + return + path.parent.mkdir(parents=True, exist_ok=True) + 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, extrasaction="ignore") + writer.writeheader() + writer.writerows(rows) + + +def json_dump(path: Path, value: Any) -> None: + path.parent.mkdir(parents=True, exist_ok=True) + path.write_text(json.dumps(value, indent=2, ensure_ascii=False) + "\n", encoding="utf-8") + + +def _as_numpy(value: Any) -> np.ndarray: + if isinstance(value, torch.Tensor): + return value.detach().cpu().numpy() + return np.asarray(value) + + +def _key_parts(key: str) -> tuple[int, int]: + left, right = str(key).split(":")[:2] + return int(left), int(right) + + +class FeatureRun: + """A normalized view of one prompt/action's saved features.""" + + def __init__( + self, + model: str, + action: str, + prompt_id: int, + path: Path, + anchors: np.ndarray, + chunks: int, + chunk_size: int, + stages: dict[tuple[str, int, int], tuple[np.ndarray, np.ndarray]], + camera: tuple[np.ndarray, np.ndarray] | None = None, + dense_path: Path | None = None, + ): + self.model = model + self.action = action + self.prompt_id = int(prompt_id) + self.path = path + self.anchors = anchors + self.chunk_size = int(chunk_size) + self.chunks = int(chunks) + self.stages = stages + self.camera = camera + self.dense_path = dense_path + self._dense_features = None + self._dense_coords = None + self._dense_stages = None + self._dense_chunks = None + self._dense_steps = None + self.grid_h = 30 + self.grid_w = 52 + for _, (_, coords) in stages.items(): + if len(coords): + self.grid_h = max(self.grid_h, int(np.max(coords[:, 1])) + 1) + self.grid_w = max(self.grid_w, int(np.max(coords[:, 2])) + 1) + + @property + def stage_names(self) -> list[str]: + return sorted({key[0] for key in self.stages}, key=stage_sort) + + def get(self, stage: str, chunk: int, step: int): + # HY's block_53 dense snapshots are large compressed archives. Keep + # them lazy so that one run is decompressed only while it is being + # analyzed, instead of retaining ~1.2 GB per prompt/action in RAM. + if stage == "block_53" and self.dense_path is not None and self.dense_path.exists(): + if self._dense_features is None: + data = np.load(self.dense_path, allow_pickle=False) + # NpzFile lazily decompresses an array on every indexing + # operation. Materialize each member once per run so the + # 6.2-GB float16 feature tensor is not decompressed once per + # chunk/timestep query. + try: + self._dense_features = np.asarray(data["features"]) + self._dense_coords = np.asarray(data["coords"], dtype=np.int32) + self._dense_stages = np.asarray(data["stages"]).astype(str) + self._dense_chunks = np.asarray(data["chunks"], dtype=np.int32) + self._dense_steps = np.asarray(data["steps"], dtype=np.int32) + finally: + data.close() + stages = self._dense_stages + chunks = self._dense_chunks + steps = self._dense_steps + matches = np.flatnonzero( + (stages == stage) & (chunks == int(chunk)) & (steps == int(step)) + ) + if len(matches): + index = int(matches[0]) + return self._dense_features[index].astype(np.float32), self._dense_coords + return self.stages.get((stage, int(chunk), int(step))) + + def release(self): + if self._dense_features is not None: + self._dense_features = None + self._dense_coords = None + self._dense_stages = None + self._dense_chunks = None + self._dense_steps = None + gc.collect() + + +def stage_sort(stage: str): + if stage == "projected": + return (10000, stage) + try: + return (int(stage.split("_")[1]), stage) + except (IndexError, ValueError): + return (9000, stage) + + +def stage_role(run: FeatureRun, stage: str) -> str: + if stage == "projected": + return "final" + block_stages = [name for name in run.stage_names if name.startswith("block_")] + if stage in block_stages: + index = block_stages.index(stage) + return ROLE_ORDER[min(index, len(ROLE_ORDER) - 1)] + return stage + + +def _anchor(path: Path) -> np.ndarray: + data = np.load(path, allow_pickle=False) + value = data["frames"] + if value.ndim != 4: + raise ValueError(f"Expected [T,H,W,3] anchors at {path}, got {value.shape}") + return value.astype(np.uint8) + + +def load_self(root: Path) -> list[FeatureRun]: + result: list[FeatureRun] = [] + for path in sorted((root / "runs").glob("prompt_*.pt")): + state = torch.load(path, map_location="cpu", weights_only=False) + anchor_path = path.with_suffix(".anchors.npz") + if not anchor_path.exists(): + raise FileNotFoundError(anchor_path) + coords_hidden = _as_numpy(state["sample_coords"]["hidden"]).astype(np.int32) + stages: dict[tuple[str, int, int], tuple[np.ndarray, np.ndarray]] = {} + for stage, values in state["records"].items(): + if not stage.endswith("_hidden"): + continue + for key, value in values.items(): + chunk, step = _key_parts(key) + stages[(stage, chunk, step)] = (_as_numpy(value).astype(np.float32), coords_hidden) + # The native recorder has a compact full-grid random projection for its + # last block. Keep it as an additional final-stage representation. + full_coords = np.stack( + np.meshgrid(np.arange(3), np.arange(30), np.arange(52), indexing="ij"), axis=-1 + ).reshape(-1, 3).astype(np.int32) + for key, value in state.get("projected", {}).items(): + chunk, step = _key_parts(key) + array = _as_numpy(value).astype(np.float32).reshape(-1, _as_numpy(value).shape[-1]) + stages[("projected", chunk, step)] = (array, full_coords) + result.append( + FeatureRun( + "self_forcing", + "none", + int(state.get("run_index", len(result))), + path, + _anchor(anchor_path), + int(state["num_frames"]) // int(state["num_frame_per_block"]), + int(state["num_frame_per_block"]), + stages, + ) + ) + del state + if len(result) != 10: + print(f"[warn] Self-Forcing runs found {len(result)} prompts, expected 10") + return result + + +def _flat_indices_to_coords(indices: np.ndarray, frame_count: int = 3) -> np.ndarray: + indices = indices.astype(np.int64).reshape(-1) + frame_size = 30 * 52 + return np.stack( + [indices // frame_size, (indices % frame_size) // 52, indices % 52], axis=1 + ).astype(np.int32) + + +def load_causal(root: Path) -> list[FeatureRun]: + result: list[FeatureRun] = [] + for run_dir in sorted((root / "runs").glob("prompt_*")): + path = run_dir / "feature_snapshots.pt" + if not path.exists(): + continue + state = torch.load(path, map_location="cpu", weights_only=False) + anchor_path = run_dir / "rgb_anchor_frames.npz" + if not anchor_path.exists(): + raise FileNotFoundError(anchor_path) + stages: dict[tuple[str, int, int], tuple[np.ndarray, np.ndarray]] = {} + index_dict = state.get("feature_indices", {}) + for key, value in state["features"].items(): + layer, chunk, step = (int(v) for v in str(key).split(":")) + indices = index_dict.get(key) + if indices is None: + token_count = 3 * 30 * 52 + indices = np.linspace(0, token_count - 1, int(state["max_tokens"])).round().astype(np.int64) + coords = _flat_indices_to_coords(_as_numpy(indices)) + stages[(f"block_{layer:02d}", chunk, step)] = (_as_numpy(value).astype(np.float32), coords) + prompt_id = int(state.get("prompt_id", run_dir.name.split("_")[-1])) + result.append( + FeatureRun( + "causal_forcing", + "none", + prompt_id, + path, + _anchor(anchor_path), + int(state["num_chunks"]), + 3, + stages, + ) + ) + del state + if len(result) != 10: + print(f"[warn] Causal-Forcing runs found {len(result)} prompts, expected 10") + return result + + +def _load_hy_action(root: Path, action: str) -> list[FeatureRun]: + result: list[FeatureRun] = [] + for case_dir in sorted((root / "runs").glob("prompt_*")): + run_dir = case_dir / action + if not run_dir.exists(): + continue + anchor_path = run_dir / "rgb_anchor_frames.npz" + dense_path = run_dir / "dense_selected_snapshots.npz" + sampled_path = run_dir / "final_hidden_snapshots.npz" + # The sampled archive contains all requested layers and is sufficient + # for the ordinary token metrics. Keep the expensive dense archive + # only as a lazy source for HY's final block (block_53). + snapshot_path = sampled_path if sampled_path.exists() else dense_path + if not anchor_path.exists() or not snapshot_path.exists(): + print(f"[warn] skip incomplete HY run {run_dir}") + continue + data = np.load(snapshot_path, allow_pickle=False) + features = data["features"].astype(np.float32) + chunks = data["chunks"].astype(int) + steps = data["steps"].astype(int) + stages_np = data["stages"].astype(str) + coords = data["coords"].astype(np.int32) + grid_shape = tuple(int(v) for v in data["grid_shape"]) + stages: dict[tuple[str, int, int], tuple[np.ndarray, np.ndarray]] = {} + for index in range(features.shape[0]): + stages[(stages_np[index], int(chunks[index]), int(steps[index]))] = ( + features[index], + coords, + ) + camera = None + camera_path = run_dir / "camera_trajectory.npz" + if camera_path.exists(): + camera_data = np.load(camera_path, allow_pickle=False) + camera = (camera_data["viewmats"].astype(np.float64), camera_data["intrinsics"].astype(np.float64)) + metadata = run_dir / "run_metadata.json" + prompt_id = int(case_dir.name.split("_")[-1]) + chunks_count = int(json.loads(metadata.read_text()).get("video_length", 45) if metadata.exists() else 45) + latent_count = (chunks_count - 1) // 4 + 1 + # Prefer the actual snapshot chunk count when available. + chunks_count = max(1, int(np.max(chunks)) + 1) + result.append( + FeatureRun( + "hy_worldplay", + action, + prompt_id, + run_dir, + _anchor(anchor_path), + chunks_count, + 4, + stages, + camera, + # Keep the cross-model comparison at the same 240-token + # sampling budget. The dense 6,240-token archive is retained + # as an optional artifact, but using it for all-pairs matching + # would make HY's final layer incomparable and unnecessarily + # expensive (6,240^2 distances per query). + dense_path=None, + ) + ) + return result + + +class FlowEstimator: + def __init__(self, backend: str, device: str): + self.backend = backend + self.device = torch.device("cuda" if device == "auto" and torch.cuda.is_available() else (device if device != "auto" else "cpu")) + self.model = None + self.transforms = None + if backend == "raft": + try: + from torchvision.models.optical_flow import Raft_Small_Weights, raft_small + + weights = Raft_Small_Weights.DEFAULT + self.model = raft_small(weights=weights, progress=True).eval().to(self.device) + self.transforms = weights.transforms() + print(f"[flow] RAFT-small on {self.device}") + except Exception as error: + print(f"[flow] RAFT unavailable ({error}); using Farneback") + self.backend = "farneback" + + @staticmethod + def _image(frame: np.ndarray) -> np.ndarray: + if frame.shape[:2] != (IMG_H, IMG_W): + return cv2.resize(frame, (IMG_W, IMG_H), interpolation=cv2.INTER_AREA) + return frame + + def _farneback(self, first: np.ndarray, second: np.ndarray) -> np.ndarray: + a = cv2.cvtColor(self._image(first), cv2.COLOR_RGB2GRAY) + b = cv2.cvtColor(self._image(second), cv2.COLOR_RGB2GRAY) + return cv2.calcOpticalFlowFarneback(a, b, None, 0.5, 3, 21, 5, 7, 1.5, 0) + + @torch.inference_mode() + def pair(self, first: np.ndarray, second: np.ndarray) -> tuple[np.ndarray, np.ndarray, np.ndarray]: + """Return first->second flow, reverse flow, and a mask on first pixels.""" + first = self._image(first) + second = self._image(second) + if self.backend != "raft" or self.model is None: + forward = self._farneback(first, second) + backward = self._farneback(second, first) + else: + x = torch.from_numpy(first).permute(2, 0, 1).float().div_(255).unsqueeze(0).to(self.device) + y = torch.from_numpy(second).permute(2, 0, 1).float().div_(255).unsqueeze(0).to(self.device) + x, y = self.transforms(x, y) + forward = self.model(x, y)[-1][0].permute(1, 2, 0).float().cpu().numpy() + backward = self.model(y, x)[-1][0].permute(1, 2, 0).float().cpu().numpy() + # The feature query is made at pixels in ``first`` (the current/target + # frame), so the consistency mask must also be indexed in that frame. + # A previous implementation used the reverse flow and a mask indexed + # in ``second``; that silently inverted the alignment direction. + h, w = forward.shape[:2] + yy, xx = np.mgrid[0:h, 0:w].astype(np.float32) + sx = xx + forward[..., 0] + sy = yy + forward[..., 1] + sampled_backward = cv2.remap(backward, sx, sy, cv2.INTER_LINEAR, borderMode=cv2.BORDER_CONSTANT) + fb = np.linalg.norm(forward + sampled_backward, axis=-1) + magnitude = np.linalg.norm(forward, axis=-1) + in_bounds = (sx >= 0) & (sx < w) & (sy >= 0) & (sy < h) + # The threshold scales mildly with motion, avoiding an overly strict + # rejection of the fast camera actions while retaining occlusion masks. + valid = in_bounds & (fb <= 1.5 + 0.05 * magnitude) + return forward, backward, valid + + +def flow_at(flow: np.ndarray, xy: np.ndarray) -> np.ndarray: + h, w = flow.shape[:2] + x = np.clip(xy[:, 0], 0, w - 1).astype(np.float32) + y = np.clip(xy[:, 1], 0, h - 1).astype(np.float32) + return np.stack( + [cv2.remap(flow[..., dim], x, y, cv2.INTER_LINEAR).reshape(-1) for dim in range(2)], axis=1 + ) + + +def mask_at(mask: np.ndarray, xy: np.ndarray) -> np.ndarray: + h, w = mask.shape[:2] + x = np.clip(xy[:, 0], 0, w - 1).astype(np.float32) + y = np.clip(xy[:, 1], 0, h - 1).astype(np.float32) + value = cv2.remap(mask.astype(np.uint8), x, y, cv2.INTER_NEAREST).reshape(-1) + return value.astype(bool) + + +def fit_transforms(flow: np.ndarray, valid: np.ndarray): + h, w = flow.shape[:2] + yy, xx = np.mgrid[0:h:8, 0:w:8].astype(np.float32) + points = np.stack([xx.reshape(-1), yy.reshape(-1)], axis=1) + selected = valid[::8, ::8].reshape(-1) + points = points[selected] + if len(points) < 6: + fallback_y, fallback_x = np.mgrid[0:h:4, 0:w:4] + points = np.stack([fallback_x.reshape(-1), fallback_y.reshape(-1)], axis=1).astype(np.float32) + selected_flow = flow[::4, ::4].reshape(-1, 2) + else: + selected_flow = flow[::8, ::8].reshape(-1, 2)[selected] + destinations = points + selected_flow + affine = None + homography = None + if len(points) >= 3: + affine, _ = cv2.estimateAffine2D(points, destinations, method=cv2.RANSAC, ransacReprojThreshold=3.0) + if len(points) >= 4: + homography, _ = cv2.findHomography(points, destinations, cv2.RANSAC, 4.0) + median = np.median(flow[valid] if bool(np.any(valid)) else flow.reshape(-1, 2), axis=0) + translation = np.asarray([[1.0, 0.0, median[0]], [0.0, 1.0, median[1]]], dtype=np.float32) + return translation, affine, homography + + +def apply_transform(points: np.ndarray, transform: np.ndarray | None) -> np.ndarray | None: + if transform is None or len(points) == 0: + return None + if transform.shape == (2, 3): + return points @ transform[:, :2].T + transform[:, 2] + homogeneous = np.concatenate([points, np.ones((len(points), 1), dtype=np.float32)], axis=1) + mapped = homogeneous @ transform.T + return mapped[:, :2] / np.clip(mapped[:, 2:3], 1e-6, None) + + +def nearest_features(source_xy: np.ndarray, source_features: np.ndarray, query_xy: np.ndarray): + if len(source_xy) == 0 or len(query_xy) == 0: + return np.empty((0, source_features.shape[-1]), dtype=np.float32), np.zeros(len(query_xy), bool), np.empty(len(query_xy)) + distances = ((query_xy[:, None, :] - source_xy[None, :, :]) ** 2).sum(axis=-1) + indices = np.argmin(distances, axis=1) + return source_features[indices], np.ones(len(query_xy), bool), np.sqrt(distances[np.arange(len(indices)), indices]) + + +def cosine_mean(left: np.ndarray, right: np.ndarray, mask: np.ndarray | None = None) -> float: + if len(left) == 0 or len(right) == 0: + return float("nan") + value = np.sum(left * right, axis=-1) / ( + np.linalg.norm(left, axis=-1) * np.linalg.norm(right, axis=-1) + 1e-8 + ) + if mask is not None: + value = value[mask] + return mean_or_nan(value) + + +def select_frame(features: np.ndarray, coords: np.ndarray, frame: int): + selected = coords[:, 0] == int(frame) + return features[selected], coords[selected, 1:3][:, ::-1].astype(np.float32) + + +def token_correspondence(target: np.ndarray, target_xy: np.ndarray, source: np.ndarray, source_xy: np.ndarray, max_tokens: int): + if len(target) == 0 or len(source) == 0: + return {"top1_cosine": float("nan"), "top5_cosine": float("nan"), "mnn_fraction": float("nan"), "match_distance": float("nan")} + def thin(array, xy): + if len(array) <= max_tokens: + return array, xy + idx = np.linspace(0, len(array) - 1, max_tokens).round().astype(int) + return array[idx], xy[idx] + target, target_xy = thin(target, target_xy) + source, source_xy = thin(source, source_xy) + target_n = target / (np.linalg.norm(target, axis=-1, keepdims=True) + 1e-8) + source_n = source / (np.linalg.norm(source, axis=-1, keepdims=True) + 1e-8) + similarity = target_n @ source_n.T + topk = np.sort(similarity, axis=1)[:, -min(5, similarity.shape[1]):] + best_source = np.argmax(similarity, axis=1) + best_target = np.argmax(similarity, axis=0) + target_ids = np.arange(len(best_source)) + mnn = best_target[best_source] == target_ids + distance = np.linalg.norm(target_xy - source_xy[best_source], axis=-1) + return { + "top1_cosine": float(np.mean(topk[:, -1])), + "top5_cosine": float(np.mean(topk)), + "mnn_fraction": float(np.mean(mnn)), + "match_distance": float(np.mean(distance)), + } + + +def camera_rotation_homography(camera, source_index: int, target_index: int, image_shape=(IMG_H, IMG_W)): + if camera is None: + return None + viewmats, intrinsics = camera + if source_index >= len(viewmats) or target_index >= len(viewmats): + return None + source_k = intrinsics[source_index].copy() + target_k = intrinsics[target_index].copy() + h, w = image_shape + for k in (source_k, target_k): + k[0, 0] *= w + k[0, 2] *= w + k[1, 1] *= h + k[1, 2] *= h + source_r = viewmats[source_index][:3, :3] + target_r = viewmats[target_index][:3, :3] + try: + return source_k @ source_r @ target_r.T @ np.linalg.inv(target_k) + except np.linalg.LinAlgError: + return None + + +def action_motion(camera, source_index: int, target_index: int): + if camera is None: + return {"action_translation": float("nan"), "action_rotation_deg": float("nan")} + viewmats, _ = camera + if source_index >= len(viewmats) or target_index >= len(viewmats): + return {"action_translation": float("nan"), "action_rotation_deg": float("nan")} + source_to_world = np.linalg.inv(viewmats[source_index]) + target_to_world = np.linalg.inv(viewmats[target_index]) + relative_rotation = source_to_world[:3, :3].T @ target_to_world[:3, :3] + trace = np.clip((np.trace(relative_rotation) - 1.0) / 2.0, -1.0, 1.0) + return { + "action_translation": float(np.linalg.norm(target_to_world[:3, 3] - source_to_world[:3, 3])), + "action_rotation_deg": float(np.degrees(np.arccos(trace))), + } + + +def token_points(run: FeatureRun, coords: np.ndarray) -> np.ndarray: + # coords are (frame,y,x) in the latent feature grid; convert to pixel + # centers in the resized RGB canvas, represented as (x,y). + return np.stack( + [ + (coords[:, 2].astype(np.float32) + 0.5) * IMG_W / run.grid_w, + (coords[:, 1].astype(np.float32) + 0.5) * IMG_H / run.grid_h, + ], + axis=1, + ) + + +def alignment_metrics(run: FeatureRun, stage: str, target_pair, source_pair, target_slot: int, flow: np.ndarray, valid_flow: np.ndarray, transforms, action_h, corr_max_tokens: int): + target_features, target_coords = target_pair + source_features, source_coords = source_pair + target_features, target_xy = select_frame(target_features, target_coords, target_slot) + source_boundary, source_boundary_xy = select_frame(source_features, source_coords, run.chunk_size - 1) + source_same, source_same_xy = select_frame(source_features, source_coords, min(target_slot, run.chunk_size - 1)) + if len(target_features) == 0 or len(source_boundary) == 0: + return None + target_pixels = token_points(run, np.concatenate([np.full((len(target_xy), 1), target_slot), target_xy[:, ::-1]], axis=1).astype(np.int32)) + source_boundary_pixels = token_points(run, np.concatenate([np.full((len(source_boundary_xy), 1), run.chunk_size - 1), source_boundary_xy[:, ::-1]], axis=1).astype(np.int32)) + source_same_pixels = token_points(run, np.concatenate([np.full((len(source_same_xy), 1), min(target_slot, run.chunk_size - 1)), source_same_xy[:, ::-1]], axis=1).astype(np.int32)) + raw_features, _, raw_dist = nearest_features(source_boundary_pixels, source_boundary, target_pixels) + same_features, _, _ = nearest_features(source_same_pixels, source_same, target_pixels) + target_flow = flow_at(flow, target_pixels) + flow_query = target_pixels + target_flow + flow_mask = mask_at(valid_flow, target_pixels) + flow_features, _, flow_dist = nearest_features(source_boundary_pixels, source_boundary, flow_query) + translation, affine, homography = transforms + def transformed_metric(transform): + query = apply_transform(target_pixels, transform) + if query is None: + return float("nan"), float("nan") + values, _, distances = nearest_features(source_boundary_pixels, source_boundary, query) + in_bounds = (query[:, 0] >= 0) & (query[:, 0] < IMG_W) & (query[:, 1] >= 0) & (query[:, 1] < IMG_H) + return cosine_mean(target_features, values, in_bounds), mean_or_nan(distances[in_bounds]) + trans_cos, trans_dist = transformed_metric(translation) + affine_cos, affine_dist = transformed_metric(affine) + homo_cos, homo_dist = transformed_metric(homography) + action_cos, action_dist = transformed_metric(action_h) + correspondence = token_correspondence(target_features, target_pixels, source_boundary, source_boundary_pixels, corr_max_tokens) + return { + "same_slot_cosine": cosine_mean(target_features, same_features), + "boundary_raw_cosine": cosine_mean(target_features, raw_features), + "translation_aligned_cosine": trans_cos, + "affine_aligned_cosine": affine_cos, + "homography_aligned_cosine": homo_cos, + "flow_aligned_cosine": cosine_mean(target_features, flow_features, flow_mask), + "action_rotation_cosine": action_cos, + "raw_match_distance": mean_or_nan(raw_dist), + "flow_match_distance": mean_or_nan(flow_dist[flow_mask]), + "translation_match_distance": trans_dist, + "affine_match_distance": affine_dist, + "homography_match_distance": homo_dist, + "action_match_distance": action_dist, + "valid_flow_ratio": float(np.mean(flow_mask)), + "occlusion_ratio": float(1.0 - np.mean(flow_mask)), + **correspondence, + } + + +def motion_rows(runs: list[FeatureRun], estimator: FlowEstimator, corr_max_tokens: int) -> tuple[list[dict[str, Any]], list[dict[str, Any]], dict[str, Any]]: + rows: list[dict[str, Any]] = [] + correspondence_rows: list[dict[str, Any]] = [] + flow_cache: dict[tuple[str, int, int, int], tuple[np.ndarray, np.ndarray, np.ndarray]] = {} + for run in runs: + for chunk in range(1, run.chunks): + source_index = (chunk - 1) * run.chunk_size + run.chunk_size - 1 + if source_index >= len(run.anchors): + continue + for slot in range(run.chunk_size): + target_index = chunk * run.chunk_size + slot + if target_index >= len(run.anchors): + continue + cache_key = (str(run.path), chunk, source_index, target_index) + if cache_key not in flow_cache: + # The first return is target->source and is therefore the + # displacement used to query the previous chunk's tokens. + forward, _, valid = estimator.pair(run.anchors[target_index], run.anchors[source_index]) + flow_cache[cache_key] = (forward, valid, np.asarray([])) + flow, valid, _ = flow_cache[cache_key] + flow_mag = np.linalg.norm(flow, axis=-1) + median = np.median(flow[valid] if bool(np.any(valid)) else flow.reshape(-1, 2), axis=0) + residual = flow - median[None, None] + source_img = cv2.resize(run.anchors[source_index], (IMG_W, IMG_H), interpolation=cv2.INTER_AREA) + target_img = cv2.resize(run.anchors[target_index], (IMG_W, IMG_H), interpolation=cv2.INTER_AREA) + scene_cut = float(np.mean(np.abs(source_img.astype(np.float32) - target_img.astype(np.float32))) / 255.0) + action_motion_values = action_motion(run.camera, source_index, target_index) + action_h = camera_rotation_homography(run.camera, source_index, target_index) + transforms = fit_transforms(flow, valid) + stage_names = [stage for stage in run.stage_names if stage.startswith("block_") or stage == "projected"] + for stage in stage_names: + common_steps = sorted( + set(step for st, ch, step in run.stages if st == stage and ch == chunk) + & set(step for st, ch, step in run.stages if st == stage and ch == chunk - 1) + ) + for step in common_steps: + target_pair = run.get(stage, chunk, step) + source_pair = run.get(stage, chunk - 1, step) + metrics = alignment_metrics(run, stage, target_pair, source_pair, slot, flow, valid, transforms, action_h, corr_max_tokens) + if metrics is None: + continue + row = { + "model": run.model, + "action": run.action, + "prompt_id": run.prompt_id, + "chunk": chunk, + "target_slot": slot, + "step": step, + "stage": stage, + "role": stage_role(run, stage), + "source_frame": source_index, + "target_frame": target_index, + "total_motion": float(np.mean(flow_mag)), + "camera_motion": float(np.linalg.norm(median)), + "object_motion": float(np.mean(np.linalg.norm(residual, axis=-1))), + "scene_cut": scene_cut, + **action_motion_values, + **metrics, + } + rows.append(row) + if row["role"] == "final": + correspondence_rows.append({ + key: row[key] + for key in ("model", "action", "prompt_id", "chunk", "target_slot", "step", "stage", "role", "total_motion") + } | {key: metrics[key] for key in ("top1_cosine", "top5_cosine", "mnn_fraction", "match_distance")}) + run.release() + values = np.asarray([row["total_motion"] for row in rows if finite(row["total_motion"])], dtype=np.float64) + if len(values) >= 3: + low, high = np.quantile(values, [1 / 3, 2 / 3]) + for row in rows: + row["motion_bin"] = "low" if row["total_motion"] <= low else "high" if row["total_motion"] > high else "medium" + for row in correspondence_rows: + matches = [item for item in rows if all(item[k] == row[k] for k in ("model", "action", "prompt_id", "chunk", "target_slot", "step", "stage", "role"))] + row["motion_bin"] = matches[0].get("motion_bin", "unknown") if matches else "unknown" + metadata = { + "flow_backend": estimator.backend, + "flow_device": str(estimator.device), + "flow_resolution": [IMG_W, IMG_H], + "forward_backward_consistency": "valid = in-bounds and FB error <= 1.5 + 0.05*|flow|", + "motion_bin_edges": [float(low), float(high)] if len(values) >= 3 else [], + } + return rows, correspondence_rows, metadata + + +def group_rows(rows: list[dict[str, Any]], keys: list[str], metrics: list[str]) -> list[dict[str, Any]]: + groups: dict[tuple[Any, ...], list[dict[str, Any]]] = defaultdict(list) + for row in rows: + groups[tuple(row.get(key) for key in keys)].append(row) + output: list[dict[str, Any]] = [] + for group, values in sorted(groups.items(), key=lambda item: tuple(map(str, item[0]))): + item = {key: value for key, value in zip(keys, group)} + item["count"] = len(values) + for metric in metrics: + item[metric] = mean_or_nan([row.get(metric) for row in values]) + output.append(item) + return output + + +def plot_outputs(rows: list[dict[str, Any]], correspondence: list[dict[str, Any]], output: Path) -> None: + output.mkdir(parents=True, exist_ok=True) + if not rows: + return + models = sorted({str(row["model"]) for row in rows}) + colors = {model: plt.cm.tab10(index) for index, model in enumerate(models)} + fig, axes = plt.subplots(2, 2, figsize=(14, 10)) + for model in models: + selected = [row for row in rows if row["model"] == model] + axes[0, 0].scatter([row["total_motion"] for row in selected], [row["boundary_raw_cosine"] for row in selected], s=10, alpha=0.45, label=model, color=colors[model]) + axes[0, 0].set(xlabel="flow magnitude", ylabel="raw boundary cosine", title="Motion vs raw cross-chunk similarity") + axes[0, 0].legend(fontsize=8) + methods = ["boundary_raw_cosine", "translation_aligned_cosine", "affine_aligned_cosine", "homography_aligned_cosine", "flow_aligned_cosine", "action_rotation_cosine"] + labels = ["raw", "translation", "affine", "homography", "dense flow", "action rotation"] + positions = np.arange(len(methods)) + for index, model in enumerate(models): + values = [mean_or_nan([row.get(method) for row in rows if row["model"] == model]) for method in methods] + axes[0, 1].plot(positions, values, marker="o", label=model, color=colors[model]) + axes[0, 1].set_xticks(positions, labels, rotation=25, ha="right") + axes[0, 1].set_ylim(-1, 1) + axes[0, 1].set_title("Alignment methods") + axes[0, 1].legend(fontsize=8) + for index, model in enumerate(models): + selected = [row for row in rows if row["model"] == model] + bins = ["low", "medium", "high"] + raw = [mean_or_nan([row["boundary_raw_cosine"] for row in selected if row.get("motion_bin") == name]) for name in bins] + flow = [mean_or_nan([row["flow_aligned_cosine"] for row in selected if row.get("motion_bin") == name]) for name in bins] + axes[1, 0].plot(bins, raw, marker="o", linestyle="--", label=f"{model} raw", color=colors[model], alpha=0.55) + axes[1, 0].plot(bins, flow, marker="o", label=f"{model} flow", color=colors[model]) + axes[1, 0].set_ylim(-1, 1) + axes[1, 0].set_title("Motion bins") + axes[1, 0].set_ylabel("cosine") + axes[1, 0].legend(fontsize=7) + if correspondence: + for model in models: + selected = [row for row in correspondence if row["model"] == model] + axes[1, 1].scatter([row["total_motion"] for row in selected], [row["top1_cosine"] for row in selected], s=12, alpha=0.5, label=model, color=colors[model]) + axes[1, 1].set(xlabel="flow magnitude", ylabel="top-1 feature match cosine", title="Token correspondence") + axes[1, 1].legend(fontsize=8) + fig.tight_layout() + fig.savefig(output / "motion_alignment_overview.png", dpi=180) + plt.close(fig) + + if correspondence: + fig, axes = plt.subplots(1, 3, figsize=(15, 4.5)) + for model in models: + selected = [row for row in correspondence if row["model"] == model] + for axis, metric, title in zip(axes, ("top1_cosine", "mnn_fraction", "match_distance"), ("Top-1 cosine", "MNN fraction", "Match distance")): + axis.scatter([row["total_motion"] for row in selected], [row[metric] for row in selected], s=12, alpha=0.5, label=model, color=colors[model]) + axis.set_title(title) + axis.set_xlabel("flow magnitude") + axes[0].set_ylabel("value") + axes[0].legend(fontsize=8) + fig.tight_layout() + fig.savefig(output / "token_correspondence_summary.png", dpi=180) + plt.close(fig) + + +def markdown_table(rows: list[dict[str, Any]], columns: list[str]) -> str: + if not rows: + return "_No rows._" + lines = ["| " + " | ".join(columns) + " |", "|" + "|".join("---" for _ in columns) + "|"] + for row in rows: + values = [] + for column in columns: + value = row.get(column, "") + values.append(f"{float(value):.4f}" if isinstance(value, (float, np.floating)) and finite(value) else str(value)) + lines.append("| " + " | ".join(values) + " |") + return "\n".join(lines) + + +def write_report(output: Path, rows: list[dict[str, Any]], correspondence: list[dict[str, Any]], metadata: dict[str, Any], runs: list[FeatureRun]) -> None: + stage_summary = group_rows(rows, ["model", "action", "role", "stage"], ["boundary_raw_cosine", "translation_aligned_cosine", "affine_aligned_cosine", "homography_aligned_cosine", "flow_aligned_cosine", "action_rotation_cosine", "total_motion", "valid_flow_ratio"]) + bin_summary = group_rows(rows, ["model", "action", "motion_bin"], ["total_motion", "boundary_raw_cosine", "flow_aligned_cosine", "homography_aligned_cosine", "object_motion", "occlusion_ratio"]) + json_dump(output / "motion_alignment_summary.json", {"metadata": metadata, "stage_summary": stage_summary, "motion_bin_summary": bin_summary, "run_count": len(runs), "row_count": len(rows), "correspondence_count": len(correspondence)}) + report = [ + "# Experiment 3: motion stratification and explicit alignment", + "", + f"Runs: {len(runs)} prompt/action runs; metric rows: {len(rows)}; correspondence rows: {len(correspondence)}.", + "", + "The raw baseline compares the current chunk token with the previous chunk's boundary frame at the same spatial coordinate. `flow_aligned` uses the target-to-source flow from the backend listed below and a forward/backward consistency mask. Affine and homography are global-transform fits to the valid flow. `action_rotation` is the WorldPlay camera-rotation homography only; it does not claim to model depth-dependent translation.", + "", + "## By model/action/stage", + "", + markdown_table(stage_summary, ["model", "action", "role", "stage", "count", "boundary_raw_cosine", "translation_aligned_cosine", "affine_aligned_cosine", "homography_aligned_cosine", "flow_aligned_cosine", "action_rotation_cosine"]), + "", + "## Motion bins", + "", + markdown_table(bin_summary, ["model", "action", "motion_bin", "count", "total_motion", "boundary_raw_cosine", "homography_aligned_cosine", "flow_aligned_cosine", "object_motion", "occlusion_ratio"]), + "", + "## Interpretation guardrails", + "", + "- A positive dense-flow gain with lower raw cosine supports spatial migration rather than loss of content information.", + "- Action rotation is a physically informed control for HY; forward translation requires depth and is therefore reported separately rather than treated as an exact warp.", + "- Correlation and alignment rows are prompt-level repeated measurements; use prompt bootstrap or a mixed-effects model for significance claims.", + "", + "## Run settings", + "", + "```json", + json.dumps(metadata, indent=2, ensure_ascii=False), + "```", + ] + (output / "REPORT.md").write_text("\n".join(report) + "\n", encoding="utf-8") + write_csv(output / "motion_alignment_metrics.csv", rows) + write_csv(output / "motion_alignment_summary.csv", stage_summary) + write_csv(output / "motion_alignment_motion_bins.csv", bin_summary) + write_csv(output / "token_correspondence.csv", correspondence) + + +def main() -> None: + args = parse_args() + output = args.output_root.resolve() + output.mkdir(parents=True, exist_ok=True) + if not args.overwrite and (output / "motion_alignment_metrics.csv").exists(): + print(f"[skip] existing analysis at {output}; pass --overwrite to recompute") + return + runs: list[FeatureRun] = [] + runs.extend(load_self(args.self_root.resolve())) + runs.extend(load_causal(args.causal_root.resolve())) + runs.extend(_load_hy_action(args.hy_root.resolve(), "forward")) + runs.extend(_load_hy_action(args.hy_root.resolve(), "static")) + if args.hy_right_root is not None: + runs.extend(_load_hy_action(args.hy_right_root.resolve(), "right")) + if not runs: + raise RuntimeError("No normalized runs found") + estimator = FlowEstimator(args.flow_backend, args.flow_device) + rows, correspondence, metadata = motion_rows(runs, estimator, args.corr_max_tokens) + metadata.update({ + "models": sorted({run.model for run in runs}), + "actions": sorted({run.action for run in runs}), + "prompt_ids": sorted({run.prompt_id for run in runs}), + "run_descriptions": [ + {"model": run.model, "action": run.action, "prompt_id": run.prompt_id, "chunk_size": run.chunk_size, "chunks": run.chunks, "anchor_frames": len(run.anchors), "stages": run.stage_names} + for run in runs + ], + }) + plot_outputs(rows, correspondence, output) + write_report(output, rows, correspondence, metadata, runs) + json_dump(output / "analysis_config.json", metadata) + print(f"[complete] {output}: {len(runs)} runs, {len(rows)} motion rows, {len(correspondence)} correspondence rows") + + +if __name__ == "__main__": + main() diff --git a/scripts/build_boundary_conditional_dataset.py b/scripts/build_boundary_conditional_dataset.py new file mode 100644 index 0000000000000000000000000000000000000000..e3496c7fae88440f0010942bab4b440638189b22 --- /dev/null +++ b/scripts/build_boundary_conditional_dataset.py @@ -0,0 +1,146 @@ +#!/usr/bin/env python3 +"""Assemble the native-feature boundary-to-all conditional-probe dataset. + +Self-Forcing and HY-static already use structured temporal/spatial sampling and +are linked read-only. Causal-Forcing is normalized from a fresh recorder run +whose features are stored as [temporal, spatial, channel]. +""" + +from __future__ import annotations + +import argparse +import json +import os +from pathlib import Path + +import numpy as np +import torch + + +ROLE_MAP = {7: "early", 14: "middle", 22: "late", 29: "final"} + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser() + parser.add_argument("--self_dir", type=Path, required=True) + parser.add_argument("--causal_root", type=Path, required=True) + parser.add_argument("--hy_dir", type=Path, required=True) + parser.add_argument("--output_root", type=Path, required=True) + parser.add_argument("--num_prompts", type=int, default=10) + parser.add_argument("--chunks", type=int, default=4) + parser.add_argument("--steps", type=int, default=4) + return parser.parse_args() + + +def ensure_link(path: Path, target: Path) -> None: + target = target.resolve() + if path.is_symlink(): + if path.resolve() != target: + raise ValueError(f"Existing symlink {path} points to {path.resolve()}, not {target}") + return + if path.exists(): + raise FileExistsError(path) + path.symlink_to(target, target_is_directory=True) + + +def atomic_save(path: Path, value: dict) -> None: + temporary = path.with_suffix(path.suffix + ".tmp") + torch.save(value, temporary) + os.replace(temporary, path) + + +def coords_from_indices(indices: torch.Tensor) -> np.ndarray: + flat = indices.detach().cpu().numpy().astype(np.int64).reshape(-1) + plane = 30 * 52 + temporal = flat // plane + remainder = flat % plane + coords = np.stack([temporal, remainder // 52, remainder % 52], axis=1) + slots = sorted(int(value) for value in np.unique(temporal)) + if slots != [0, 1, 2]: + raise ValueError(f"Expected three temporal slots, got {slots}") + reference = coords[coords[:, 0] == slots[-1], 1:] + for slot in slots: + if not np.array_equal(coords[coords[:, 0] == slot, 1:], reference): + raise ValueError(f"Causal slot {slot} does not share the structured spatial grid") + return coords + + +def normalize_causal(path: Path, chunks: int, steps: int) -> dict: + state = torch.load(path, map_location="cpu", weights_only=False) + features = {} + for layer, role in ROLE_MAP.items(): + chunk_rows = [] + for chunk in range(chunks): + step_rows = [] + for step in range(steps): + value = state["features"][f"{layer}:{chunk}:{step}"] + step_rows.append(value.reshape(-1, value.shape[-1]).to(torch.float16)) + chunk_rows.append(torch.stack(step_rows, dim=0)) + features[role] = torch.stack(chunk_rows, dim=0).contiguous() + index_key = f"{next(iter(ROLE_MAP))}:0:0" + coords = coords_from_indices(state["feature_indices"][index_key]) + token_counts = {int(value.shape[2]) for value in features.values()} + if token_counts != {len(coords)}: + raise ValueError(f"Feature/coordinate mismatch: tokens={token_counts}, coords={len(coords)}") + return { + "prompt_id": int(state["prompt_id"]), + "prompt": state["prompt"], + "seed": int(state["seed"]), + "model_family": "causal_forcing", + "model_variant": "dmd4_boundary_structured", + "features": features, + "timesteps": np.asarray([1000.0, 937.5, 833.3333, 625.0], dtype=np.float32), + "coords": coords, + "grid_shape": np.asarray([3, 30, 52], dtype=np.int64), + } + + +def main() -> None: + args = parse_args() + output = args.output_root.resolve() + output.mkdir(parents=True, exist_ok=True) + ensure_link(output / "self_forcing", args.self_dir) + ensure_link(output / "hy_worldplay", args.hy_dir) + causal_output = output / "causal_forcing" + causal_output.mkdir(exist_ok=True) + inventory = [] + for prompt_id in range(args.num_prompts): + source = ( + args.causal_root + / "runs" + / f"prompt_{prompt_id:04d}" + / "feature_snapshots.pt" + ) + if not source.exists(): + raise FileNotFoundError(source) + item = normalize_causal(source, args.chunks, args.steps) + if item["prompt_id"] != prompt_id: + raise ValueError(f"Prompt mismatch in {source}: {item['prompt_id']}") + destination = causal_output / f"prompt_{prompt_id:04d}.pt" + atomic_save(destination, item) + inventory.append({ + "prompt_id": prompt_id, + "source": str(source), + "destination": str(destination), + "tokens": int(item["features"]["early"].shape[2]), + }) + manifest = { + "dataset_version": 1, + "chunk_pairing": "boundary_to_all", + "prompt_ids": list(range(args.num_prompts)), + "chunks": args.chunks, + "steps": args.steps, + "self_forcing": str(args.self_dir.resolve()), + "hy_worldplay": str(args.hy_dir.resolve()), + "causal_source": str(args.causal_root.resolve()), + "causal_inventory": inventory, + } + (output / "manifest.json").write_text( + json.dumps(manifest, indent=2, ensure_ascii=False) + "\n", + encoding="utf-8", + ) + print(f"[complete] {output} prompts={args.num_prompts}", flush=True) + + +if __name__ == "__main__": + main() diff --git a/scripts/build_conditional_probe_dataset.py b/scripts/build_conditional_probe_dataset.py new file mode 100644 index 0000000000000000000000000000000000000000..dd75868dd106bd174063dab43371ddb7fbb6b04c --- /dev/null +++ b/scripts/build_conditional_probe_dataset.py @@ -0,0 +1,247 @@ +#!/usr/bin/env python3 +"""Normalize three-model hidden snapshots into one offline probe dataset. + +The recorder implementations live in their respective projects. This script +only converts their per-prompt snapshots to a small common CPU format; it does +not run a model or manufacture control examples. +""" + +from __future__ import annotations + +import argparse +import json +import os +import sys +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 numpy as np +import torch + + +ROLES = { + "self_forcing": {7: "early", 14: "middle", 22: "late", 29: "final"}, + "causal_forcing": {7: "early", 14: "middle", 22: "late", 29: "final"}, + "hy_worldplay": {13: "early", 26: "middle", 40: "late", 53: "final"}, +} + + +def regular_coords(frames: int = 3, height: int = 30, width: int = 52, max_tokens: int = 240): + total = frames * height * width + if total <= max_tokens: + flat = np.arange(total, dtype=np.int64) + else: + per_frame = max(1, max_tokens // frames) + h_count = min(height, max(1, int(round((per_frame * height / width) ** 0.5)))) + w_count = min(width, max(1, per_frame // h_count)) + while frames * h_count * w_count > max_tokens and w_count > 1: + w_count -= 1 + while frames * h_count * w_count > max_tokens and h_count > 1: + h_count -= 1 + hs = np.unique(np.rint(np.linspace(0, height - 1, h_count)).astype(np.int64)) + ws = np.unique(np.rint(np.linspace(0, width - 1, w_count)).astype(np.int64)) + flat = np.asarray( + [t * height * width + h * width + w for t in range(frames) for h in hs for w in ws], + dtype=np.int64, + ) + t = flat // (height * width) + rem = flat % (height * width) + return np.stack([t, rem // width, rem % width], axis=1) + + +def ensure_stack(values: dict[tuple[int, int], torch.Tensor], layer: int, chunks: int, steps: int): + rows = [] + for chunk in range(chunks): + step_rows = [] + for step in range(steps): + key = (chunk, step) + if key not in values: + raise ValueError(f"Missing layer={layer} chunk={chunk} step={step}") + step_rows.append(values[key].detach().cpu().to(torch.float16)) + rows.append(torch.stack(step_rows, dim=0)) + return torch.stack(rows, dim=0).contiguous() + + +def load_self(path: Path, layers: list[int], chunks: int, steps: int) -> dict[str, Any]: + run = torch.load(path, map_location="cpu", weights_only=False) + features = {} + for layer in layers: + stage = f"block_{layer}_hidden" + values = {} + for key, value in run["records"][stage].items(): + c, s = (int(part) for part in key.split(":")) + if c < chunks and s < steps: + values[(c, s)] = value + features[ROLES["self_forcing"][layer]] = ensure_stack(values, layer, chunks, steps) + return { + "prompt_id": int(run["run_index"]), + "prompt": run["prompt"], + "seed": int(run["seed"]), + "model_family": "self_forcing", + "model_variant": "dmd4", + "features": features, + "timesteps": np.asarray([1000.0, 937.5, 833.3333, 625.0], dtype=np.float32), + "coords": regular_coords(), + } + + +def load_causal(path: Path, layers: list[int], chunks: int, steps: int) -> dict[str, Any]: + run = torch.load(path, map_location="cpu", weights_only=False) + raw = {} + for key, value in run["features"].items(): + layer, chunk, step = (int(part) for part in key.split(":")) + if layer in layers and chunk < chunks and step < steps: + raw.setdefault(layer, {})[(chunk, step)] = value + features = { + ROLES["causal_forcing"][layer]: ensure_stack(raw.get(layer, {}), layer, chunks, steps) + for layer in layers + } + return { + "prompt_id": int(run["prompt_id"]), + "prompt": run["prompt"], + "seed": int(run["seed"]), + "model_family": "causal_forcing", + "model_variant": "dmd4", + "features": features, + "timesteps": np.asarray([1000.0, 937.5, 833.3333, 625.0], dtype=np.float32), + "coords": regular_coords(), + } + + +def load_hy(path: Path, layers: list[int], chunks: int, steps: int) -> dict[str, Any]: + data = np.load(path, allow_pickle=False) + stages = [str(value) for value in data["stages"]] + raw = {} + for index, stage in enumerate(stages): + if not stage.startswith("block_"): + continue + layer = int(stage.split("_")[-1]) + chunk = int(data["chunks"][index]) + step = int(data["steps"][index]) + if layer in layers and chunk < chunks and step < steps: + raw.setdefault(layer, {})[(chunk, step)] = torch.from_numpy(data["features"][index]) + features = { + ROLES["hy_worldplay"][layer]: ensure_stack(raw.get(layer, {}), layer, chunks, steps) + for layer in layers + } + return { + "features": features, + "timesteps": np.asarray(data["timesteps"], dtype=np.float32), + "coords": np.asarray(data["coords"], dtype=np.int64), + } + + +def atomic_save(path: Path, value: dict[str, Any]) -> None: + path.parent.mkdir(parents=True, exist_ok=True) + temporary = path.with_suffix(path.suffix + ".tmp") + torch.save(value, temporary) + os.replace(temporary, path) + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser() + parser.add_argument("--self_root", type=Path, required=True) + parser.add_argument("--causal_root", type=Path, required=True) + parser.add_argument("--hy_root", type=Path, required=True) + parser.add_argument("--output_root", type=Path, required=True) + parser.add_argument("--chunks", type=int, default=4) + parser.add_argument("--steps", type=int, default=4) + parser.add_argument("--max_prompts", type=int, default=10) + return parser.parse_args() + + +def main() -> None: + args = parse_args() + args.output_root.mkdir(parents=True, exist_ok=True) + specs = { + "self_forcing": ([7, 14, 22, 29], args.self_root / "runs", "self"), + "causal_forcing": ([7, 14, 22, 29], args.causal_root / "runs", "causal"), + } + inventory = [] + for family, (layers, run_root, prefix) in specs.items(): + out_dir = args.output_root / family + out_dir.mkdir(parents=True, exist_ok=True) + for prompt_id in range(args.max_prompts): + if family == "self_forcing": + source = run_root / f"prompt_{prompt_id:02d}.pt" + if not source.exists(): + raise FileNotFoundError(source) + item = load_self(source, layers, args.chunks, args.steps) + else: + source = run_root / f"prompt_{prompt_id:04d}" / "feature_snapshots.pt" + if not source.exists(): + raise FileNotFoundError(source) + item = load_causal(source, layers, args.chunks, args.steps) + destination = out_dir / f"prompt_{prompt_id:04d}.pt" + atomic_save(destination, item) + inventory.append({ + "family": family, + "prompt_id": prompt_id, + "path": str(destination), + "bytes": destination.stat().st_size, + "roles": sorted(item["features"]), + }) + + hy_files = sorted(args.hy_root.glob("shard_gpu*/runs/prompt_*/forward/final_hidden_snapshots.npz")) + hy_by_prompt = {} + for source in hy_files: + prompt_id = int(source.parts[-3].split("_")[-1]) + if prompt_id < args.max_prompts: + hy_by_prompt[prompt_id] = source + out_dir = args.output_root / "hy_worldplay" + out_dir.mkdir(parents=True, exist_ok=True) + for prompt_id in range(args.max_prompts): + source = hy_by_prompt.get(prompt_id) + if source is None: + raise FileNotFoundError(f"HY snapshot for prompt {prompt_id}") + item = load_hy(source, [13, 26, 40, 53], args.chunks, args.steps) + item.update({ + "prompt_id": prompt_id, + "model_family": "hy_worldplay", + "model_variant": "ar4", + "seed": 0, + "prompt": f"prompt_{prompt_id:04d}", + }) + destination = out_dir / f"prompt_{prompt_id:04d}.pt" + atomic_save(destination, item) + inventory.append({ + "family": "hy_worldplay", + "prompt_id": prompt_id, + "path": str(destination), + "bytes": destination.stat().st_size, + "roles": sorted(item["features"]), + }) + + manifest = { + "dataset_version": 1, + "prompt_ids": list(range(args.max_prompts)), + "chunks": args.chunks, + "steps": args.steps, + "max_tokens": 240, + "roles": ["early", "middle", "late", "final"], + "source_roots": { + "self_forcing": str(args.self_root), + "causal_forcing": str(args.causal_root), + "hy_worldplay": str(args.hy_root), + }, + "inventory": inventory, + } + (args.output_root / "manifest.json").write_text( + json.dumps(manifest, indent=2, ensure_ascii=False) + "\n", encoding="utf-8" + ) + print(f"[complete] {args.output_root} prompts={args.max_prompts} files={len(inventory)}", flush=True) + + +if __name__ == "__main__": + main() diff --git a/scripts/build_predictor_offline_data.py b/scripts/build_predictor_offline_data.py new file mode 100644 index 0000000000000000000000000000000000000000..757ebdebfab5b3434233cb239446ed830fde7226 --- /dev/null +++ b/scripts/build_predictor_offline_data.py @@ -0,0 +1,878 @@ +#!/usr/bin/env python3 +"""Build offline FFFF trajectories for the lightweight Self-Forcing predictor. + +The dataset is prompt-sharded and resumable. Common denoising tensors are +stored once per prompt, while clean self-attention prefeatures are stored in +one sidecar per Teacher block. This layout avoids duplicating history tensors +across the 18 adjacent-step training samples produced by each prompt. +""" + +from __future__ import annotations + +import argparse +import hashlib +import json +import os +import random +import shutil +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 torch +from omegaconf import OmegaConf +from safetensors.torch import save_file + +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 utils.misc import set_seed +from utils.wan_wrapper import WanDiffusionWrapper, WanTextEncoder + + +DATASET_VERSION = 2 +EXCLUDED_CHUNKS = (0,) +LATENT_CHANNELS = 16 +LATENT_HEIGHT = 60 +LATENT_WIDTH = 104 + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser( + description="Build full-step Self-Forcing predictor trajectories" + ) + 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( + "--prompt_path", + type=Path, + default=Path("prompts/vidprom_filtered_extended.txt"), + ) + parser.add_argument( + "--validation_prompt_path", + type=Path, + default=Path("prompts/MovieGenVideoBench_extended.txt"), + ) + parser.add_argument("--output_dir", type=Path, required=True) + parser.add_argument("--num_prompts", type=int, default=100) + parser.add_argument( + "--prompt_ids", + type=int, + nargs="*", + default=None, + help=( + "Only materialize these selected prompt IDs. The manifest still " + "records the full deterministic prompt selection." + ), + ) + parser.add_argument("--num_frames", type=int, default=21) + parser.add_argument("--selection_seed", type=int, default=0) + parser.add_argument("--generation_seed", type=int, default=0) + parser.add_argument( + "--layers", + type=int, + nargs="*", + default=None, + help="Teacher blocks to cache. Omit to cache every block.", + ) + parser.add_argument( + "--max_new_prompts", + type=int, + default=None, + help="Stop after this many new prompt shards; use 1 for the dry run.", + ) + parser.add_argument( + "--min_free_gib", + type=float, + default=50.0, + help="Stop before a new prompt if free disk space falls below this value.", + ) + parser.add_argument("--overwrite", action="store_true") + args = parser.parse_args() + + if args.num_prompts < 1: + parser.error("--num_prompts must be positive") + if args.prompt_ids is not None and any( + value < 0 or value >= args.num_prompts for value in args.prompt_ids + ): + parser.error("--prompt_ids must be within [0, --num_prompts)") + if args.num_frames < 1 or args.num_frames % 3: + parser.error("--num_frames must be a positive multiple of 3") + if args.max_new_prompts is not None and args.max_new_prompts < 0: + parser.error("--max_new_prompts must be non-negative") + return args + + +def resolve_path(path: Path) -> Path: + path = path.expanduser() + return path.resolve() if path.is_absolute() else (REPO_ROOT / path).resolve() + + +def atomic_write_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) + "\n", + encoding="utf-8", + ) + os.replace(temporary, path) + + +def file_sha256(path: Path) -> str: + digest = hashlib.sha256() + with path.open("rb") as handle: + while chunk := handle.read(8 * 1024 * 1024): + digest.update(chunk) + return digest.hexdigest() + + +def read_nonempty_lines(path: Path) -> list[str]: + with path.open("r", encoding="utf-8") as handle: + return [line.strip() for line in handle if line.strip()] + + +def select_prompts( + prompt_path: Path, + validation_prompt_path: Path, + num_prompts: int, + seed: int, +) -> list[dict[str, Any]]: + source = read_nonempty_lines(prompt_path) + validation = set(read_nonempty_lines(validation_prompt_path)[:100]) + eligible = [ + {"source_index": index, "prompt": prompt} + for index, prompt in enumerate(source) + if prompt not in validation + ] + if len(eligible) < num_prompts: + raise ValueError( + f"Only {len(eligible)} eligible prompts remain after excluding " + f"the first 100 validation prompts; requested {num_prompts}" + ) + return random.Random(seed).sample(eligible, num_prompts) + + +def tensor_to_bf16_cpu(value: torch.Tensor) -> torch.Tensor: + return value.detach().to(device="cpu", dtype=torch.bfloat16).contiguous() + + +def atomic_save_safetensors( + tensors: dict[str, torch.Tensor], + path: Path, + 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 directory_size(path: Path) -> int: + return sum(item.stat().st_size for item in path.rglob("*") if item.is_file()) + + +class TrajectoryRecorder: + """Capture final hidden states and clean K-projection inputs.""" + + def __init__(self, model: torch.nn.Module, layers: list[int]) -> None: + self.model = model + self.layers = layers + self.mode: str | None = None + self.final_hidden: torch.Tensor | None = None + self.current_clean: dict[int, torch.Tensor] = {} + self.clean_prefeatures: dict[int, list[torch.Tensor]] = { + layer: [] for layer in layers + } + self.handles: list[Any] = [] + + self.handles.append( + model.head.register_forward_pre_hook(self._head_pre_hook) + ) + for layer in layers: + self.handles.append( + model.blocks[layer].self_attn.k.register_forward_pre_hook( + self._make_clean_prefeature_hook(layer) + ) + ) + + def close(self) -> None: + for handle in self.handles: + handle.remove() + self.handles.clear() + + def start_denoising_step(self) -> None: + self.mode = "denoise" + self.final_hidden = None + + def finish_denoising_step(self) -> torch.Tensor: + if self.final_hidden is None: + raise RuntimeError("The Teacher head hook did not capture final_hidden") + value = self.final_hidden + self.final_hidden = None + self.mode = None + return value + + def start_clean_pass(self) -> None: + self.mode = "clean" + self.current_clean = {} + + def finish_clean_pass(self, *, store: bool = True) -> dict[int, torch.Tensor]: + missing = sorted(set(self.layers) - set(self.current_clean)) + if missing: + raise RuntimeError( + f"Clean pass did not capture prefeatures for blocks {missing}" + ) + captured = self.current_clean + if store: + for layer in self.layers: + self.clean_prefeatures[layer].append(captured[layer]) + self.current_clean = {} + self.mode = None + return captured + + def _head_pre_hook( + self, _module: torch.nn.Module, inputs: tuple[torch.Tensor, ...] + ) -> None: + if self.mode != "denoise": + return + if self.final_hidden is not None: + raise RuntimeError("Captured final_hidden more than once in one step") + if not inputs or not isinstance(inputs[0], torch.Tensor): + raise RuntimeError("Unexpected Teacher head inputs") + self.final_hidden = tensor_to_bf16_cpu(inputs[0]) + + def _make_clean_prefeature_hook(self, layer: int): + def hook( + _module: torch.nn.Module, inputs: tuple[torch.Tensor, ...] + ) -> None: + if self.mode != "clean": + return + if layer in self.current_clean: + raise RuntimeError( + f"Captured block {layer} clean prefeature more than once" + ) + if not inputs or not isinstance(inputs[0], torch.Tensor): + raise RuntimeError(f"Unexpected block {layer} K inputs") + self.current_clean[layer] = tensor_to_bf16_cpu(inputs[0]) + + return hook + + +def build_pipeline( + config: Any, checkpoint_path: Path, device: torch.device +) -> CausalInferencePipeline: + generator = WanDiffusionWrapper( + **getattr(config, "model_kwargs", {}), is_causal=True + ) + text_encoder = WanTextEncoder() + pipeline = CausalInferencePipeline( + config, + device=device, + generator=generator, + text_encoder=text_encoder, + vae=torch.nn.Identity(), + ) + + checkpoint = torch.load( + checkpoint_path, map_location="cpu", weights_only=False, mmap=True + ) + if set(checkpoint) != {"generator_ema"}: + raise KeyError( + f"Expected checkpoint key generator_ema, found {sorted(checkpoint)}" + ) + pipeline.generator.load_state_dict(checkpoint["generator_ema"], strict=True) + del checkpoint + + pipeline.to(dtype=torch.bfloat16) + pipeline.text_encoder.to(device=device) + pipeline.generator.to(device=device) + pipeline.eval() + pipeline.generator.model.requires_grad_(False) + pipeline.text_encoder.requires_grad_(False) + return pipeline + + +def reset_caches( + pipeline: CausalInferencePipeline, + batch_size: int, + dtype: torch.dtype, + device: torch.device, +) -> None: + if pipeline.kv_cache1 is None: + pipeline._initialize_kv_cache(batch_size, dtype, device) + pipeline._initialize_crossattn_cache(batch_size, dtype, device) + return + + 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 + + +def collect_cross_attention_cache( + pipeline: CausalInferencePipeline, + layers: list[int], +) -> dict[str, torch.Tensor]: + output: dict[str, torch.Tensor] = {} + for layer in layers: + cache = pipeline.crossattn_cache[layer] + if not cache["is_init"]: + raise RuntimeError(f"Cross-attention cache for block {layer} is empty") + output[f"block_{layer:02d}_k"] = tensor_to_bf16_cpu(cache["k"]) + output[f"block_{layer:02d}_v"] = tensor_to_bf16_cpu(cache["v"]) + return output + + +@torch.inference_mode() +def generate_prompt( + pipeline: CausalInferencePipeline, + recorder: TrajectoryRecorder, + prompt: str, + num_frames: int, + generation_seed: int, + device: torch.device, +) -> tuple[ + dict[str, torch.Tensor], + dict[int, list[torch.Tensor]], + dict[str, torch.Tensor], + dict[int, torch.Tensor], + dict[str, torch.Tensor], + float, + float, +]: + set_seed(generation_seed) + reset_caches(pipeline, 1, torch.bfloat16, device) + recorder.clean_prefeatures = {layer: [] for layer in recorder.layers} + + conditional_dict = pipeline.text_encoder(text_prompts=[prompt]) + noise = torch.randn( + 1, + num_frames, + LATENT_CHANNELS, + LATENT_HEIGHT, + LATENT_WIDTH, + dtype=torch.bfloat16, + device=device, + ) + trajectory: dict[str, torch.Tensor] = {} + chunk0_trajectory: dict[str, torch.Tensor] = {} + chunk0_prefeatures: dict[int, torch.Tensor] = {} + chunk_size = pipeline.num_frame_per_block + num_chunks = num_frames // chunk_size + timesteps = pipeline.denoising_step_list.to(device=device) + + torch.cuda.reset_peak_memory_stats() + torch.cuda.synchronize() + start_time = time.perf_counter() + + current_start_frame = 0 + for chunk in range(num_chunks): + noisy_input = noise[ + :, current_start_frame : current_start_frame + chunk_size + ] + timestep: torch.Tensor | None = None + denoised_pred: torch.Tensor | None = None + + for step, current_timestep in enumerate(timesteps): + timestep = ( + torch.ones( + [1, chunk_size], + device=device, + dtype=torch.int64, + ) + * current_timestep + ) + prefix = f"chunk_{chunk:02d}_step_{step:02d}" + if chunk not in EXCLUDED_CHUNKS: + trajectory[f"{prefix}_noisy_latent"] = tensor_to_bf16_cpu( + noisy_input + ) + trajectory[f"{prefix}_timestep"] = ( + timestep.detach() + .to(device="cpu", dtype=torch.float32) + .contiguous() + ) + + recorder.start_denoising_step() + 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=current_start_frame * pipeline.frame_seq_length, + ) + final_hidden = recorder.finish_denoising_step() + if chunk not in EXCLUDED_CHUNKS: + trajectory[f"{prefix}_final_hidden"] = final_hidden + trajectory[f"{prefix}_flow"] = tensor_to_bf16_cpu(flow) + else: + chunk0_trajectory[f"{prefix}_final_hidden"] = final_hidden + + if step < len(timesteps) - 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( + [chunk_size], device=device, dtype=torch.long + ), + ).unflatten(0, denoised_pred.shape[:2]) + + if denoised_pred is None or timestep is None: + raise RuntimeError("Denoising loop produced no output") + + if chunk not in EXCLUDED_CHUNKS: + trajectory[f"chunk_{chunk:02d}_clean_latent"] = ( + tensor_to_bf16_cpu(denoised_pred) + ) + + recorder.start_clean_pass() + 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=current_start_frame * pipeline.frame_seq_length, + ) + captured_clean = recorder.finish_clean_pass( + store=chunk not in EXCLUDED_CHUNKS + ) + if chunk in EXCLUDED_CHUNKS: + if chunk != 0: + raise RuntimeError(f"Unsupported excluded context chunk {chunk}") + chunk0_prefeatures = captured_clean + current_start_frame += chunk_size + + cross_attention = collect_cross_attention_cache(pipeline, recorder.layers) + torch.cuda.synchronize() + elapsed = time.perf_counter() - start_time + peak_gib = torch.cuda.max_memory_allocated() / (1024**3) + + del conditional_dict, noise + return ( + trajectory, + recorder.clean_prefeatures, + chunk0_trajectory, + chunk0_prefeatures, + cross_attention, + elapsed, + peak_gib, + ) + + +def save_prompt_shard( + output_dir: Path, + prompt_index: int, + selection: dict[str, Any], + trajectory: dict[str, torch.Tensor], + clean_prefeatures: dict[int, list[torch.Tensor]], + chunk0_trajectory: dict[str, torch.Tensor], + chunk0_prefeatures: dict[int, torch.Tensor], + cross_attention: dict[str, torch.Tensor], + elapsed_s: float, + peak_gpu_gib: float, + layers: list[int], + generation_seed: int, + num_chunks: int, +) -> Path: + destination = output_dir / f"prompt_{prompt_index:04d}" + partial = output_dir / f"prompt_{prompt_index:04d}.partial" + if partial.exists(): + shutil.rmtree(partial) + partial.mkdir(parents=True) + + shared_metadata = { + "dataset_version": str(DATASET_VERSION), + "dtype": "bfloat16", + "prompt_index": str(prompt_index), + } + atomic_save_safetensors( + trajectory, + partial / "trajectory.safetensors", + {**shared_metadata, "kind": "trajectory"}, + ) + atomic_save_safetensors( + cross_attention, + partial / "cross_attention.safetensors", + {**shared_metadata, "kind": "cross_attention_kv"}, + ) + atomic_save_safetensors( + chunk0_trajectory, + partial / "chunk0_context" / "trajectory.safetensors", + {**shared_metadata, "kind": "chunk0_context_final_hidden"}, + ) + for layer in layers: + atomic_save_safetensors( + {"chunk_00": chunk0_prefeatures[layer]}, + partial + / "chunk0_context" + / "clean_prefeatures" + / f"block_{layer:02d}.safetensors", + { + **shared_metadata, + "kind": "chunk0_context_clean_self_attention_k_input", + "block_id": str(layer), + }, + ) + atomic_write_json( + partial / "chunk0_context" / "metadata.json", + { + "kind": "context_only", + "chunk": 0, + "is_training_target": False, + "hidden_steps": [0, 1, 2, 3], + "layers": layers, + }, + ) + (partial / "chunk0_context" / "_SUCCESS").write_text( + "ok\n", encoding="utf-8" + ) + + prefeature_shapes: dict[str, list[int]] = {} + for layer in layers: + values = clean_prefeatures[layer] + stored_chunks = [ + chunk for chunk in range(num_chunks) if chunk not in EXCLUDED_CHUNKS + ] + if len(values) != len(stored_chunks): + raise RuntimeError( + f"Expected {len(stored_chunks)} stored chunks, got {len(values)}" + ) + tensors = { + f"chunk_{chunk:02d}": value + for chunk, value in zip(stored_chunks, values) + } + atomic_save_safetensors( + tensors, + partial / "clean_prefeatures" / f"block_{layer:02d}.safetensors", + { + **shared_metadata, + "kind": "clean_self_attention_k_input", + "block_id": str(layer), + }, + ) + if values: + prefeature_shapes[str(layer)] = list(values[0].shape) + + metadata = { + "dataset_version": DATASET_VERSION, + "prompt_index": prompt_index, + "source_index": selection["source_index"], + "prompt": selection["prompt"], + "generation_seed": generation_seed, + "dtype": "bfloat16", + "layers": layers, + "num_clean_chunks": len(clean_prefeatures[layers[0]]), + "excluded_chunks": list(EXCLUDED_CHUNKS), + "stored_chunks": [ + chunk for chunk in range(num_chunks) if chunk not in EXCLUDED_CHUNKS + ], + "prefeature_shapes": prefeature_shapes, + "elapsed_s": elapsed_s, + "peak_gpu_gib": peak_gpu_gib, + } + atomic_write_json(partial / "metadata.json", metadata) + (partial / "_SUCCESS").write_text("ok\n", encoding="utf-8") + os.replace(partial, destination) + return destination + + +def prepare_manifest( + args: argparse.Namespace, + config: Any, + prompt_path: Path, + validation_prompt_path: Path, + checkpoint_path: Path, + output_dir: Path, +) -> tuple[dict[str, Any], list[dict[str, Any]]]: + output_dir.mkdir(parents=True, exist_ok=True) + prompt_selection_path = output_dir / "prompt_selection.json" + selected = select_prompts( + prompt_path, + validation_prompt_path, + args.num_prompts, + args.selection_seed, + ) + selection_document = { + "selection_seed": args.selection_seed, + "num_prompts": args.num_prompts, + "prompt_source": str(prompt_path), + "prompt_source_sha256": file_sha256(prompt_path), + "validation_source": str(validation_prompt_path), + "validation_source_sha256": file_sha256(validation_prompt_path), + "excluded_validation_count": 100, + "prompts": selected, + } + + if prompt_selection_path.exists() and not args.overwrite: + existing = json.loads(prompt_selection_path.read_text(encoding="utf-8")) + if existing != selection_document: + raise RuntimeError( + "Existing prompt_selection.json differs from the requested " + "selection. Use another output directory or --overwrite." + ) + else: + atomic_write_json(prompt_selection_path, selection_document) + + manifest = { + "dataset_version": DATASET_VERSION, + "config_path": str(resolve_path(args.config_path)), + "checkpoint_path": str(checkpoint_path), + "checkpoint_key": "generator_ema", + "checkpoint_sha256": file_sha256(checkpoint_path), + "model": "Wan2.1-T2V-1.3B causal generator_ema", + "model_hidden_dim": 1536, + "num_teacher_blocks": 30, + "cached_layers": args.layers, + "storage_dtype": "bfloat16", + "num_prompts": args.num_prompts, + "num_frames": args.num_frames, + "num_chunks": args.num_frames // int(config.num_frame_per_block), + "excluded_chunks": list(EXCLUDED_CHUNKS), + "stored_chunks": [ + chunk + for chunk in range( + args.num_frames // int(config.num_frame_per_block) + ) + if chunk not in EXCLUDED_CHUNKS + ], + "num_frame_per_block": int(config.num_frame_per_block), + "local_attention_latents": ( + int(config.model_kwargs.local_attn_size) + if args.num_frames > 21 else None + ), + "denoising_step_source": list(config.denoising_step_list), + "selection_seed": args.selection_seed, + "generation_seed_reset_per_prompt": args.generation_seed, + "prompt_selection_file": str(prompt_selection_path), + "schema": { + "trajectory": "prompt_NNNN/trajectory.safetensors", + "cross_attention": "prompt_NNNN/cross_attention.safetensors", + "clean_prefeature": ( + "prompt_NNNN/clean_prefeatures/block_XX.safetensors" + ), + "chunk0_context": "prompt_NNNN/chunk0_context/", + }, + } + atomic_write_json(output_dir / "manifest.json", manifest) + return manifest, selected + + +def update_progress(output_dir: Path, num_prompts: int) -> None: + completed = [] + total_bytes = 0 + for index in range(num_prompts): + prompt_dir = output_dir / f"prompt_{index:04d}" + if (prompt_dir / "_SUCCESS").exists(): + completed.append(index) + total_bytes += directory_size(prompt_dir) + atomic_write_json( + output_dir / "progress.json", + { + "completed_prompts": completed, + "completed_count": len(completed), + "num_prompts": num_prompts, + "stored_bytes": total_bytes, + "stored_gib": total_bytes / (1024**3), + }, + ) + + +def main() -> None: + args = parse_args() + args.config_path = resolve_path(args.config_path) + args.checkpoint_path = resolve_path(args.checkpoint_path) + args.prompt_path = resolve_path(args.prompt_path) + args.validation_prompt_path = resolve_path(args.validation_prompt_path) + args.output_dir = resolve_path(args.output_dir) + + config = OmegaConf.merge( + OmegaConf.load(REPO_ROOT / "configs/default_config.yaml"), + OmegaConf.load(args.config_path), + ) + if int(config.num_frame_per_block) != 3: + raise ValueError("This dataset builder currently expects 3-frame chunks") + if args.num_frames > 21: + # Keep the released model's 21-latent training horizon as a rolling + # attention window while global RoPE positions continue increasing. + config.model_kwargs.local_attn_size = 21 + + checkpoint = torch.load( + args.checkpoint_path, map_location="cpu", weights_only=False, mmap=True + ) + state_dict = checkpoint.get("generator_ema") + if state_dict is None: + raise KeyError("Checkpoint does not contain generator_ema") + checkpoint_layers = sorted( + { + int(key.split(".")[2]) + for key in state_dict + if key.startswith("model.blocks.") + } + ) + del checkpoint, state_dict + if checkpoint_layers != list(range(30)): + raise ValueError( + f"Expected checkpoint blocks 0..29, found {checkpoint_layers}" + ) + + layers = ( + list(range(30)) + if args.layers is None or len(args.layers) == 0 + else sorted(set(args.layers)) + ) + invalid = [layer for layer in layers if layer not in checkpoint_layers] + if invalid: + raise ValueError(f"Invalid requested block IDs: {invalid}") + args.layers = layers + + _, selected = prepare_manifest( + args, + config, + args.prompt_path, + args.validation_prompt_path, + args.checkpoint_path, + args.output_dir, + ) + update_progress(args.output_dir, args.num_prompts) + if args.max_new_prompts == 0: + print("[prepare] prompt selection and manifest are ready", flush=True) + return + + requested_prompt_ids = ( + set(range(args.num_prompts)) + if args.prompt_ids is None + else set(args.prompt_ids) + ) + pending = [] + for index, selection in enumerate(selected): + if index not in requested_prompt_ids: + continue + destination = args.output_dir / f"prompt_{index:04d}" + if (destination / "_SUCCESS").exists() and not args.overwrite: + continue + pending.append((index, selection)) + if not pending: + print("[dataset] all prompt shards already exist", flush=True) + return + + device = torch.device("cuda") + pipeline = build_pipeline(config, args.checkpoint_path, device) + if len(pipeline.generator.model.blocks) != 30: + raise ValueError( + f"Loaded generator has {len(pipeline.generator.model.blocks)} blocks" + ) + recorder = TrajectoryRecorder(pipeline.generator.model, layers) + + generated = 0 + try: + for index, selection in pending: + if ( + args.max_new_prompts is not None + and generated >= args.max_new_prompts + ): + break + free_gib = shutil.disk_usage(args.output_dir).free / (1024**3) + if free_gib < args.min_free_gib: + raise RuntimeError( + f"Only {free_gib:.1f} GiB free, below --min_free_gib " + f"{args.min_free_gib:.1f}" + ) + + destination = args.output_dir / f"prompt_{index:04d}" + if destination.exists(): + if not args.overwrite: + raise RuntimeError( + f"Incomplete destination exists: {destination}" + ) + shutil.rmtree(destination) + + print( + f"[dataset] prompt {index + 1}/{args.num_prompts}, " + f"free={free_gib:.1f} GiB", + flush=True, + ) + ( + trajectory, + clean_prefeatures, + chunk0_trajectory, + chunk0_prefeatures, + cross_attention, + elapsed_s, + peak_gpu_gib, + ) = generate_prompt( + pipeline, + recorder, + selection["prompt"], + args.num_frames, + args.generation_seed, + device, + ) + destination = save_prompt_shard( + args.output_dir, + index, + selection, + trajectory, + clean_prefeatures, + chunk0_trajectory, + chunk0_prefeatures, + cross_attention, + elapsed_s, + peak_gpu_gib, + layers, + args.generation_seed, + args.num_frames // int(config.num_frame_per_block), + ) + shard_gib = directory_size(destination) / (1024**3) + print( + f"[dataset] saved {destination.name}: {shard_gib:.3f} GiB, " + f"{elapsed_s:.1f}s, peak={peak_gpu_gib:.1f} GiB", + flush=True, + ) + generated += 1 + update_progress(args.output_dir, args.num_prompts) + del ( + trajectory, + clean_prefeatures, + chunk0_trajectory, + chunk0_prefeatures, + cross_attention, + ) + torch.cuda.empty_cache() + finally: + recorder.close() + + print(f"[dataset] generated {generated} new prompt shards", flush=True) + + +if __name__ == "__main__": + main() diff --git a/scripts/build_vbench8_extended_mapping.py b/scripts/build_vbench8_extended_mapping.py new file mode 100644 index 0000000000000000000000000000000000000000..b4f148152457b103fdf5e350fd0fa1fe62893ecf --- /dev/null +++ b/scripts/build_vbench8_extended_mapping.py @@ -0,0 +1,170 @@ +#!/usr/bin/env python3 +"""Build the auditable VBench-8 extended-prompt subset mapping. + +The standard VBench metadata is the source of truth for the 946-item order and +for the prompt-suite membership. Self-Forcing's short prompt file is checked +against that order before the selected indices are transferred to the +extended prompt file. +""" + +from __future__ import annotations + +import argparse +import copy +import hashlib +import json +from collections import Counter +from pathlib import Path +from typing import Any + + +REPO_ROOT = Path(__file__).resolve().parents[1] +SELECTED_SUITES = ("subject_consistency", "overall_consistency", "scene") +EXPECTED_COUNTS = { + "subject_consistency": 72, + "overall_consistency": 93, + "scene": 86, +} + + +def read_prompts(path: Path) -> list[str]: + lines = path.read_text(encoding="utf-8").splitlines() + prompts = [line.strip() for line in lines] + if any(not prompt for prompt in prompts): + raise ValueError(f"Prompt file contains an empty line: {path}") + return prompts + + +def sha256(path: Path) -> str: + digest = hashlib.sha256() + with path.open("rb") as handle: + for block in iter(lambda: handle.read(1024 * 1024), b""): + digest.update(block) + return digest.hexdigest() + + +def build_mapping( + *, + short_prompts: list[str], + extended_prompts: list[str], + vbench_info: list[dict[str, Any]], +) -> list[dict[str, Any]]: + if len(short_prompts) != 946: + raise ValueError(f"Expected 946 short prompts, got {len(short_prompts)}") + if len(extended_prompts) != 946: + raise ValueError( + f"Expected 946 extended prompts, got {len(extended_prompts)}" + ) + if len(vbench_info) != 946: + raise ValueError(f"Expected 946 VBench metadata rows, got {len(vbench_info)}") + + canonical = [row.get("prompt_en") for row in vbench_info] + if any(not isinstance(prompt, str) or not prompt.strip() for prompt in canonical): + raise ValueError("VBench metadata contains a missing prompt_en") + mismatches = [ + index + for index, (short, official) in enumerate(zip(short_prompts, canonical)) + if short != official + ] + if mismatches: + preview = mismatches[:10] + raise ValueError( + "Self-Forcing all_dimension.txt does not preserve VBench ordering; " + f"mismatching indices include {preview}" + ) + + counters: Counter[str] = Counter() + mapping: list[dict[str, Any]] = [] + for global_index, row in enumerate(vbench_info): + suites = [suite for suite in SELECTED_SUITES if suite in row["dimension"]] + if not suites: + continue + if len(suites) != 1: + raise ValueError( + f"Metadata row {global_index} belongs to multiple selected suites: {suites}" + ) + suite = suites[0] + suite_index = counters[suite] + counters[suite] += 1 + item: dict[str, Any] = { + "global_index": global_index, + "prompt_suite": suite, + "suite_index": suite_index, + "original_prompt": short_prompts[global_index], + "extended_prompt": extended_prompts[global_index], + "official_dimensions": list(row["dimension"]), + } + if "auxiliary_info" in row: + item["auxiliary_info"] = copy.deepcopy(row["auxiliary_info"]) + mapping.append(item) + + if counters != Counter(EXPECTED_COUNTS): + raise ValueError( + f"Unexpected selected-suite counts: {dict(counters)}; " + f"expected {EXPECTED_COUNTS}" + ) + if len(mapping) != 251: + raise ValueError(f"Expected 251 selected prompts, got {len(mapping)}") + if len({item["global_index"] for item in mapping}) != len(mapping): + raise ValueError("Duplicate global indices in mapping") + for suite, expected in EXPECTED_COUNTS.items(): + indices = [item["suite_index"] for item in mapping if item["prompt_suite"] == suite] + if sorted(indices) != list(range(expected)): + raise ValueError(f"Suite indices for {suite} are not contiguous and unique") + return mapping + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument( + "--short-prompts", + type=Path, + default=REPO_ROOT / "prompts/vbench/all_dimension.txt", + ) + parser.add_argument( + "--extended-prompts", + type=Path, + default=REPO_ROOT / "prompts/vbench/all_dimension_extended.txt", + ) + parser.add_argument( + "--vbench-info", + type=Path, + default=Path( + "/data3/chenzhuo/workspace/HY-WorldPlay-light-interaction-run-DEV/" + ".venv-vbench/lib/python3.10/site-packages/vbench/VBench_full_info.json" + ), + ) + parser.add_argument( + "--output", + type=Path, + default=REPO_ROOT / "assets/vbench8_extended_subset_mapping.json", + ) + return parser.parse_args() + + +def main() -> None: + args = parse_args() + short_path = args.short_prompts.expanduser().resolve() + extended_path = args.extended_prompts.expanduser().resolve() + info_path = args.vbench_info.expanduser().resolve() + output_path = args.output.expanduser().resolve() + mapping = build_mapping( + short_prompts=read_prompts(short_path), + extended_prompts=read_prompts(extended_path), + vbench_info=json.loads(info_path.read_text(encoding="utf-8")), + ) + output_path.parent.mkdir(parents=True, exist_ok=True) + output_path.write_text( + json.dumps(mapping, ensure_ascii=False, indent=2) + "\n", encoding="utf-8" + ) + print(f"mapping={output_path}") + print(f"counts={dict(Counter(item['prompt_suite'] for item in mapping))}") + print(f"total={len(mapping)}") + print(f"short_sha256={sha256(short_path)}") + print(f"extended_sha256={sha256(extended_path)}") + print(f"vbench_info_sha256={sha256(info_path)}") + print(f"mapping_sha256={sha256(output_path)}") + + +if __name__ == "__main__": + main() diff --git a/scripts/calibrate_atc_confidence_token_beta2_vbench3.py b/scripts/calibrate_atc_confidence_token_beta2_vbench3.py new file mode 100644 index 0000000000000000000000000000000000000000..3f7fba86dfba3412990e8537a962baf01390eb87 --- /dev/null +++ b/scripts/calibrate_atc_confidence_token_beta2_vbench3.py @@ -0,0 +1,591 @@ +#!/usr/bin/env python3 +"""Calibrate fixed-beta=2 thresholds for the ATC Confidence-token head.""" + +from __future__ import annotations + +import argparse +import json +import math +import os +import sys +import tempfile +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) + os.environ.setdefault( + "MPLCONFIGDIR", tempfile.mkdtemp(prefix="atc_confidence_calibration_") + ) + return str(args.gpu) + + +PHYSICAL_GPU = preparse_gpu() + +import torch +from omegaconf import OmegaConf +from safetensors import safe_open +from safetensors.torch import load_file + +ROOT = Path(__file__).resolve().parents[1] +if str(ROOT) not in sys.path: + sys.path.insert(0, str(ROOT)) + +from predictor_training.confidence import ConfidenceTokenHead +from scripts import evaluate_single_block_fppf as base +from scripts import generate_vbench8_extended_strategies as extended +from utils.wan_wrapper import WanTextEncoder + + +BETA = 2.0 +CANDIDATE_STEPS = [1, 2, 3] +TARGETS = [6, 9, 12, 15] +GLOBAL_INDICES = [293, 729, 819] +MAPPING_DEFAULT = ROOT / "assets/vbench8_extended_subset_mapping.json" +PREDICTOR_DEFAULT = ( + ROOT + / "training_runs/layer17_atc_chunk_stage1_1000p_4gpu_b16_2000steps" + / "checkpoint_step_2000/predictor.safetensors" +) +HEAD_DEFAULT = ( + ROOT + / "training_runs" + / "layer17_atc_chunk_confidence_token_900train_100val_4gpu_b64_warmup_cosine" + / "confidence_latest.safetensors" +) +OUTPUT_DEFAULT = ( + ROOT + / "confidence_experiments" + / "layer17_atc_chunk_confidence_token_beta2_vbench3_20260902" +) + + +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, ensure_ascii=False, indent=2, allow_nan=False) + "\n", + encoding="utf-8", + ) + os.replace(temporary, path) + + +def quantile(values: list[float], fraction: float) -> float: + ordered = sorted(values) + position = fraction * (len(ordered) - 1) + lower = math.floor(position) + upper = math.ceil(position) + if lower == upper: + return ordered[lower] + weight = position - lower + return ordered[lower] * (1.0 - weight) + ordered[upper] * weight + + +def mapping_row(path: Path, global_index: int) -> dict[str, Any]: + matches = [ + row + for row in extended.read_mapping(path) + if int(row["global_index"]) == global_index + ] + if len(matches) != 1: + raise ValueError(f"Expected one mapping row for global={global_index}") + return matches[0] + + +def read_head_config(path: Path) -> dict[str, Any]: + with safe_open(path, framework="pt", device="cpu") as handle: + metadata = handle.metadata() or {} + raw = metadata.get("head_config") + if raw is None: + raise ValueError(f"Missing head_config metadata: {path}") + config = json.loads(raw) + expected = { + "architecture": "ConfidenceTokenHead", + "dim": 1536, + "token_dim": 512, + "context_dim": 64, + "num_heads": 8, + "ffn_dim": 2048, + "num_steps": 3, + } + for key, value in expected.items(): + if config.get(key) != value: + raise ValueError(f"Unexpected head config {key}={config.get(key)}") + return config + + +def read_predictor_config( + path: Path, + input_variant: str, + atc_previous_scope: str, +) -> dict[str, Any]: + with safe_open(path, framework="pt", device="cpu") as handle: + metadata = handle.metadata() or {} + raw = metadata.get("predictor_config") + if raw is None: + if input_variant != "self_forcing": + raise ValueError( + f"Missing predictor_config metadata: {path}; only an explicitly " + "selected legacy self_forcing/concat checkpoint may omit it" + ) + return { + "source_layer": 17, + "input_variant": "self_forcing", + "gate_mode": "baseline", + } + config = json.loads(raw) + actual_variant = str(config.get("input_variant", "self_forcing")) + if actual_variant != input_variant: + raise ValueError( + f"Predictor variant mismatch: expected={input_variant} " + f"actual={actual_variant}" + ) + if input_variant == "atc": + actual_scope = str(config.get("atc_previous_scope", "chunk")) + if actual_scope != atc_previous_scope: + raise ValueError( + f"ATC scope mismatch: expected={atc_previous_scope} " + f"actual={actual_scope}" + ) + return config + + +def load_models( + predictor_weights: Path, + head_weights: Path, + input_variant: str, + atc_previous_scope: str, + device: torch.device, +) -> tuple[Any, Any, Any, ConfidenceTokenHead]: + predictor_config = read_predictor_config( + predictor_weights, input_variant, atc_previous_scope + ) + config = OmegaConf.merge( + OmegaConf.load(ROOT / "configs/default_config.yaml"), + OmegaConf.load(ROOT / "configs/self_forcing_sid.yaml"), + ) + pipeline = base.build_pipeline( + config, + ROOT / "checkpoints/self_forcing_dmd.pt", + torch.nn.Identity(), + device, + ) + predictor = base.load_predictor( + pipeline.generator.model, + { + **predictor_config, + "weights": predictor_weights, + "atc_collect_diagnostics": False, + }, + device, + ) + text_encoder = WanTextEncoder().to(device=device, dtype=torch.bfloat16).eval() + text_encoder.requires_grad_(False) + head_config = read_head_config(head_weights) + head = ConfidenceTokenHead( + dim=int(head_config["dim"]), + token_dim=int(head_config["token_dim"]), + context_dim=int(head_config["context_dim"]), + num_heads=int(head_config["num_heads"]), + ffn_dim=int(head_config["ffn_dim"]), + dropout=float(head_config.get("dropout", 0.1)), + num_steps=int(head_config["num_steps"]), + ).to(device=device).eval() + head.load_state_dict(load_file(str(head_weights), device="cpu"), strict=True) + head.requires_grad_(False) + return pipeline, predictor, text_encoder, head + + +def rollout_config(name: str, threshold: float) -> dict[str, Any]: + return { + "name": name, + "policy": "dynamic", + "candidate_steps": CANDIDATE_STEPS, + "beta": BETA, + "threshold": threshold, + "target_accepts": None, + "head": "confidence_token", + "risk_mode": "outgoing_span", + "allow_chunk0_predictor": False, + } + + +def latent_nrmse(predicted: torch.Tensor, reference: torch.Tensor) -> float: + numerator = (predicted.float() - reference.float()).square().sum() + denominator = reference.float().square().sum().clamp_min(1e-8) + return float(torch.sqrt(numerator / denominator)) + + +def run_probe(args: argparse.Namespace) -> None: + if args.global_index not in GLOBAL_INDICES: + raise ValueError(f"Probe index must be one of {GLOBAL_INDICES}") + destination = args.output_root / "probes" / f"global_{args.global_index:04d}.json" + if destination.is_file() and not args.overwrite: + print(f"[cached] {destination}", flush=True) + return + row = mapping_row(args.mapping, args.global_index) + device = torch.device("cuda") + torch.set_grad_enabled(False) + pipeline, predictor, text_encoder, head = load_models( + args.predictor_weights, + args.head_weights, + args.predictor_input_variant, + args.atc_previous_scope, + device, + ) + conditional = text_encoder(text_prompts=[str(row["extended_prompt"])]) + _, diagnostic = extended.generate_rollout( + pipeline=pipeline, + conditional_dict=conditional, + seed=args.seed, + device=device, + predictor=predictor, + head=head, + config=rollout_config("confidence_token_probe", -math.inf), + ) + decisions = diagnostic.pop("decisions") + if len(decisions) != 18 or any(item["accepted"] for item in decisions): + raise RuntimeError("Full-fallback probe did not produce 18 rejected decisions") + atomic_json( + destination, + { + "status": "complete", + "physical_gpu": PHYSICAL_GPU, + "global_index": args.global_index, + "prompt_suite": row["prompt_suite"], + "suite_index": int(row["suite_index"]), + "prompt": row["extended_prompt"], + "seed": args.seed, + "generation": diagnostic, + "decisions": decisions, + }, + ) + print(f"[probe] global={args.global_index} decisions=18", flush=True) + + +def prepare_scan(args: argparse.Namespace) -> None: + paths = sorted((args.output_root / "probes").glob("global_*.json")) + if len(paths) != 3: + raise ValueError(f"Expected three probes, found {len(paths)}") + probes = [json.loads(path.read_text(encoding="utf-8")) for path in paths] + if sorted(int(item["global_index"]) for item in probes) != GLOBAL_INDICES: + raise ValueError("Probe indices differ from the requested calibration set") + decisions = [decision for probe in probes for decision in probe["decisions"]] + risks = [] + spans: dict[int, set[float]] = {1: set(), 2: set(), 3: set()} + for decision in decisions: + step = int(decision["step"]) + spans[step].add(float(decision["outgoing_span"])) + risks.append( + float(decision["predicted_local_error"]) + * float(decision["span_weight"]) + * (1.0 + BETA * float(decision["chunk_alpha"])) + ) + expected_spans = {1: 104.1666, 2: 208.3333, 3: 625.0} + for step, expected in expected_spans.items(): + if any(abs(value - expected) > 0.02 for value in spans[step]): + raise ValueError(f"Unexpected step-{step} spans: {spans[step]}") + configurations = [] + for rank in range(1, 18): + threshold = quantile(risks, rank / 18) + config = rollout_config(f"confidence_token_beta2_thr{rank:02d}", threshold) + config["threshold_rank"] = rank + configurations.append(config) + atomic_json( + args.output_root / "scan_configs.json", + { + "status": "ready", + "beta": BETA, + "risk_definition": ( + "exp(predicted_log_error) * (outgoing_span / 1000) * " + "(1 + 2.0 * chunk_alpha)" + ), + "candidate_steps": CANDIDATE_STEPS, + "targets": TARGETS, + "probe_global_indices": GLOBAL_INDICES, + "head_weights": str(args.head_weights), + "predictor_weights": str(args.predictor_weights), + "predictor_input_variant": args.predictor_input_variant, + "atc_previous_scope": ( + args.atc_previous_scope + if args.predictor_input_variant == "atc" + else None + ), + "outgoing_spans": { + str(step): sorted(values) for step, values in spans.items() + }, + "configs": configurations, + }, + ) + print(f"[prepare] thresholds={len(configurations)}", flush=True) + + +def run_scan(args: argparse.Namespace) -> None: + if args.global_index not in GLOBAL_INDICES: + raise ValueError(f"Scan index must be one of {GLOBAL_INDICES}") + destination = args.output_root / "scans" / f"global_{args.global_index:04d}.json" + existing: dict[str, Any] = {} + if destination.is_file() and not args.overwrite: + existing = json.loads(destination.read_text(encoding="utf-8")) + if existing.get("status") == "complete": + print(f"[cached] {destination}", flush=True) + return + completed = {item["name"]: item for item in existing.get("results", [])} + manifest = json.loads((args.output_root / "scan_configs.json").read_text()) + configurations = manifest["configs"] + row = mapping_row(args.mapping, args.global_index) + device = torch.device("cuda") + torch.set_grad_enabled(False) + pipeline, predictor, text_encoder, head = load_models( + args.predictor_weights, + args.head_weights, + args.predictor_input_variant, + args.atc_previous_scope, + device, + ) + conditional = text_encoder(text_prompts=[str(row["extended_prompt"])]) + reference, reference_diagnostic = extended.generate_rollout( + pipeline=pipeline, + conditional_dict=conditional, + seed=args.seed, + device=device, + predictor=predictor, + head=None, + config={ + "name": "ffff", + "policy": "ffff", + "candidate_steps": [], + "beta": None, + "threshold": None, + "target_accepts": 0, + "head": None, + "risk_mode": "outgoing_span", + "allow_chunk0_predictor": False, + }, + ) + for position, config in enumerate(configurations, start=1): + name = str(config["name"]) + if name in completed: + continue + latent, diagnostic = extended.generate_rollout( + pipeline=pipeline, + conditional_dict=conditional, + seed=args.seed, + device=device, + predictor=predictor, + head=head, + config=config, + ) + decisions = diagnostic.pop("decisions") + record = { + **config, + "accepted_predictor_calls": diagnostic["accepted_predictor_calls"], + "full_calls": diagnostic["full_calls"], + "policy_latency_ms": diagnostic["policy_latency_ms"], + "model_path_time_ms": diagnostic["model_path_time_ms"], + "confidence_head_time_ms": diagnostic["confidence_head_time_ms"], + "latent_nrmse_vs_ffff": latent_nrmse(latent, reference), + "decisions": decisions, + } + completed[name] = record + atomic_json( + destination, + { + "status": "running", + "physical_gpu": PHYSICAL_GPU, + "global_index": args.global_index, + "prompt_suite": row["prompt_suite"], + "suite_index": int(row["suite_index"]), + "prompt": row["extended_prompt"], + "seed": args.seed, + "ffff": { + key: value + for key, value in reference_diagnostic.items() + if key != "decisions" + }, + "results": [completed[key] for key in sorted(completed)], + }, + ) + print( + f"[scan] {position}/17 global={args.global_index} " + f"threshold={config['threshold']:.8f} " + f"K={record['accepted_predictor_calls']} " + f"nrmse={record['latent_nrmse_vs_ffff']:.6f}", + flush=True, + ) + del latent + atomic_json( + destination, + { + "status": "complete", + "physical_gpu": PHYSICAL_GPU, + "global_index": args.global_index, + "prompt_suite": row["prompt_suite"], + "suite_index": int(row["suite_index"]), + "prompt": row["extended_prompt"], + "seed": args.seed, + "ffff": { + key: value + for key, value in reference_diagnostic.items() + if key != "decisions" + }, + "results": [completed[key] for key in sorted(completed)], + }, + ) + + +def summarize(args: argparse.Namespace) -> None: + paths = sorted((args.output_root / "scans").glob("global_*.json")) + if len(paths) != 3: + raise ValueError(f"Expected three scans, found {len(paths)}") + scans = [json.loads(path.read_text(encoding="utf-8")) for path in paths] + if any(item.get("status") != "complete" for item in scans): + raise RuntimeError("At least one scan is incomplete") + ffff_latency = sum( + float(scan["ffff"]["policy_latency_ms"]) for scan in scans + ) / len(scans) + by_scan = [ + {record["name"]: record for record in scan["results"]} + for scan in scans + ] + names = sorted(by_scan[0]) + summary = [] + for name in names: + rows = [records[name] for records in by_scan] + mean_latency = sum(float(row["policy_latency_ms"]) for row in rows) / 3 + summary.append( + { + "name": name, + "beta": BETA, + "threshold": float(rows[0]["threshold"]), + "threshold_rank": int(rows[0]["threshold_rank"]), + "mean_accepted_predictor_calls": sum( + float(row["accepted_predictor_calls"]) for row in rows + ) + / 3, + "mean_full_calls": sum(float(row["full_calls"]) for row in rows) / 3, + "mean_policy_latency_ms": mean_latency, + "speedup_vs_ffff": ffff_latency / mean_latency, + "mean_confidence_head_time_ms": sum( + float(row["confidence_head_time_ms"]) for row in rows + ) + / 3, + "mean_latent_nrmse_vs_ffff": sum( + float(row["latent_nrmse_vs_ffff"]) for row in rows + ) + / 3, + "per_prompt_accepts": [ + int(row["accepted_predictor_calls"]) for row in rows + ], + "per_prompt_latent_nrmse": [ + float(row["latent_nrmse_vs_ffff"]) for row in rows + ], + } + ) + selected = [] + for target in TARGETS: + within = [ + row + for row in summary + if abs(float(row["mean_accepted_predictor_calls"]) - target) <= 0.5 + ] + pool = within or summary + choice = min( + pool, + key=lambda row: ( + float(row["mean_latent_nrmse_vs_ffff"]), + abs(float(row["mean_accepted_predictor_calls"]) - target), + float(row["threshold"]), + ) + if within + else ( + abs(float(row["mean_accepted_predictor_calls"]) - target), + float(row["mean_latent_nrmse_vs_ffff"]), + float(row["threshold"]), + ), + ) + selected.append({**choice, "target_accepts": target}) + atomic_json( + args.output_root / "threshold_summary.json", + { + "status": "complete", + "beta": BETA, + "candidate_steps": CANDIDATE_STEPS, + "global_indices": GLOBAL_INDICES, + "ffff_mean_policy_latency_ms": ffff_latency, + "selection_rule": ( + "For each target K, require mean K within +/-0.5 when possible; " + "then choose the lowest mean latent nRMSE" + ), + "head_weights": str(args.head_weights), + "predictor_weights": str(args.predictor_weights), + "predictor_input_variant": args.predictor_input_variant, + "atc_previous_scope": ( + args.atc_previous_scope + if args.predictor_input_variant == "atc" + else None + ), + "summary": summary, + "selected": selected, + }, + ) + print(f"[summary] FFFF latency={ffff_latency:.2f}ms", flush=True) + for row in selected: + print( + f"[selected] K{row['target_accepts']:02d} " + f"threshold={row['threshold']:.8f} " + f"actual_K={row['mean_accepted_predictor_calls']:.3f} " + f"speedup={row['speedup_vs_ffff']:.3f}x " + f"nrmse={row['mean_latent_nrmse_vs_ffff']:.6f}", + flush=True, + ) + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument( + "--mode", choices=("probe", "prepare", "scan", "summarize"), required=True + ) + parser.add_argument("--gpu", default=PHYSICAL_GPU) + parser.add_argument("--global-index", type=int, default=None) + parser.add_argument("--mapping", type=Path, default=MAPPING_DEFAULT) + parser.add_argument("--predictor-weights", type=Path, default=PREDICTOR_DEFAULT) + parser.add_argument("--head-weights", type=Path, default=HEAD_DEFAULT) + parser.add_argument( + "--predictor-input-variant", + choices=("self_forcing", "atc"), + default="atc", + ) + parser.add_argument( + "--atc-previous-scope", + choices=("chunk", "last_frame"), + default="chunk", + ) + parser.add_argument("--output-root", type=Path, default=OUTPUT_DEFAULT) + parser.add_argument("--seed", type=int, default=0) + parser.add_argument("--overwrite", action="store_true") + args = parser.parse_args() + for name in ("mapping", "predictor_weights", "head_weights", "output_root"): + setattr(args, name, getattr(args, name).expanduser().resolve()) + return args + + +def main() -> None: + args = parse_args() + args.output_root.mkdir(parents=True, exist_ok=True) + if args.mode == "probe": + run_probe(args) + elif args.mode == "prepare": + prepare_scan(args) + elif args.mode == "scan": + run_scan(args) + else: + summarize(args) + + +if __name__ == "__main__": + main() diff --git a/scripts/calibrate_vbench251_span_risk.py b/scripts/calibrate_vbench251_span_risk.py new file mode 100644 index 0000000000000000000000000000000000000000..5ffc59693efd5c4421bd9ac03814caa860b88de0 --- /dev/null +++ b/scripts/calibrate_vbench251_span_risk.py @@ -0,0 +1,705 @@ +#!/usr/bin/env python3 +"""Calibrate outgoing-span confidence risk on three Extended-251 prompts.""" + +from __future__ import annotations + +import argparse +import json +import math +import os +import sys +import tempfile +from pathlib import Path +from typing import Any + + +def preparse_gpu() -> str: + parser = argparse.ArgumentParser(add_help=False) + parser.add_argument("--gpu", default="5") + args, _ = parser.parse_known_args() + os.environ["CUDA_VISIBLE_DEVICES"] = str(args.gpu) + os.environ.setdefault( + "MPLCONFIGDIR", tempfile.mkdtemp(prefix="self_forcing_span_calibration_") + ) + return str(args.gpu) + + +PHYSICAL_GPU = preparse_gpu() + +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_single_block_fppf as base +from scripts import generate_vbench8_extended_strategies as extended +from utils.wan_wrapper import WanTextEncoder + + +MAPPING_DEFAULT = REPO_ROOT / "assets/vbench8_extended_subset_mapping.json" +EXPERIMENT_DEFAULT = ( + REPO_ROOT / "confidence_experiments/layer17_stage1_step2000_20260831" +) +OUTPUT_DEFAULT = ( + REPO_ROOT + / "confidence_experiments/layer17_stage1_step2000_spanrisk_vbench3_20260901" +) +BETAS = (0.0, 0.5, 1.0, 1.5, 2.0) +TARGETS = { + "step12": (6, 8, 10), + "step123": (6, 9, 12, 15), +} +CANDIDATE_STEPS = { + "step12": (1, 2), + "step123": (1, 2, 3), +} + + +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, ensure_ascii=False, indent=2, allow_nan=False) + "\n", + encoding="utf-8", + ) + os.replace(temporary, path) + + +def quantile(values: list[float], fraction: float) -> float: + if not values: + raise ValueError("Cannot take a quantile of an empty list") + 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 beta_name(beta: float) -> str: + return str(beta).replace(".", "p") + + +def read_mapping_row(path: Path, global_index: int) -> dict[str, Any]: + rows = extended.read_mapping(path) + matches = [row for row in rows if int(row["global_index"]) == global_index] + if len(matches) != 1: + raise ValueError(f"Expected one mapping row for global index {global_index}") + return matches[0] + + +def load_models( + experiment_root: Path, device: torch.device +) -> tuple[Any, Any, Any, Any, Any]: + config = OmegaConf.merge( + OmegaConf.load(REPO_ROOT / "configs/default_config.yaml"), + OmegaConf.load(REPO_ROOT / "configs/self_forcing_sid.yaml"), + ) + pipeline = base.build_pipeline( + config, + REPO_ROOT / "checkpoints/self_forcing_dmd.pt", + torch.nn.Identity(), + device, + ) + predictor = base.load_predictor( + pipeline.generator.model, + { + "source_layer": 17, + "weights": REPO_ROOT + / "training_runs/layer17_stage1_1000p_4gpu_b16_2000steps" + / "checkpoint_step_2000/predictor.safetensors", + "gate_mode": "baseline", + }, + device, + ) + text_encoder = WanTextEncoder().to(device=device, dtype=torch.bfloat16).eval() + text_encoder.requires_grad_(False) + step12_head = PredictorConfidenceHead(num_steps=2).to(device=device).eval() + step12_head.load_state_dict( + load_file( + str(experiment_root / "confidence_step12/confidence_best.safetensors"), + device="cpu", + ), + strict=True, + ) + step12_head.requires_grad_(False) + step123_head = PredictorConfidenceHead(num_steps=3).to(device=device).eval() + step123_head.load_state_dict( + load_file( + str(experiment_root / "confidence_step123/confidence_best.safetensors"), + device="cpu", + ), + strict=True, + ) + step123_head.requires_grad_(False) + return pipeline, predictor, text_encoder, step12_head, step123_head + + +def probe(args: argparse.Namespace) -> None: + if args.global_index is None: + raise ValueError("--global-index is required in probe mode") + row = read_mapping_row(args.mapping, args.global_index) + destination = args.output_root / "probes" / f"global_{args.global_index:04d}.json" + if destination.is_file() and not args.overwrite: + print(f"[cached] {destination}", flush=True) + return + device = torch.device("cuda") + torch.set_grad_enabled(False) + pipeline, predictor, text_encoder, step12_head, step123_head = load_models( + args.experiment_root, device + ) + prompt = str(row["extended_prompt"]) + conditional = text_encoder(text_prompts=[prompt]) + probes: dict[str, Any] = {} + for family, head in (("step12", step12_head), ("step123", step123_head)): + config = { + "name": f"{family}_span_probe", + "policy": "dynamic", + "candidate_steps": list(CANDIDATE_STEPS[family]), + "beta": 0.0, + "threshold": -math.inf, + "target_accepts": 0, + "head": family, + "risk_mode": "outgoing_span", + } + _, diagnostic = extended.generate_rollout( + pipeline=pipeline, + conditional_dict=conditional, + seed=args.seed, + device=device, + predictor=predictor, + head=head, + config=config, + ) + decisions = diagnostic.pop("decisions") + if any(decision["accepted"] for decision in decisions): + raise RuntimeError("Probe unexpectedly accepted a Predictor result") + probes[family] = {"generation": diagnostic, "decisions": decisions} + print( + f"[probe] global={args.global_index} family={family} " + f"decisions={len(decisions)}", + flush=True, + ) + atomic_json( + destination, + { + "status": "complete", + "physical_gpu": PHYSICAL_GPU, + "global_index": int(row["global_index"]), + "prompt_suite": row["prompt_suite"], + "suite_index": int(row["suite_index"]), + "prompt": prompt, + "seed": args.seed, + "probes": probes, + }, + ) + + +def prepare(args: argparse.Namespace) -> None: + paths = sorted((args.output_root / "probes").glob("global_*.json")) + if len(paths) != 3: + raise ValueError(f"Expected exactly three probe files, found {len(paths)}") + probes = [json.loads(path.read_text(encoding="utf-8")) for path in paths] + configurations: list[dict[str, Any]] = [] + span_values: dict[int, set[float]] = {1: set(), 2: set(), 3: set()} + for family in ("step12", "step123"): + decisions = [ + decision + for probe_record in probes + for decision in probe_record["probes"][family]["decisions"] + ] + expected = 3 * 6 * len(CANDIDATE_STEPS[family]) + if len(decisions) != expected: + raise ValueError( + f"Expected {expected} {family} decisions, got {len(decisions)}" + ) + for decision in decisions: + span_values[int(decision["step"])].add(float(decision["outgoing_span"])) + max_calls = 6 * len(CANDIDATE_STEPS[family]) + for beta in BETAS: + risks = [ + float(decision["predicted_local_error"]) + * float(decision["span_weight"]) + * (1.0 + beta * float(decision["chunk_alpha"])) + for decision in decisions + ] + for target in TARGETS[family]: + threshold = quantile(risks, target / max_calls) + configurations.append( + { + "name": f"{family}_span_b{beta_name(beta)}_k{target:02d}", + "policy": "dynamic", + "candidate_steps": list(CANDIDATE_STEPS[family]), + "beta": beta, + "threshold": threshold, + "target_accepts": target, + "head": family, + "risk_mode": "outgoing_span", + } + ) + actual_spans = { + str(step): sorted(values) for step, values in span_values.items() if values + } + expected_spans = {1: 104.1666, 2: 208.3333, 3: 625.0} + for step, expected in expected_spans.items(): + values = span_values[step] + if values and any(abs(value - expected) > 0.02 for value in values): + raise ValueError(f"Unexpected outgoing span at step {step}: {values}") + atomic_json( + args.output_root / "scan_configs.json", + { + "status": "ready", + "risk_definition": ( + "predicted_local_error * (outgoing_span / 1000) * " + "(1 + beta * chunk_alpha)" + ), + "probe_global_indices": [int(value["global_index"]) for value in probes], + "outgoing_spans": actual_spans, + "betas": list(BETAS), + "targets": {key: list(value) for key, value in TARGETS.items()}, + "num_configs": len(configurations), + "configs": configurations, + }, + ) + print(f"[prepare] configurations={len(configurations)}", flush=True) + + +def prepare_threshold_only(args: argparse.Namespace) -> None: + probe_root = args.probe_root or args.output_root + paths = sorted((probe_root / "probes").glob("global_*.json")) + if len(paths) != 3: + raise ValueError(f"Expected exactly three probe files, found {len(paths)}") + probes = [json.loads(path.read_text(encoding="utf-8")) for path in paths] + configurations: list[dict[str, Any]] = [] + span_values: dict[int, set[float]] = {1: set(), 2: set(), 3: set()} + beta = 2.0 + for family in ("step12", "step123"): + decisions = [ + decision + for probe_record in probes + for decision in probe_record["probes"][family]["decisions"] + ] + expected = 3 * 6 * len(CANDIDATE_STEPS[family]) + if len(decisions) != expected: + raise ValueError( + f"Expected {expected} {family} decisions, got {len(decisions)}" + ) + for decision in decisions: + span_values[int(decision["step"])].add(float(decision["outgoing_span"])) + max_calls = 6 * len(CANDIDATE_STEPS[family]) + risks = [ + float(decision["predicted_local_error"]) + * float(decision["span_weight"]) + * (1.0 + beta * float(decision["chunk_alpha"])) + for decision in decisions + ] + for rank in range(1, max_calls): + configurations.append( + { + "name": f"{family}_span_beta2_thr{rank:02d}", + "policy": "dynamic", + "candidate_steps": list(CANDIDATE_STEPS[family]), + "beta": beta, + "threshold": quantile(risks, rank / max_calls), + "target_accepts": rank, + "threshold_rank": rank, + "head": family, + "risk_mode": "outgoing_span", + } + ) + atomic_json( + args.output_root / "scan_configs.json", + { + "status": "ready", + "risk_definition": ( + "predicted_local_error * (outgoing_span / 1000) * " + "(1 + 2.0 * chunk_alpha)" + ), + "beta": beta, + "probe_global_indices": [int(value["global_index"]) for value in probes], + "outgoing_spans": { + str(step): sorted(values) + for step, values in span_values.items() + if values + }, + "threshold_search": "quantile rank 1..max_candidate_calls-1", + "targets": {key: list(value) for key, value in TARGETS.items()}, + "num_configs": len(configurations), + "configs": configurations, + }, + ) + print( + f"[prepare-threshold-only] beta={beta} configurations={len(configurations)}", + flush=True, + ) + + +def latent_nrmse(predicted: torch.Tensor, reference: torch.Tensor) -> float: + numerator = (predicted.float() - reference.float()).square().sum() + denominator = reference.float().square().sum().clamp_min(1e-8) + return float(torch.sqrt(numerator / denominator)) + + +def scan(args: argparse.Namespace) -> None: + if args.global_index is None: + raise ValueError("--global-index is required in scan mode") + row = read_mapping_row(args.mapping, args.global_index) + config_path = args.output_root / "scan_configs.json" + config_manifest = json.loads(config_path.read_text(encoding="utf-8")) + configurations = list(config_manifest["configs"]) + destination = args.output_root / "scans" / f"global_{args.global_index:04d}.json" + existing: dict[str, Any] = {} + if destination.is_file() and not args.overwrite: + existing = json.loads(destination.read_text(encoding="utf-8")) + if existing.get("status") == "complete": + print(f"[cached] {destination}", flush=True) + return + completed = {record["name"]: record for record in existing.get("results", [])} + device = torch.device("cuda") + torch.set_grad_enabled(False) + pipeline, predictor, text_encoder, step12_head, step123_head = load_models( + args.experiment_root, device + ) + prompt = str(row["extended_prompt"]) + conditional = text_encoder(text_prompts=[prompt]) + reference, reference_diagnostic = extended.generate_rollout( + pipeline=pipeline, + conditional_dict=conditional, + seed=args.seed, + device=device, + predictor=predictor, + head=None, + config={ + "name": "ffff", + "policy": "ffff", + "candidate_steps": [], + "beta": None, + "threshold": None, + "target_accepts": 0, + "head": None, + "risk_mode": "outgoing_span", + }, + ) + for position, config in enumerate(configurations, start=1): + name = str(config["name"]) + if name in completed: + continue + head = step12_head if config["head"] == "step12" else step123_head + latent, diagnostic = extended.generate_rollout( + pipeline=pipeline, + conditional_dict=conditional, + seed=args.seed, + device=device, + predictor=predictor, + head=head, + config=config, + ) + decisions = diagnostic.pop("decisions") + record = { + **config, + "accepted_predictor_calls": diagnostic["accepted_predictor_calls"], + "full_calls": diagnostic["full_calls"], + "policy_latency_ms": diagnostic["policy_latency_ms"], + "context_dit_time_ms": diagnostic["context_dit_time_ms"], + "latent_nrmse_vs_ffff": latent_nrmse(latent, reference), + "decisions": decisions, + } + completed[name] = record + atomic_json( + destination, + { + "status": "running", + "physical_gpu": PHYSICAL_GPU, + "global_index": int(row["global_index"]), + "prompt_suite": row["prompt_suite"], + "suite_index": int(row["suite_index"]), + "prompt": prompt, + "seed": args.seed, + "ffff": { + key: value + for key, value in reference_diagnostic.items() + if key != "decisions" + }, + "results": [completed[key] for key in sorted(completed)], + }, + ) + print( + f"[scan] global={args.global_index} {position}/{len(configurations)} " + f"name={name} K={record['accepted_predictor_calls']} " + f"latent_nrmse={record['latent_nrmse_vs_ffff']:.6f}", + flush=True, + ) + del latent + atomic_json( + destination, + { + "status": "complete", + "physical_gpu": PHYSICAL_GPU, + "global_index": int(row["global_index"]), + "prompt_suite": row["prompt_suite"], + "suite_index": int(row["suite_index"]), + "prompt": prompt, + "seed": args.seed, + "ffff": { + key: value + for key, value in reference_diagnostic.items() + if key != "decisions" + }, + "results": [completed[key] for key in sorted(completed)], + }, + ) + + +def summarize(args: argparse.Namespace) -> None: + paths = sorted((args.output_root / "scans").glob("global_*.json")) + if len(paths) != 3: + raise ValueError(f"Expected exactly three scan files, found {len(paths)}") + scans = [json.loads(path.read_text(encoding="utf-8")) for path in paths] + if any(scan.get("status") != "complete" for scan in scans): + raise RuntimeError("At least one scan is incomplete") + by_name = [ + {record["name"]: record for record in scan["results"]} for scan in scans + ] + names = sorted(set(by_name[0])) + if any(set(records) != set(names) for records in by_name[1:]): + raise RuntimeError("Scan configuration sets differ") + summary: list[dict[str, Any]] = [] + for name in names: + rows = [records[name] for records in by_name] + first = rows[0] + summary.append( + { + "name": name, + "head": first["head"], + "candidate_steps": first["candidate_steps"], + "beta": first["beta"], + "threshold": first["threshold"], + "target_accepts": first["target_accepts"], + "mean_accepted_predictor_calls": sum( + float(row["accepted_predictor_calls"]) for row in rows + ) + / len(rows), + "mean_full_calls": sum(float(row["full_calls"]) for row in rows) + / len(rows), + "mean_policy_latency_ms": sum( + float(row["policy_latency_ms"]) for row in rows + ) + / len(rows), + "mean_latent_nrmse_vs_ffff": sum( + float(row["latent_nrmse_vs_ffff"]) for row in rows + ) + / len(rows), + "per_prompt_accepts": [ + int(row["accepted_predictor_calls"]) for row in rows + ], + "per_prompt_latent_nrmse": [ + float(row["latent_nrmse_vs_ffff"]) for row in rows + ], + } + ) + selected: list[dict[str, Any]] = [] + for family, targets in TARGETS.items(): + for target in targets: + candidates = [ + row + for row in summary + if row["head"] == family and int(row["target_accepts"]) == target + ] + within_budget = [ + row + for row in candidates + if abs(float(row["mean_accepted_predictor_calls"]) - target) <= 0.5 + ] + pool = within_budget or candidates + choice = min( + pool, + key=lambda row: ( + float(row["mean_latent_nrmse_vs_ffff"]), + abs(float(row["mean_accepted_predictor_calls"]) - target), + float(row["mean_policy_latency_ms"]), + float(row["beta"]), + ) + if within_budget + else ( + abs(float(row["mean_accepted_predictor_calls"]) - target), + float(row["mean_latent_nrmse_vs_ffff"]), + float(row["mean_policy_latency_ms"]), + float(row["beta"]), + ), + ) + selected.append(choice) + atomic_json( + args.output_root / "calibration_summary.json", + { + "status": "complete", + "global_indices": [int(scan["global_index"]) for scan in scans], + "selection_rule": ( + "within mean K +/-0.5 choose lowest mean latent nRMSE; if no " + "candidate is within budget, choose nearest K first" + ), + "summary": summary, + "selected": selected, + }, + ) + print("[selected]", flush=True) + for row in selected: + print( + f" {row['head']} K{row['target_accepts']:02d}: beta={row['beta']} " + f"threshold={row['threshold']:.8f} " + f"actual_K={row['mean_accepted_predictor_calls']:.3f} " + f"latent_nrmse={row['mean_latent_nrmse_vs_ffff']:.6f}", + flush=True, + ) + + +def summarize_threshold_only(args: argparse.Namespace) -> None: + paths = sorted((args.output_root / "scans").glob("global_*.json")) + if len(paths) != 3: + raise ValueError(f"Expected exactly three scan files, found {len(paths)}") + scans = [json.loads(path.read_text(encoding="utf-8")) for path in paths] + if any(scan.get("status") != "complete" for scan in scans): + raise RuntimeError("At least one scan is incomplete") + by_name = [ + {record["name"]: record for record in scan["results"]} for scan in scans + ] + names = sorted(set(by_name[0])) + if any(set(records) != set(names) for records in by_name[1:]): + raise RuntimeError("Scan configuration sets differ") + summary: list[dict[str, Any]] = [] + for name in names: + rows = [records[name] for records in by_name] + first = rows[0] + summary.append( + { + "name": name, + "head": first["head"], + "candidate_steps": first["candidate_steps"], + "beta": first["beta"], + "threshold": first["threshold"], + "threshold_rank": first.get("threshold_rank"), + "mean_accepted_predictor_calls": sum( + float(row["accepted_predictor_calls"]) for row in rows + ) + / len(rows), + "mean_full_calls": sum(float(row["full_calls"]) for row in rows) + / len(rows), + "mean_policy_latency_ms": sum( + float(row["policy_latency_ms"]) for row in rows + ) + / len(rows), + "mean_latent_nrmse_vs_ffff": sum( + float(row["latent_nrmse_vs_ffff"]) for row in rows + ) + / len(rows), + "per_prompt_accepts": [ + int(row["accepted_predictor_calls"]) for row in rows + ], + } + ) + selected: list[dict[str, Any]] = [] + for family, targets in TARGETS.items(): + for target in targets: + candidates = [row for row in summary if row["head"] == family] + within_budget = [ + row + for row in candidates + if abs(float(row["mean_accepted_predictor_calls"]) - target) <= 0.5 + ] + pool = within_budget or candidates + choice = min( + pool, + key=lambda row: ( + float(row["mean_latent_nrmse_vs_ffff"]), + abs(float(row["mean_accepted_predictor_calls"]) - target), + float(row["threshold"]), + ) + if within_budget + else ( + abs(float(row["mean_accepted_predictor_calls"]) - target), + float(row["mean_latent_nrmse_vs_ffff"]), + float(row["threshold"]), + ), + ) + selected.append( + { + **choice, + "target_accepts": target, + } + ) + atomic_json( + args.output_root / "threshold_summary.json", + { + "status": "complete", + "beta": 2.0, + "global_indices": [int(scan["global_index"]) for scan in scans], + "selection_rule": ( + "fixed beta=2; for each target K choose threshold with mean K " + "within +/-0.5 and lowest latent nRMSE" + ), + "summary": summary, + "selected": selected, + }, + ) + print("[selected-threshold-only]", flush=True) + for row in selected: + print( + f" {row['head']} K{row['target_accepts']:02d}: " + f"threshold={row['threshold']:.8f} " + f"actual_K={row['mean_accepted_predictor_calls']:.3f} " + f"latent_nrmse={row['mean_latent_nrmse_vs_ffff']:.6f}", + flush=True, + ) + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument( + "--mode", + choices=( + "probe", "prepare", "prepare_threshold", "scan", "summarize", + "summarize_threshold", + ), + required=True, + ) + parser.add_argument("--gpu", default=PHYSICAL_GPU) + parser.add_argument("--global-index", type=int, default=None) + parser.add_argument("--mapping", type=Path, default=MAPPING_DEFAULT) + parser.add_argument("--experiment-root", type=Path, default=EXPERIMENT_DEFAULT) + parser.add_argument("--output-root", type=Path, default=OUTPUT_DEFAULT) + parser.add_argument("--probe-root", type=Path, default=None) + parser.add_argument("--seed", type=int, default=0) + parser.add_argument("--overwrite", action="store_true") + args = parser.parse_args() + args.mapping = args.mapping.resolve() + args.experiment_root = args.experiment_root.resolve() + args.output_root = args.output_root.resolve() + args.probe_root = args.probe_root.resolve() if args.probe_root is not None else None + return args + + +def main() -> None: + args = parse_args() + args.output_root.mkdir(parents=True, exist_ok=True) + if args.mode == "probe": + probe(args) + elif args.mode == "prepare": + prepare(args) + elif args.mode == "prepare_threshold": + prepare_threshold_only(args) + elif args.mode == "scan": + scan(args) + else: + if args.mode == "summarize_threshold": + summarize_threshold_only(args) + else: + summarize(args) + + +if __name__ == "__main__": + main() diff --git a/scripts/create_lmdb_14b_shards.py b/scripts/create_lmdb_14b_shards.py new file mode 100644 index 0000000000000000000000000000000000000000..bb0a76f3aec5bb8e6e09197dbc43a61affc37da8 --- /dev/null +++ b/scripts/create_lmdb_14b_shards.py @@ -0,0 +1,101 @@ +""" +python create_lmdb_14b_shards.py \ +--data_path /mnt/localssd/wanx_14b_data \ +--lmdb_path /mnt/localssd/wanx_14B_shift-3.0_cfg-5.0_lmdb +""" +from tqdm import tqdm +import numpy as np +import argparse +import torch +import lmdb +import glob +import os + +from utils.lmdb import store_arrays_to_lmdb, process_data_dict + + +def main(): + """ + Aggregate all ode pairs inside a folder into a lmdb dataset. + Each pt file should contain a (key, value) pair representing a + video's ODE trajectories. + """ + parser = argparse.ArgumentParser() + parser.add_argument("--data_path", type=str, + required=True, help="path to ode pairs") + parser.add_argument("--lmdb_path", type=str, + required=True, help="path to lmdb") + parser.add_argument("--num_shards", type=int, + default=16, help="num_shards") + + args = parser.parse_args() + + all_dirs = sorted(os.listdir(args.data_path)) + + # figure out the maximum map size needed + map_size = int(1e12) # adapt to your need, set to 1TB by default + os.makedirs(args.lmdb_path, exist_ok=True) + # 1) Open one LMDB env per shard + envs = [] + num_shards = args.num_shards + for shard_id in range(num_shards): + print("shard_id ", shard_id) + path = os.path.join(args.lmdb_path, f"shard_{shard_id}") + env = lmdb.open(path, + map_size=map_size, + subdir=True, # set to True if you want a directory per env + readonly=False, + metasync=True, + sync=True, + lock=True, + readahead=False, + meminit=False) + envs.append(env) + + counters = [0] * num_shards + seen_prompts = set() # for deduplication + total_samples = 0 + all_files = [] + + for part_dir in all_dirs: + all_files += sorted(glob.glob(os.path.join(args.data_path, part_dir, "*.pt"))) + + # 2) Prepare a write transaction for each shard + for idx, file in tqdm(enumerate(all_files)): + try: + data_dict = torch.load(file) + data_dict = process_data_dict(data_dict, seen_prompts) + except Exception as e: + print(f"Error processing {file}: {e}") + continue + + if data_dict["latents"].shape != (1, 21, 16, 60, 104): + continue + + shard_id = idx % num_shards + # write to lmdb file + store_arrays_to_lmdb(envs[shard_id], data_dict, start_index=counters[shard_id]) + counters[shard_id] += len(data_dict['prompts']) + data_shape = data_dict["latents"].shape + + total_samples += len(all_files) + + print(len(seen_prompts)) + + # save each entry's shape to lmdb + for shard_id, env in enumerate(envs): + with env.begin(write=True) as txn: + for key, val in (data_dict.items()): + assert len(data_shape) == 5 + array_shape = np.array(data_shape) # val.shape) + array_shape[0] = counters[shard_id] + shape_key = f"{key}_shape".encode() + print(shape_key, array_shape) + shape_str = " ".join(map(str, array_shape)) + txn.put(shape_key, shape_str.encode()) + + print(f"Finished writing {total_samples} examples into {num_shards} shards under {args.lmdb_path}") + + +if __name__ == "__main__": + main() diff --git a/scripts/create_lmdb_iterative.py b/scripts/create_lmdb_iterative.py new file mode 100644 index 0000000000000000000000000000000000000000..f77c2d4ff2b7559b93ed474f6082426562d4a41a --- /dev/null +++ b/scripts/create_lmdb_iterative.py @@ -0,0 +1,60 @@ +from tqdm import tqdm +import numpy as np +import argparse +import torch +import lmdb +import glob +import os + +from utils.lmdb import store_arrays_to_lmdb, process_data_dict + + +def main(): + """ + Aggregate all ode pairs inside a folder into a lmdb dataset. + Each pt file should contain a (key, value) pair representing a + video's ODE trajectories. + """ + parser = argparse.ArgumentParser() + parser.add_argument("--data_path", type=str, + required=True, help="path to ode pairs") + parser.add_argument("--lmdb_path", type=str, + required=True, help="path to lmdb") + + args = parser.parse_args() + + all_files = sorted(glob.glob(os.path.join(args.data_path, "*.pt"))) + + # figure out the maximum map size needed + total_array_size = 5000000000000 # adapt to your need, set to 5TB by default + + env = lmdb.open(args.lmdb_path, map_size=total_array_size * 2) + + counter = 0 + + seen_prompts = set() # for deduplication + + for index, file in tqdm(enumerate(all_files)): + # read from disk + data_dict = torch.load(file) + + data_dict = process_data_dict(data_dict, seen_prompts) + + # write to lmdb file + store_arrays_to_lmdb(env, data_dict, start_index=counter) + counter += len(data_dict['prompts']) + + # save each entry's shape to lmdb + with env.begin(write=True) as txn: + for key, val in data_dict.items(): + print(key, val) + array_shape = np.array(val.shape) + array_shape[0] = counter + + shape_key = f"{key}_shape".encode() + shape_str = " ".join(map(str, array_shape)) + txn.put(shape_key, shape_str.encode()) + + +if __name__ == "__main__": + main() diff --git a/scripts/eval_vbench8_extended_naive_baselines.py b/scripts/eval_vbench8_extended_naive_baselines.py new file mode 100644 index 0000000000000000000000000000000000000000..f404c9296e4180189e8068e64a9e8e5ba77faf84 --- /dev/null +++ b/scripts/eval_vbench8_extended_naive_baselines.py @@ -0,0 +1,176 @@ +#!/usr/bin/env python3 +"""Score generated naive baselines with VBench-8 and build final summaries. + +This script does not generate videos. Run +``generate_vbench8_extended_naive_baselines.py`` first, then use this entry +point to score FFFF plus any selected naive strategies. Strategies are scored +sequentially on the requested GPU; separate processes may be used on different +GPUs by passing disjoint ``--strategy`` sets and ``--no-summarize``. +""" + +from __future__ import annotations + +import argparse +import os +import subprocess +import sys +from pathlib import Path + + +REPO_ROOT = Path(__file__).resolve().parents[1] +if str(REPO_ROOT) not in sys.path: + sys.path.insert(0, str(REPO_ROOT)) + +from scripts.naive_vbench_policies import ( + ALL_STRATEGY_NAMES, + EVALUATION_STRATEGY_NAMES, +) + + +OUTPUT_DEFAULT = REPO_ROOT / "evaluation_runs/vbench8_extended_naive_baselines" +MAPPING_DEFAULT = REPO_ROOT / "assets/vbench8_extended_subset_mapping.json" +EXTENDED_PROMPTS_DEFAULT = REPO_ROOT / "prompts/vbench/all_dimension_extended.txt" +VBENCH_SITE_DEFAULT = REPO_ROOT / ".evaluation_env/vbench_site2" +VBENCH_INFO_DEFAULT = VBENCH_SITE_DEFAULT / "vbench/VBench_full_info.json" +VBENCH_CACHE_DEFAULT = Path("/data3/chenzhuo/.cache/vbench") + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--gpu", default="0") + parser.add_argument("--output-root", type=Path, default=OUTPUT_DEFAULT) + parser.add_argument("--mapping", type=Path, default=MAPPING_DEFAULT) + parser.add_argument( + "--extended-prompts", type=Path, default=EXTENDED_PROMPTS_DEFAULT + ) + parser.add_argument("--vbench-site", type=Path, default=VBENCH_SITE_DEFAULT) + parser.add_argument("--vbench-info", type=Path, default=VBENCH_INFO_DEFAULT) + parser.add_argument("--vbench-cache", type=Path, default=VBENCH_CACHE_DEFAULT) + parser.add_argument( + "--strategy", + action="append", + choices=EVALUATION_STRATEGY_NAMES, + default=None, + help="Naive strategy to score; repeat as needed. Default: all six.", + ) + parser.add_argument( + "--score-ffff", + action=argparse.BooleanOptionalAction, + default=True, + help="Also score the matched FFFF videos.", + ) + parser.add_argument( + "--skip-existing", action=argparse.BooleanOptionalAction, default=True + ) + parser.add_argument( + "--summarize", + action=argparse.BooleanOptionalAction, + default=True, + help="Build pixel and final summaries for this exact strategy set.", + ) + return parser.parse_args() + + +def checked_file(path: Path, label: str) -> Path: + resolved = path.resolve() + if not resolved.is_file(): + raise FileNotFoundError(f"Missing {label}: {resolved}") + return resolved + + +def run(command: list[str], env: dict[str, str]) -> None: + print("[run] " + " ".join(command), flush=True) + subprocess.run(command, cwd=REPO_ROOT, env=env, check=True) + + +def main() -> None: + args = parse_args() + selected = args.strategy or list(EVALUATION_STRATEGY_NAMES) + if len(set(selected)) != len(selected): + raise ValueError("--strategy values must be unique") + strategies = (["ffff"] if args.score_ffff else []) + selected + if args.summarize and "ffff" not in strategies: + raise ValueError("Final summaries require --score-ffff") + unknown = set(strategies) - set(ALL_STRATEGY_NAMES) + if unknown: + raise ValueError(f"Unknown strategies: {sorted(unknown)}") + + output_root = args.output_root.resolve() + mapping = checked_file(args.mapping, "mapping") + extended_prompts = checked_file(args.extended_prompts, "extended prompts") + vbench_site = args.vbench_site.resolve() + if not vbench_site.is_dir(): + raise FileNotFoundError(f"Missing VBench site directory: {vbench_site}") + vbench_info = checked_file(args.vbench_info, "VBench full info") + videos_root = output_root / "generated_videos" + if not videos_root.is_dir(): + raise FileNotFoundError(f"Missing generated videos: {videos_root}") + + env = os.environ.copy() + existing_pythonpath = env.get("PYTHONPATH") + pythonpath_parts = [str(vbench_site), str(REPO_ROOT)] + if existing_pythonpath: + pythonpath_parts.append(existing_pythonpath) + env["PYTHONPATH"] = os.pathsep.join(pythonpath_parts) + env["VBENCH_CACHE_DIR"] = str(args.vbench_cache.resolve()) + env["VBENCH_BERT_MODEL_DIR"] = str( + (args.vbench_cache / "bert-base-uncased").resolve() + ) + + if args.summarize: + run( + [ + sys.executable, + str(REPO_ROOT / "scripts/summarize_vbench8_generation.py"), + "--output-root", + str(output_root), + "--strategies", + *strategies, + ], + env, + ) + + for strategy in strategies: + command = [ + sys.executable, + str(REPO_ROOT / "scripts/eval_vbench8_extended_subset.py"), + "--gpu", + str(args.gpu), + "--strategy", + strategy, + "--mapping", + str(mapping), + "--videos-root", + str(videos_root), + "--output-root", + str(output_root), + "--vbench-info", + str(vbench_info), + ] + if args.skip_existing: + command.append("--skip-existing") + run(command, env) + + if args.summarize: + run( + [ + sys.executable, + str(REPO_ROOT / "scripts/summarize_vbench8_extended.py"), + "--output-root", + str(output_root), + "--mapping", + str(mapping), + "--extended-prompts", + str(extended_prompts), + "--vbench-info", + str(vbench_info), + "--strategies", + *strategies, + ], + env, + ) + print(f"[complete] strategies={','.join(strategies)}", flush=True) + + +if __name__ == "__main__": + main() diff --git a/scripts/eval_vbench8_extended_subset.py b/scripts/eval_vbench8_extended_subset.py new file mode 100644 index 0000000000000000000000000000000000000000..851f60ed0171fb6dbbbaed7b2cd6724ccc484b11 --- /dev/null +++ b/scripts/eval_vbench8_extended_subset.py @@ -0,0 +1,256 @@ +#!/usr/bin/env python3 +"""Run the standard VBench metrics for one generated strategy. + +The official VBench metric modules are used unchanged. Only the metadata +builder is replaced so that the evaluator reads the exact extended prompt used +for generation and the auditable mapping can use suite-indexed filenames. +""" + +from __future__ import annotations + +import argparse +import hashlib +import json +import os +import sys +import tempfile +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) + os.environ.setdefault("MPLCONFIGDIR", tempfile.mkdtemp(prefix="vbench_mpl_")) + return str(args.gpu) + + +PHYSICAL_GPU = preparse_gpu() + +_vbench_cache = Path( + os.environ.get("VBENCH_CACHE_DIR", "/data3/chenzhuo/.cache/vbench") +).expanduser() +os.environ.setdefault( + "VBENCH_BERT_MODEL_DIR", str(_vbench_cache / "bert-base-uncased") +) + +from vbench import VBench + +REPO_ROOT = Path(__file__).resolve().parents[1] +if str(REPO_ROOT) not in sys.path: + sys.path.insert(0, str(REPO_ROOT)) + +from scripts.vbench8_protocol import ( + DIMENSIONS, + PROTOCOL_NAME, + SUITE_COUNTS, + aggregate_selected_score, +) + + +def sha256(path: Path) -> str: + digest = hashlib.sha256() + with path.open("rb") as handle: + for block in iter(lambda: handle.read(1024 * 1024), b""): + digest.update(block) + return digest.hexdigest() + + +def read_mapping(path: Path) -> list[dict[str, Any]]: + value = json.loads(path.read_text(encoding="utf-8")) + if not isinstance(value, list) or len(value) != 251: + raise ValueError(f"Expected a 251-row mapping list: {path}") + counts = {suite: 0 for suite in SUITE_COUNTS} + globals_seen: set[int] = set() + suite_seen: dict[str, set[int]] = {suite: set() for suite in SUITE_COUNTS} + for row in value: + suite = str(row["prompt_suite"]) + global_index = int(row["global_index"]) + suite_index = int(row["suite_index"]) + if suite not in counts: + raise ValueError(f"Unknown suite {suite}") + if global_index in globals_seen: + raise ValueError(f"Duplicate global index {global_index}") + if suite_index in suite_seen[suite]: + raise ValueError(f"Duplicate suite index {suite}/{suite_index}") + globals_seen.add(global_index) + suite_seen[suite].add(suite_index) + counts[suite] += 1 + if counts != SUITE_COUNTS: + raise ValueError(f"Unexpected suite counts: {counts}") + return sorted(value, key=lambda row: int(row["global_index"])) + + +def build_metadata( + *, + mapping: list[dict[str, Any]], + videos_root: Path, + strategy: str, + metadata_path: Path, +) -> None: + records: list[dict[str, Any]] = [] + for row in mapping: + suite = str(row["prompt_suite"]) + suite_index = int(row["suite_index"]) + video_path = ( + videos_root / strategy / suite / f"{suite_index:03d}.mp4" + ).resolve() + if not video_path.is_file(): + raise FileNotFoundError(video_path) + official_dimensions = list(row["official_dimensions"]) + if not set(official_dimensions).issubset(set(DIMENSIONS)): + raise ValueError( + f"Unsupported dimensions at global index {row['global_index']}: " + f"{official_dimensions}" + ) + record: dict[str, Any] = { + "prompt_en": str(row["extended_prompt"]), + "dimension": official_dimensions, + "video_list": [str(video_path)], + } + # The standard scene implementation in vbench==0.1.5 needs the + # official scene keyword in auxiliary_info. It is not a replacement + # for prompt_en: overall_consistency still receives the extended text. + if "auxiliary_info" in row: + record["auxiliary_info"] = row["auxiliary_info"] + records.append(record) + if len(records) != 251: + raise ValueError(f"Expected 251 metadata records, got {len(records)}") + metadata_path.parent.mkdir(parents=True, exist_ok=True) + metadata_path.write_text( + json.dumps(records, ensure_ascii=False, indent=2) + "\n", encoding="utf-8" + ) + + +class FixedMetadataVBench(VBench): + """Use prepared metadata while retaining VBench's official evaluator.""" + + def __init__(self, *, device: str, metadata_path: Path, output_path: Path): + super().__init__( + device=device, + full_info_dir=str(metadata_path), + output_path=str(output_path), + ) + self.prepared_metadata_path = metadata_path + + def build_full_info_json(self, *args: Any, **kwargs: Any) -> str: + return str(self.prepared_metadata_path) + + +def parse_result(value: Any) -> tuple[float, Any]: + if isinstance(value, (list, tuple)) and len(value) == 2: + score, details = value + else: + score, details = value, None + score = float(score) + if not 0.0 <= score <= 1.0: + raise ValueError(f"VBench raw score is outside [0, 1]: {score}") + return score, details + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--gpu", default=PHYSICAL_GPU) + parser.add_argument("--strategy", required=True) + parser.add_argument("--mapping", type=Path, required=True) + parser.add_argument("--videos-root", type=Path, required=True) + parser.add_argument("--output-root", type=Path, required=True) + parser.add_argument("--vbench-info", type=Path, required=True) + parser.add_argument("--skip-existing", action="store_true") + return parser.parse_args() + + +def main() -> None: + args = parse_args() + mapping_path = args.mapping.resolve() + videos_root = args.videos_root.resolve() + output_root = args.output_root.resolve() + mapping = read_mapping(mapping_path) + metadata_path = output_root / "vbench/metadata" / f"{args.strategy}_full_info.json" + raw_dir = output_root / "vbench/raw_results" / args.strategy + raw_name = f"vbench8_{args.strategy}" + raw_result_path = raw_dir / f"{raw_name}_eval_results.json" + score_path = output_root / "vbench/scores" / f"{args.strategy}.json" + if args.skip_existing and raw_result_path.is_file() and score_path.is_file(): + print(f"[cached] strategy={args.strategy}", flush=True) + return + + build_metadata( + mapping=mapping, + videos_root=videos_root, + strategy=args.strategy, + metadata_path=metadata_path, + ) + raw_dir.mkdir(parents=True, exist_ok=True) + print( + f"[vbench] gpu={args.gpu} strategy={args.strategy} videos=251 " + f"dimensions={','.join(DIMENSIONS)}", + flush=True, + ) + bench = FixedMetadataVBench( + device="cuda", metadata_path=metadata_path, output_path=raw_dir + ) + bench.evaluate( + videos_path=str(videos_root / args.strategy), + name=raw_name, + dimension_list=list(DIMENSIONS), + local=True, + mode="vbench_standard", + ) + if not raw_result_path.is_file(): + raise FileNotFoundError(raw_result_path) + raw_json = json.loads(raw_result_path.read_text(encoding="utf-8")) + raw_scores: dict[str, float] = {} + details: dict[str, Any] = {} + for dimension in DIMENSIONS: + if dimension not in raw_json: + raise KeyError(f"Missing {dimension} in {raw_result_path}") + score, dimension_details = parse_result(raw_json[dimension]) + raw_scores[dimension] = score + if dimension_details is not None: + details[dimension] = dimension_details + aggregate = aggregate_selected_score(raw_scores) + result = { + "protocol": PROTOCOL_NAME, + "strategy": args.strategy, + "physical_gpu": str(args.gpu), + "benchmark": "VBench", + "vbench_version": "0.1.5", + "vbench_long": False, + "vbench_info": str(args.vbench_info.resolve()), + "vbench_info_sha256": sha256(args.vbench_info.resolve()), + "mapping": str(mapping_path), + "mapping_sha256": sha256(mapping_path), + "videos_root": str(videos_root / args.strategy), + "num_videos": 251, + "dimensions": list(DIMENSIONS), + "raw_scores": raw_scores, + **aggregate, + "scene_prompt_note": ( + "prompt_en is the Self-Forcing extended prompt; vbench==0.1.5 " + "scene uses the official auxiliary scene keyword by design." + ), + } + if details: + details_path = output_root / "vbench/details" / f"{args.strategy}.json" + details_path.parent.mkdir(parents=True, exist_ok=True) + details_path.write_text( + json.dumps(details, ensure_ascii=False, indent=2) + "\n", + encoding="utf-8", + ) + result["details"] = str(details_path) + score_path.parent.mkdir(parents=True, exist_ok=True) + score_path.write_text( + json.dumps(result, ensure_ascii=False, indent=2) + "\n", encoding="utf-8" + ) + print( + f"[complete] strategy={args.strategy} " + f"selected={result['selected_vbench_percent']:.4f}%", + flush=True, + ) + + +if __name__ == "__main__": + main() diff --git a/scripts/evaluate_layer17_chunk_impact.py b/scripts/evaluate_layer17_chunk_impact.py new file mode 100644 index 0000000000000000000000000000000000000000..65b00f26efa5d625a285263df911659aa56d093f --- /dev/null +++ b/scripts/evaluate_layer17_chunk_impact.py @@ -0,0 +1,782 @@ +#!/usr/bin/env python3 +"""Measure local and downstream impact of isolated Layer-17 Predictor calls. + +For every prompt, chunk 1..6 is independently evaluated with FPFF, FFPF, and +FPPF while every other chunk remains FFFF. A shadow Full forward is executed +at each selected Predictor step to measure hidden/flow/x0 error on the exact +rollout state. The shadow output is never used by the generated trajectory. +""" + +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 + +REPO_ROOT = Path(__file__).resolve().parents[1] +if str(REPO_ROOT) not in sys.path: + sys.path.insert(0, str(REPO_ROOT)) + +from scripts import evaluate_single_block_fppf as base +from utils.misc import set_seed +from utils.wan_wrapper import WanVAEWrapper + + +SCHEDULES = ("FPFF", "FFPF", "FPPF") +EPS = 1e-8 + + +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( + "--reference_root", + type=Path, + default=Path("outputs/single_block_fppf_eval"), + ) + parser.add_argument( + "--output_dir", + type=Path, + default=Path("outputs/layer17_chunk_impact_pilot"), + ) + parser.add_argument( + "--prompt_ids", type=int, nargs="*", default=list(range(80, 90)) + ) + parser.add_argument("--schedules", nargs="*", choices=SCHEDULES, default=list(SCHEDULES)) + parser.add_argument( + "--chunks", type=int, nargs="*", default=list(range(1, base.NUM_CHUNKS)) + ) + parser.add_argument("--max_prompts", 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( + "--skip_lpips", action=argparse.BooleanOptionalAction, default=False + ) + parser.add_argument("--overwrite", action="store_true") + args = parser.parse_args() + 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]") + if args.metric_batch_size < 1: + parser.error("--metric_batch_size must be positive") + if not args.schedules: + parser.error("At least one schedule is required") + if not args.chunks or any(chunk < 1 or chunk >= base.NUM_CHUNKS for chunk in args.chunks): + parser.error("--chunks must contain values in [1, 6]") + return args + + +def resolve(path: Path) -> Path: + return path.resolve() if path.is_absolute() else (REPO_ROOT / path).resolve() + + +def rms(value: torch.Tensor) -> torch.Tensor: + return value.float().square().mean().sqrt() + + +def nrmse(prediction: torch.Tensor, target: torch.Tensor) -> float: + return float(rms(prediction.float() - target.float()) / rms(target).clamp_min(EPS)) + + +def chunk_frame_slice(chunk: int) -> slice: + if chunk == 0: + return slice(0, base.PIXEL_FRAMES_FIRST_CHUNK) + start = base.PIXEL_FRAMES_FIRST_CHUNK + 12 * (chunk - 1) + return slice(start, start + 12) + + +def summarize_frame_range(metrics: dict[str, Any], selected: slice) -> dict[str, float]: + mse_values = metrics["mse_per_frame"][selected] + ssim_values = metrics["ssim_per_frame"][selected] + lpips_values = metrics["lpips_per_frame"][selected] + mean_mse = sum(mse_values) / len(mse_values) + return { + "pixel_mse": mean_mse, + "psnr": -10.0 * math.log10(max(mean_mse, 1e-12)), + "ssim": sum(ssim_values) / len(ssim_values), + "lpips": ( + sum(lpips_values) / len(lpips_values) if lpips_values else float("nan") + ), + } + + +@torch.inference_mode() +def generate_intervention( + *, + pipeline: Any, + dataset_root: Path, + prompt_id: int, + generation_seed: int, + device: torch.device, + predictor: Any, + source_layer: int, + intervention_chunk: int, + intervention_schedule: str, +) -> tuple[torch.Tensor, dict[str, Any]]: + if intervention_schedule not in SCHEDULES: + raise ValueError(intervention_schedule) + if intervention_chunk < 1 or intervention_chunk >= base.NUM_CHUNKS: + raise ValueError("Predictor intervention chunk must be 1..6") + + 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) + output_chunks: list[torch.Tensor] = [] + previous_chunk_hidden: list[torch.Tensor | None] | None = None + teacher = pipeline.generator.model + capture = base.FinalHiddenCapture(teacher) + full_calls = 0 + predictor_calls = 0 + shadow_full_calls = 0 + local_errors: list[dict[str, float | int]] = [] + 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 + selected = ( + chunk == intervention_chunk + and intervention_schedule[step] == "P" + ) + + if selected: + anchor_hidden = current_hidden[step - 1] + if anchor_hidden is None or previous_chunk_hidden is None: + raise RuntimeError("Predictor inputs are unavailable") + previous_hidden = previous_chunk_hidden[step] + if previous_hidden is None: + raise RuntimeError("Previous-chunk hidden is unavailable") + history = pipeline.kv_cache1[source_layer] + cross = pipeline.crossattn_cache[source_layer] + pred_hidden, pred_flow, _ = base.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 * base.TOKENS_PER_CHUNK, + ) + 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]) + + # Shadow Full measures the exact counterfactual target on + # this rollout state. Its x0/hidden are never accepted. + capture.start() + full_flow, full_x0 = pipeline.generator( + noisy_image_or_video=noisy_input, + conditional_dict={ + "prompt_embeds": torch.zeros( + 1, + 1, + int(teacher.text_embedding[0].in_features), + dtype=torch.bfloat16, + device=device, + ) + }, + timestep=timestep, + kv_cache=pipeline.kv_cache1, + crossattn_cache=pipeline.crossattn_cache, + current_start=chunk * base.TOKENS_PER_CHUNK, + ) + full_hidden = capture.finish() + local_errors.append( + { + "step": step, + "timestep": float(current_timestep), + "hidden_nrmse": nrmse(pred_hidden, full_hidden), + "flow_nrmse": nrmse(pred_flow, full_flow), + "x0_nrmse": nrmse(pred_x0, full_x0), + } + ) + current_hidden[step] = pred_hidden + denoised_pred = pred_x0 + predictor_calls += 1 + shadow_full_calls += 1 + del full_flow, full_x0, full_hidden + else: + capture.start() + _, denoised_pred = pipeline.generator( + noisy_image_or_video=noisy_input, + conditional_dict={ + "prompt_embeds": torch.zeros( + 1, + 1, + int(teacher.text_embedding[0].in_features), + dtype=torch.bfloat16, + device=device, + ) + }, + timestep=timestep, + kv_cache=pipeline.kv_cache1, + crossattn_cache=pipeline.crossattn_cache, + current_start=chunk * base.TOKENS_PER_CHUNK, + ) + current_hidden[step] = capture.finish() + full_calls += 1 + + 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 + pipeline.generator( + noisy_image_or_video=denoised_pred, + conditional_dict={ + "prompt_embeds": torch.zeros( + 1, + 1, + int(teacher.text_embedding[0].in_features), + dtype=torch.bfloat16, + device=device, + ) + }, + timestep=context_timestep, + kv_cache=pipeline.kv_cache1, + crossattn_cache=pipeline.crossattn_cache, + current_start=chunk * base.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, + "shadow_full_calls": shadow_full_calls, + "local_errors": local_errors, + } + + +def average(values: list[float]) -> float: + return sum(values) / len(values) + + +def finite_average(values: list[Any]) -> float: + numeric = [ + float(value) + for value in values + if value is not None and math.isfinite(float(value)) + ] + return average(numeric) if numeric else float("nan") + + +def rankdata(values: list[float]) -> list[float]: + order = sorted(range(len(values)), key=values.__getitem__) + ranks = [0.0] * len(values) + start = 0 + while start < len(order): + end = start + 1 + while end < len(order) and values[order[end]] == values[order[start]]: + end += 1 + rank = 0.5 * (start + end - 1) + for position in range(start, end): + ranks[order[position]] = rank + start = end + return ranks + + +def pearson(left: list[float], right: list[float]) -> float: + left_mean, right_mean = average(left), average(right) + left_centered = [value - left_mean for value in left] + right_centered = [value - right_mean for value in right] + numerator = sum(a * b for a, b in zip(left_centered, right_centered)) + denominator = math.sqrt( + sum(value * value for value in left_centered) + * sum(value * value for value in right_centered) + ) + return numerator / denominator if denominator > 0 else float("nan") + + +def spearman(left: list[float], right: list[float]) -> float: + return pearson(rankdata(left), rankdata(right)) + + +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 group_centered_values( + records: list[dict[str, Any]], field: str +) -> list[float]: + """Remove schedule-by-chunk means to isolate prompt/state variation.""" + groups: dict[tuple[str, int], list[float]] = {} + for row in records: + key = (str(row["schedule"]), int(row["chunk"])) + groups.setdefault(key, []).append(float(row[field])) + means = {key: average(values) for key, values in groups.items()} + return [ + float(row[field]) - means[(str(row["schedule"]), int(row["chunk"]))] + for row in records + ] + + +def aggregate(records: list[dict[str, Any]], output_dir: Path) -> None: + numeric_fields = [ + "hidden_nrmse_mean", + "flow_nrmse_mean", + "x0_nrmse_mean", + "latent_all_nrmse", + "latent_current_nrmse", + "latent_tail_nrmse", + "psnr", + "ssim", + "lpips", + "current_psnr", + "current_ssim", + "current_lpips", + "tail_psnr", + "tail_ssim", + "tail_lpips", + "generation_time_s", + ] + summary_rows: list[dict[str, Any]] = [] + available_schedules = [ + schedule for schedule in SCHEDULES if any(row["schedule"] == schedule for row in records) + ] + available_chunks = sorted({int(row["chunk"]) for row in records}) + for schedule in available_schedules: + for chunk in available_chunks: + selected = [ + row + for row in records + if row["schedule"] == schedule and row["chunk"] == chunk + ] + row: dict[str, Any] = { + "schedule": schedule, + "chunk": chunk, + "num_prompts": len(selected), + "alpha_eligible_linear": (base.NUM_CHUNKS - 1 - chunk) + / (base.NUM_CHUNKS - 2), + } + for field in numeric_fields: + row[field] = finite_average([item[field] for item in selected]) + summary_rows.append(row) + + summary_fields = [ + "schedule", + "chunk", + "num_prompts", + "alpha_eligible_linear", + *numeric_fields, + ] + write_csv(output_dir / "summary_by_schedule_chunk.csv", summary_rows, summary_fields) + + correlation_rows: list[dict[str, Any]] = [] + for schedule in (*available_schedules, "ALL"): + selected = ( + records if schedule == "ALL" else [r for r in records if r["schedule"] == schedule] + ) + for local in ("hidden_nrmse_mean", "flow_nrmse_mean", "x0_nrmse_mean"): + for downstream in ("tail_lpips", "latent_tail_nrmse", "tail_pixel_mse"): + pairs = [ + (float(row[local]), float(row[downstream])) + for row in selected + if row[local] is not None + and row[downstream] is not None + and math.isfinite(float(row[local])) + and math.isfinite(float(row[downstream])) + ] + correlation_rows.append( + { + "schedule": schedule, + "local_metric": local, + "downstream_metric": downstream, + "num_observations": len(pairs), + "spearman": ( + spearman( + [pair[0] for pair in pairs], + [pair[1] for pair in pairs], + ) + if len(pairs) >= 2 + else float("nan") + ), + } + ) + correlation_fields = [ + "schedule", + "local_metric", + "downstream_metric", + "num_observations", + "spearman", + ] + write_csv(output_dir / "local_downstream_correlations.csv", correlation_rows, correlation_fields) + + controlled_rows: list[dict[str, Any]] = [] + for local in ("hidden_nrmse_mean", "flow_nrmse_mean", "x0_nrmse_mean"): + local_residual = group_centered_values(records, local) + downstream_residual = group_centered_values(records, "tail_lpips") + within_cell = [] + for schedule in available_schedules: + for chunk in available_chunks: + selected = [ + row + for row in records + if row["schedule"] == schedule and row["chunk"] == chunk + ] + if len(selected) >= 2: + within_cell.append( + spearman( + [float(row[local]) for row in selected], + [float(row["tail_lpips"]) for row in selected], + ) + ) + controlled_rows.append( + { + "local_metric": local, + "downstream_metric": "tail_lpips", + "controls": "schedule+chunk", + "residual_spearman": spearman(local_residual, downstream_residual), + "mean_within_cell_spearman": finite_average(within_cell), + "positive_cells": sum(value > 0 for value in within_cell), + "num_cells": len(within_cell), + } + ) + controlled_fields = [ + "local_metric", + "downstream_metric", + "controls", + "residual_spearman", + "mean_within_cell_spearman", + "positive_cells", + "num_cells", + ] + write_csv( + output_dir / "controlled_local_downstream_correlations.csv", + controlled_rows, + controlled_fields, + ) + + report = [ + "# Layer-17 Predictor chunk-impact pilot", + "", + f"Prompts: {len(set(row['prompt_id'] for row in records))}; seed 0; " + "all non-intervened chunks use FFFF.", + "", + "A shadow Full call measures local error at each Predictor decision, but the " + "generated trajectory always consumes the Predictor output.", + "", + ] + for schedule in available_schedules: + report.extend( + [ + f"## {schedule}", + "", + "| Chunk | x0 nRMSE | Tail PSNR | Tail LPIPS | Tail latent nRMSE |", + "|---:|---:|---:|---:|---:|", + ] + ) + for row in summary_rows: + if row["schedule"] != schedule: + continue + report.append( + f"| {row['chunk']} | {row['x0_nrmse_mean']:.6f} | " + f"{row['tail_psnr']:.4f} | {row['tail_lpips']:.6f} | " + f"{row['latent_tail_nrmse']:.6f} |" + ) + report.append("") + report.extend( + [ + "## Local-to-downstream Spearman correlations", + "", + "| Schedule | Local metric | Downstream metric | Spearman |", + "|---|---|---|---:|", + ] + ) + for row in correlation_rows: + if row["downstream_metric"] == "tail_lpips": + report.append( + f"| {row['schedule']} | {row['local_metric']} | tail LPIPS | " + f"{row['spearman']:.4f} |" + ) + report.extend( + [ + "", + "## Correlations after controlling schedule and chunk", + "", + "Residual correlations remove each schedule-by-chunk mean, so they test " + "whether local error explains prompt/state risk beyond the position prior.", + "", + "| Local metric | Residual Spearman | Mean within-cell Spearman | Positive cells |", + "|---|---:|---:|---:|", + ] + ) + for row in controlled_rows: + report.append( + f"| {row['local_metric']} | {row['residual_spearman']:.4f} | " + f"{row['mean_within_cell_spearman']:.4f} | " + f"{row['positive_cells']}/{row['num_cells']} |" + ) + (output_dir / "REPORT.md").write_text("\n".join(report) + "\n", encoding="utf-8") + + +def main() -> None: + args = parse_args() + for name in ( + "config_path", + "checkpoint_path", + "dataset_root", + "sweep_dir", + "reference_root", + "output_dir", + ): + setattr(args, name, resolve(getattr(args, name))) + 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] + + config = OmegaConf.merge( + OmegaConf.load(REPO_ROOT / "configs/default_config.yaml"), + OmegaConf.load(args.config_path), + ) + device = torch.device("cuda") + torch.set_grad_enabled(False) + set_seed(args.generation_seed) + manifest = { + "status": "running", + "gpu": str(args.gpu), + "prompt_ids": prompt_ids, + "generation_seed": args.generation_seed, + "checkpoint_path": str(args.checkpoint_path), + "predictor_weights": str( + args.sweep_dir / "teacher_layer_17" / "predictor_final.safetensors" + ), + "dataset_root": str(args.dataset_root), + "schedules": list(args.schedules), + "chunks": list(args.chunks), + "shadow_full_target": True, + } + base.atomic_json(args.output_dir / "manifest.json", manifest) + + print("[setup] loading VAE", flush=True) + vae = WanVAEWrapper().to(device=device, dtype=torch.bfloat16).eval() + 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, + ) + + print("[setup] loading frozen generator and Layer-17 Predictor", flush=True) + pipeline = base.build_pipeline(config, args.checkpoint_path, vae, device) + experiment = base.discover_experiments( + args.sweep_dir, ["teacher_layer_17"], None + )[0] + predictor = base.load_predictor(pipeline.generator.model, experiment, device) + lpips_model = None + if not args.skip_lpips: + lpips_model = lpips.LPIPS(net="alex", verbose=False).to(device).eval() + lpips_model.requires_grad_(False) + + records: list[dict[str, Any]] = [] + total = len(prompt_ids) * len(args.schedules) * len(args.chunks) + completed = 0 + for prompt_id in prompt_ids: + 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) + for schedule in args.schedules: + for chunk in args.chunks: + destination = ( + args.output_dir + / "per_intervention" + / f"prompt_{prompt_id:04d}_{schedule}_chunk_{chunk:02d}.json" + ) + if destination.exists() and not args.overwrite: + record = json.loads(destination.read_text(encoding="utf-8")) + records.append(record) + completed += 1 + print(f"[cached] {completed}/{total} {destination.stem}", flush=True) + continue + started = time.perf_counter() + latent, diagnostics = generate_intervention( + pipeline=pipeline, + dataset_root=args.dataset_root, + prompt_id=prompt_id, + generation_seed=args.generation_seed, + device=device, + predictor=predictor, + source_layer=17, + intervention_chunk=chunk, + intervention_schedule=schedule, + ) + 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) + frame = base.frame_metrics( + reference_u8=reference_u8, + prediction_u8=prediction_u8, + lpips_model=lpips_model, + batch_size=args.metric_batch_size, + device=device, + ) + current_slice = chunk_frame_slice(chunk) + current = summarize_frame_range(frame, current_slice) + tail = summarize_frame_range(frame, slice(current_slice.start, None)) + latent_chunk_start = chunk * base.FRAMES_PER_CHUNK + latent_chunk_end = latent_chunk_start + base.FRAMES_PER_CHUNK + errors = diagnostics.pop("local_errors") + record = { + "prompt_id": prompt_id, + "schedule": schedule, + "chunk": chunk, + "predictor_steps": [int(item["step"]) for item in errors], + "hidden_nrmse_mean": average( + [float(item["hidden_nrmse"]) for item in errors] + ), + "flow_nrmse_mean": average( + [float(item["flow_nrmse"]) for item in errors] + ), + "x0_nrmse_mean": average( + [float(item["x0_nrmse"]) for item in errors] + ), + "local_errors": errors, + "latent_all_nrmse": nrmse(latent, reference_latent), + "latent_current_nrmse": nrmse( + latent[:, latent_chunk_start:latent_chunk_end], + reference_latent[:, latent_chunk_start:latent_chunk_end], + ), + "latent_tail_nrmse": nrmse( + latent[:, latent_chunk_start:], + reference_latent[:, latent_chunk_start:], + ), + "psnr": frame["psnr"], + "ssim": frame["ssim"], + "lpips": frame["lpips"], + "current_pixel_mse": current["pixel_mse"], + "current_psnr": current["psnr"], + "current_ssim": current["ssim"], + "current_lpips": current["lpips"], + "tail_pixel_mse": tail["pixel_mse"], + "tail_psnr": tail["psnr"], + "tail_ssim": tail["ssim"], + "tail_lpips": tail["lpips"], + **diagnostics, + "total_time_s": time.perf_counter() - started, + } + base.atomic_json(destination, record) + records.append(record) + completed += 1 + print( + f"[run] {completed}/{total} p={prompt_id} {schedule} c={chunk} " + f"x0={record['x0_nrmse_mean']:.5f} " + 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 latent, pixels, prediction_u8, frame + torch.cuda.empty_cache() + del reference_latent, reference_u8 + + fields = sorted({key for record in records for key in record if key != "local_errors"}) + flattened = [{key: row.get(key) for key in fields} for row in records] + write_csv(args.output_dir / "interventions.csv", flattened, fields) + aggregate(records, args.output_dir) + manifest["status"] = "complete" + base.atomic_json(args.output_dir / "manifest.json", manifest) + print(f"[complete] results -> {args.output_dir}", flush=True) + + +if __name__ == "__main__": + main() diff --git a/scripts/evaluate_layer17_dynamic_gate.py b/scripts/evaluate_layer17_dynamic_gate.py new file mode 100644 index 0000000000000000000000000000000000000000..20729f54b4a3b95a77575e24d2e86ef873a8cce1 --- /dev/null +++ b/scripts/evaluate_layer17_dynamic_gate.py @@ -0,0 +1,904 @@ +#!/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() diff --git a/scripts/evaluate_layer17_moviebench_step2000.py b/scripts/evaluate_layer17_moviebench_step2000.py new file mode 100644 index 0000000000000000000000000000000000000000..0aaeafc4167f3fda77fceb38694ce423777fe2e5 --- /dev/null +++ b/scripts/evaluate_layer17_moviebench_step2000.py @@ -0,0 +1,271 @@ +#!/usr/bin/env python3 +"""Evaluate a trained Self-Forcing Layer-17 predictor on MovieBench prompts.""" + +from __future__ import annotations + +import argparse +import json +import os +import sys +import time +from pathlib import Path + + +def preparse_gpu() -> str: + parser = argparse.ArgumentParser(add_help=False) + parser.add_argument("--gpu", required=True) + args, _ = parser.parse_known_args() + os.environ["CUDA_VISIBLE_DEVICES"] = args.gpu + return args.gpu + + +GPU = preparse_gpu() + +import lpips +import torch +from omegaconf import OmegaConf + +ROOT = Path(__file__).resolve().parents[1] +if str(ROOT) not in sys.path: + sys.path.insert(0, str(ROOT)) + +from predictor_training.offline_data import TOKENS_PER_CHUNK +from scripts.evaluate_single_block_fppf import ( + FinalHiddenCapture, + atomic_json, + build_pipeline, + frame_metrics, + load_predictor, + pixels_to_u8, + predictor_step, + save_mp4, +) +from utils.misc import set_seed +from utils.wan_wrapper import WanTextEncoder, WanVAEWrapper + + +NUM_CHUNKS = 7 +FRAMES_PER_CHUNK = 3 +NUM_STEPS = 4 +LATENT_CHANNELS = 16 +LATENT_HEIGHT = 60 +LATENT_WIDTH = 104 + + +def read_lines(path: Path) -> list[str]: + return [line.strip() for line in path.read_text(encoding="utf-8").splitlines() if line.strip()] + + +def reset_caches(pipeline, 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_() + for cache in pipeline.crossattn_cache: + cache["is_init"] = False + + +@torch.inference_mode() +def rollout(pipeline, conditional_dict, seed: int, device: torch.device, predictor=None): + reset_caches(pipeline, device) + set_seed(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 + timesteps = pipeline.denoising_step_list.to(device=device) + outputs = [] + previous_chunk_hidden = None + capture = FinalHiddenCapture(teacher) + full_calls = predictor_calls = 0 + started = time.perf_counter() + try: + for chunk in range(NUM_CHUNKS): + noisy_input = noise[:, chunk * 3 : (chunk + 1) * 3] + current_hidden = [None] * NUM_STEPS + denoised_pred = timestep = None + for step, current_timestep in enumerate(timesteps): + timestep = torch.ones([1, 3], dtype=torch.int64, device=device) * current_timestep + use_predictor = predictor is not None and chunk > 0 and step in {1, 2} + if use_predictor: + hidden, flow, _ = predictor_step( + predictor=predictor, + teacher=teacher, + noisy_input=noisy_input, + timestep=timestep, + anchor_hidden=current_hidden[step - 1], + previous_hidden=previous_chunk_hidden[step], + history_cache=pipeline.kv_cache1[17], + cross_cache=pipeline.crossattn_cache[17], + 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] = 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_STEPS - 1: + flat = denoised_pred.flatten(0, 1) + noisy_input = pipeline.scheduler.add_noise( + flat, + torch.randn_like(flat), + timesteps[step + 1] * torch.ones([3], dtype=torch.long, device=device), + ).unflatten(0, denoised_pred.shape[:2]) + outputs.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(outputs, dim=1), { + "generation_time_s": time.perf_counter() - started, + "full_calls": full_calls, + "predictor_calls": predictor_calls, + } + + +def main() -> None: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--gpu", default=GPU) + parser.add_argument("--prompt_ids", type=int, nargs="+", required=True) + parser.add_argument("--original_prompts", type=Path, required=True) + parser.add_argument("--extended_prompts", type=Path, required=True) + parser.add_argument("--output_dir", type=Path, required=True) + parser.add_argument("--weights", type=Path, required=True) + parser.add_argument("--seed", type=int, default=0) + parser.add_argument("--metric_batch_size", type=int, default=4) + args = parser.parse_args() + args.output_dir.mkdir(parents=True, exist_ok=True) + original = read_lines(args.original_prompts) + extended = read_lines(args.extended_prompts) + if len(original) != len(extended) or min(args.prompt_ids) < 0 or max(args.prompt_ids) >= len(original): + raise ValueError("MovieBench original/extended prompt pairing is invalid") + + atomic_json( + args.output_dir / "manifest.json", + { + "status": "running", + "gpu": args.gpu, + "prompt_ids": args.prompt_ids, + "weights": str(args.weights.resolve()), + "teacher": str((ROOT / "checkpoints/self_forcing_dmd.pt").resolve()), + "generation_seed_reset_per_prompt": args.seed, + "generation_prompts": str(args.extended_prompts.resolve()), + "evaluation_prompts": str(args.original_prompts.resolve()), + "schedule": "chunk0=FFFF; chunks1-6=FPPF", + "metrics": ["PSNR", "SSIM", "LPIPS", "rollout-only PSNR/SSIM/LPIPS"], + }, + ) + + device = torch.device("cuda") + torch.set_grad_enabled(False) + config = OmegaConf.merge( + OmegaConf.load(ROOT / "configs/default_config.yaml"), + OmegaConf.load(ROOT / "configs/self_forcing_sid.yaml"), + ) + vae = WanVAEWrapper().to(device=device, dtype=torch.bfloat16).eval() + pipeline = build_pipeline(config, ROOT / "checkpoints/self_forcing_dmd.pt", vae, device) + text_encoder = WanTextEncoder().to(device=device, dtype=torch.bfloat16).eval() + text_encoder.requires_grad_(False) + predictor = load_predictor( + pipeline.generator.model, + {"source_layer": 17, "weights": args.weights, "gate_mode": "baseline"}, + device, + ) + lpips_model = lpips.LPIPS(net="alex", verbose=False).to(device).eval() + lpips_model.requires_grad_(False) + + reference_dir = args.output_dir / "videos" / "ffff" + prediction_dir = args.output_dir / "videos" / "fppf_step2000" + for offset, prompt_id in enumerate(args.prompt_ids, start=1): + result_path = args.output_dir / "per_prompt" / f"prompt_{prompt_id:04d}.json" + if ( + result_path.exists() + and (reference_dir / f"{prompt_id:05d}.mp4").exists() + and (prediction_dir / f"{prompt_id:05d}.mp4").exists() + ): + print(f"[skip] {offset}/{len(args.prompt_ids)} id={prompt_id}", flush=True) + continue + print(f"[encode] {offset}/{len(args.prompt_ids)} id={prompt_id}", flush=True) + conditional = text_encoder(text_prompts=[extended[prompt_id]]) + reference_latent, ffff_counts = rollout(pipeline, conditional, args.seed, device) + prediction_latent, fppf_counts = rollout( + pipeline, conditional, args.seed, device, predictor=predictor + ) + with torch.autocast(device_type="cuda", dtype=torch.bfloat16): + reference_pixels = vae.decode_to_pixel(reference_latent, use_cache=False) + prediction_pixels = vae.decode_to_pixel(prediction_latent, use_cache=False) + reference_u8 = pixels_to_u8(reference_pixels) + prediction_u8 = pixels_to_u8(prediction_pixels) + save_mp4(reference_u8, reference_dir / f"{prompt_id:05d}.mp4") + save_mp4(prediction_u8, prediction_dir / f"{prompt_id:05d}.mp4") + metrics = frame_metrics( + reference_u8=reference_u8, + prediction_u8=prediction_u8, + lpips_model=lpips_model, + batch_size=args.metric_batch_size, + device=device, + ) + atomic_json( + result_path, + { + "status": "complete", + "prompt_id": prompt_id, + "original_prompt": original[prompt_id], + "generation_prompt": extended[prompt_id], + "seed": args.seed, + "latent_frames": NUM_CHUNKS * FRAMES_PER_CHUNK, + "decoded_frames": metrics["num_frames"], + "schedule": "chunk0=FFFF; chunks1-6=FPPF", + "ffff": ffff_counts, + "fppf": fppf_counts, + **metrics, + }, + ) + print( + f"[result] id={prompt_id} psnr={metrics['psnr']:.4f} " + f"ssim={metrics['ssim']:.6f} lpips={metrics['lpips']:.6f}", + flush=True, + ) + if hasattr(vae.model, "clear_cache"): + vae.model.clear_cache() + del conditional, reference_latent, prediction_latent, reference_pixels + del prediction_pixels, reference_u8, prediction_u8 + torch.cuda.empty_cache() + + manifest = json.loads((args.output_dir / "manifest.json").read_text(encoding="utf-8")) + manifest["status"] = "complete" + atomic_json(args.output_dir / "manifest.json", manifest) + + +if __name__ == "__main__": + main() diff --git a/scripts/evaluate_long_video_fppf.py b/scripts/evaluate_long_video_fppf.py new file mode 100644 index 0000000000000000000000000000000000000000..0ef584bf7f06b0959461f6576a238c5a56c7b7ea --- /dev/null +++ b/scripts/evaluate_long_video_fppf.py @@ -0,0 +1,295 @@ +#!/usr/bin/env python3 +"""Evaluate 2x/4x long FFFF and one-block Layer-17 FPPF rollouts.""" + +from __future__ import annotations + +import argparse +import json +import os +import sys +import time +from pathlib import Path + + +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 +from omegaconf import OmegaConf +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 predictor_training.offline_data import TOKENS_PER_CHUNK +from scripts.evaluate_single_block_fppf import ( + DEFAULT_PROMPT_IDS, + FRAMES_PER_CHUNK, + LATENT_CHANNELS, + LATENT_HEIGHT, + LATENT_WIDTH, + NUM_DENOISING_STEPS, + FinalHiddenCapture, + atomic_json, + build_pipeline, + discover_experiments, + frame_metrics, + load_predictor, + load_prompt_metadata, + pixels_to_u8, + predictor_step, + reset_kv_and_load_cross_cache, +) +from utils.misc import set_seed +from utils.wan_wrapper import WanVAEWrapper + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--gpu", default=PHYSICAL_GPU) + parser.add_argument("--prompt_ids", type=int, nargs="+", required=True) + parser.add_argument("--latent_lengths", type=int, nargs="+", default=[42, 84]) + 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/long_video_2x4x_eval") + ) + parser.add_argument("--metric_batch_size", type=int, default=4) + parser.add_argument("--generation_seed", type=int, default=0) + args = parser.parse_args() + if any(length <= 0 or length % FRAMES_PER_CHUNK for length in args.latent_lengths): + parser.error("Latent lengths must be positive multiples of 3") + if any(prompt not in DEFAULT_PROMPT_IDS for prompt in args.prompt_ids): + parser.error("This evaluation is restricted to validation prompt IDs 80..99") + return args + + +def resolve(path: Path) -> Path: + return path.resolve() if path.is_absolute() else (REPO_ROOT / path).resolve() + + +@torch.inference_mode() +def generate_rollout( + *, pipeline, dataset_root: Path, prompt_id: int, latent_length: int, + generation_seed: int, device: torch.device, predictor, 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 the Predictor") + num_chunks = latent_length // FRAMES_PER_CHUNK + reset_kv_and_load_cross_cache(pipeline, dataset_root, prompt_id, device) + set_seed(generation_seed) + noise = torch.randn( + 1, latent_length, 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 = [] + previous_chunk_hidden = None + capture = FinalHiddenCapture(teacher) + full_calls = predictor_calls = 0 + started = time.perf_counter() + try: + for chunk in range(num_chunks): + noisy_input = noise[:, chunk * 3:(chunk + 1) * 3] + current_hidden = [None] * NUM_DENOISING_STEPS + denoised_pred = timestep = 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: + pred_hidden, flow, _ = predictor_step( + predictor=predictor, + teacher=teacher, + noisy_input=noisy_input, + timestep=timestep, + anchor_hidden=current_hidden[step - 1], + previous_hidden=previous_chunk_hidden[step], + history_cache=pipeline.kv_cache1[source_layer], + cross_cache=pipeline.crossattn_cache[source_layer], + 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: + flat = denoised_pred.flatten(0, 1) + noisy_input = pipeline.scheduler.add_noise( + flat, torch.randn_like(flat), + timesteps[step + 1] * torch.ones( + [FRAMES_PER_CHUNK], dtype=torch.long, device=device + ), + ).unflatten(0, denoised_pred.shape[:2]) + 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, + "num_chunks": num_chunks, + } + + +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 main() -> None: + args = parse_args() + for field in ("config_path", "checkpoint_path", "dataset_root", "sweep_dir", "output_dir"): + setattr(args, field, resolve(getattr(args, field))) + args.output_dir.mkdir(parents=True, exist_ok=True) + atomic_json(args.output_dir / "manifest.json", { + "status": "running", "physical_gpu": str(args.gpu), + "prompt_ids": args.prompt_ids, "latent_lengths": args.latent_lengths, + "methods": ["FFFF", "FPPF_teacher_layer_17"], + "generation_seed": args.generation_seed, + }) + 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), + ) + # The released checkpoint uses full attention over its 21-latent training + # horizon. Long inference keeps exactly that horizon as a rolling window; + # within the first 21 latents this is numerically the same attention span. + config.model_kwargs.local_attn_size = 21 + vae = WanVAEWrapper().to(device=device, dtype=torch.bfloat16).eval() + pipeline = build_pipeline(config, args.checkpoint_path, vae, device) + experiment = discover_experiments( + args.sweep_dir, ["teacher_layer_17"], None + )[0] + predictor = load_predictor(pipeline.generator.model, experiment, device) + lpips_model = lpips.LPIPS(net="alex", verbose=False).to(device).eval() + lpips_model.requires_grad_(False) + + for prompt_id in args.prompt_ids: + prompt = load_prompt_metadata(args.dataset_root, prompt_id)["prompt"] + for latent_length in args.latent_lengths: + run_dir = args.output_dir / f"latent_{latent_length}" / f"prompt_{prompt_id:04d}" + result_path = run_dir / "metrics.json" + if result_path.exists(): + existing = json.loads(result_path.read_text(encoding="utf-8")) + if existing.get("status") == "complete": + print(f"[skip] latent={latent_length} prompt={prompt_id}", flush=True) + continue + print(f"[run] latent={latent_length} prompt={prompt_id} FFFF", flush=True) + reference_latent, ffff_counts = generate_rollout( + pipeline=pipeline, dataset_root=args.dataset_root, + prompt_id=prompt_id, latent_length=latent_length, + generation_seed=args.generation_seed, device=device, + predictor=None, source_layer=None, schedule="FFFF", + ) + with torch.autocast(device_type="cuda", dtype=torch.bfloat16): + reference_pixels = vae.decode_to_pixel(reference_latent, use_cache=False) + reference_u8 = pixels_to_u8(reference_pixels) + save_mp4(reference_u8, run_dir / "ffff.mp4") + del reference_latent, reference_pixels + if hasattr(vae.model, "clear_cache"): + vae.model.clear_cache() + torch.cuda.empty_cache() + + print(f"[run] latent={latent_length} prompt={prompt_id} FPPF", flush=True) + prediction_latent, fppf_counts = generate_rollout( + pipeline=pipeline, dataset_root=args.dataset_root, + prompt_id=prompt_id, latent_length=latent_length, + generation_seed=args.generation_seed, device=device, + predictor=predictor, source_layer=17, schedule="FPPF", + ) + with torch.autocast(device_type="cuda", dtype=torch.bfloat16): + prediction_pixels = vae.decode_to_pixel(prediction_latent, use_cache=False) + prediction_u8 = pixels_to_u8(prediction_pixels) + save_mp4(prediction_u8, run_dir / "fppf_layer17.mp4") + metrics = frame_metrics( + reference_u8=reference_u8, prediction_u8=prediction_u8, + lpips_model=lpips_model, batch_size=args.metric_batch_size, + device=device, + ) + atomic_json(result_path, { + "status": "complete", "prompt_id": prompt_id, "prompt": prompt, + "latent_length": latent_length, "decoded_frames": metrics["num_frames"], + "reference": "FFFF same prompt/seed/noise", + "predictor": "single_block_teacher_layer_17", + "ffff": ffff_counts, "fppf": fppf_counts, **metrics, + }) + print( + f"[result] latent={latent_length} prompt={prompt_id} " + f"psnr={metrics['psnr']:.4f} ssim={metrics['ssim']:.6f} " + f"lpips={metrics['lpips']:.6f}", flush=True, + ) + del prediction_latent, prediction_pixels, reference_u8, prediction_u8 + if hasattr(vae.model, "clear_cache"): + vae.model.clear_cache() + torch.cuda.empty_cache() + + manifest = json.loads((args.output_dir / "manifest.json").read_text(encoding="utf-8")) + manifest["status"] = "complete" + atomic_json(args.output_dir / "manifest.json", manifest) + + +if __name__ == "__main__": + main() diff --git a/scripts/evaluate_long_video_vbench.py b/scripts/evaluate_long_video_vbench.py new file mode 100644 index 0000000000000000000000000000000000000000..de3fc4313fd49b67df94c2dbdbe497e948971975 --- /dev/null +++ b/scripts/evaluate_long_video_vbench.py @@ -0,0 +1,53 @@ +#!/usr/bin/env python3 +"""Run all standard VBench dimensions for one long-video condition.""" + +from __future__ import annotations + +import argparse +import os +import tempfile +from pathlib import Path + + +def preparse_gpu() -> str: + parser = argparse.ArgumentParser(add_help=False) + parser.add_argument("--gpu", required=True) + args, _ = parser.parse_known_args() + os.environ["CUDA_VISIBLE_DEVICES"] = args.gpu + os.environ.setdefault("MPLCONFIGDIR", tempfile.mkdtemp(prefix="vbench_mpl_")) + return args.gpu + + +GPU = preparse_gpu() + +from vbench import VBench + + +DIMENSIONS = [ + "subject_consistency", "background_consistency", "motion_smoothness", + "aesthetic_quality", "imaging_quality", +] + + +def main() -> None: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--gpu", default=GPU) + parser.add_argument("--video_dir", type=Path, required=True) + parser.add_argument("--output_dir", type=Path, required=True) + parser.add_argument("--name", required=True) + args = parser.parse_args() + args.video_dir = args.video_dir.resolve() + args.output_dir = args.output_dir.resolve() + args.output_dir.mkdir(parents=True, exist_ok=True) + bench = VBench( + device="cuda", full_info_dir=str(args.video_dir / "full_info.json"), + output_path=str(args.output_dir), + ) + bench.evaluate( + videos_path=str(args.video_dir), name=args.name, + dimension_list=DIMENSIONS, mode="custom_input", + ) + + +if __name__ == "__main__": + main() diff --git a/scripts/evaluate_single_block_fppf.py b/scripts/evaluate_single_block_fppf.py new file mode 100644 index 0000000000000000000000000000000000000000..3111cc59d2a086092e89115d909d11ced11dc2ee --- /dev/null +++ b/scripts/evaluate_single_block_fppf.py @@ -0,0 +1,1135 @@ +#!/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() diff --git a/scripts/evaluate_single_chunk_frrr.py b/scripts/evaluate_single_chunk_frrr.py new file mode 100644 index 0000000000000000000000000000000000000000..d71a8dcebdd705afa85a9596ee203534f7a3424c --- /dev/null +++ b/scripts/evaluate_single_chunk_frrr.py @@ -0,0 +1,838 @@ +#!/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() diff --git a/scripts/evaluate_three_block_fppf.py b/scripts/evaluate_three_block_fppf.py new file mode 100644 index 0000000000000000000000000000000000000000..27bea14a5b6d0c5418de12ca92ce90e0ff9b1a28 --- /dev/null +++ b/scripts/evaluate_three_block_fppf.py @@ -0,0 +1,25 @@ +#!/usr/bin/env python3 +"""Evaluate consecutive three-block Predictors with the FPPF protocol.""" + +from __future__ import annotations + +import sys + +from evaluate_two_block_fppf import main + + +DEFAULTS = { + "--architecture": "three_block", + "--sweep_dir": "outputs/three_block_consecutive_sweep", + "--output_dir": "outputs/three_block_fppf_eval", + "--reference_root": "outputs/single_block_fppf_eval", + "--schedule": "FPPF", +} + + +if __name__ == "__main__": + present = set(sys.argv[1:]) + for option, value in DEFAULTS.items(): + if option not in present: + sys.argv.extend([option, value]) + main() diff --git a/scripts/evaluate_trained_long_predictors.py b/scripts/evaluate_trained_long_predictors.py new file mode 100644 index 0000000000000000000000000000000000000000..fcda150188cc9bcbcbdddfe6136eba38b89bd6d7 --- /dev/null +++ b/scripts/evaluate_trained_long_predictors.py @@ -0,0 +1,157 @@ +#!/usr/bin/env python3 +"""Evaluate 2x/4x-trained Layer-17 predictors at 1x, 2x, and 4x.""" + +from __future__ import annotations + +import argparse +import json +import os +import sys +from pathlib import Path + + +def preparse_gpu() -> str: + parser = argparse.ArgumentParser(add_help=False) + parser.add_argument("--gpu", required=True) + args, _ = parser.parse_known_args() + os.environ["CUDA_VISIBLE_DEVICES"] = args.gpu + return args.gpu + + +GPU = preparse_gpu() + +import lpips +import torch +from omegaconf import OmegaConf + +ROOT = Path(__file__).resolve().parents[1] +if str(ROOT) not in sys.path: + sys.path.insert(0, str(ROOT)) + +from scripts.evaluate_long_video_fppf import generate_rollout, save_mp4 +from scripts.evaluate_single_block_fppf import ( + atomic_json, build_pipeline, frame_metrics, load_predictor, + load_prompt_metadata, pixels_to_u8, +) +from utils.misc import set_seed +from utils.wan_wrapper import WanVAEWrapper + + +PREDICTORS = { + "trained_2x": ROOT / "outputs/layer17_long_training_four_gpu_v2/2x/predictor_final.safetensors", + "trained_4x": ROOT / "outputs/layer17_long_training_four_gpu_v2/4x/predictor_final.safetensors", +} + + +def main() -> None: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--gpu", default=GPU) + parser.add_argument("--prompt_ids", type=int, nargs="+", required=True) + parser.add_argument("--latent_lengths", type=int, nargs="+", default=[21, 42, 84]) + parser.add_argument( + "--dataset_root", type=Path, + default=Path("outputs/predictor_offline_100_all_blocks"), + ) + parser.add_argument( + "--output_dir", type=Path, + default=Path("outputs/layer17_long_training_eval"), + ) + parser.add_argument("--generation_seed", type=int, default=0) + parser.add_argument("--metric_batch_size", type=int, default=4) + args = parser.parse_args() + args.dataset_root = (ROOT / args.dataset_root).resolve() if not args.dataset_root.is_absolute() else args.dataset_root + args.output_dir = (ROOT / args.output_dir).resolve() if not args.output_dir.is_absolute() else args.output_dir + args.output_dir.mkdir(parents=True, exist_ok=True) + + device = torch.device("cuda") + set_seed(args.generation_seed) + config = OmegaConf.merge( + OmegaConf.load(ROOT / "configs/default_config.yaml"), + OmegaConf.load(ROOT / "configs/self_forcing_sid.yaml"), + ) + config.model_kwargs.local_attn_size = 21 + vae = WanVAEWrapper().to(device=device, dtype=torch.bfloat16).eval() + pipeline = build_pipeline( + config, ROOT / "checkpoints/self_forcing_dmd.pt", vae, device, + ) + predictors = { + name: load_predictor( + pipeline.generator.model, + {"source_layer": 17, "weights": path}, + device, + ) + for name, path in PREDICTORS.items() + } + lpips_model = lpips.LPIPS(net="alex", verbose=False).to(device).eval() + lpips_model.requires_grad_(False) + + for prompt_id in args.prompt_ids: + prompt = load_prompt_metadata(args.dataset_root, prompt_id)["prompt"] + for latent_length in args.latent_lengths: + run_dir = args.output_dir / f"latent_{latent_length}" / f"prompt_{prompt_id:04d}" + result_path = run_dir / "metrics.json" + if result_path.exists(): + existing = json.loads(result_path.read_text()) + if existing.get("status") == "complete": + print(f"[skip] prompt={prompt_id} latent={latent_length}", flush=True) + continue + + print(f"[run] prompt={prompt_id} latent={latent_length} FFFF", flush=True) + reference_latent, ffff_counts = generate_rollout( + pipeline=pipeline, dataset_root=args.dataset_root, + prompt_id=prompt_id, latent_length=latent_length, + generation_seed=args.generation_seed, device=device, + predictor=None, source_layer=None, schedule="FFFF", + ) + with torch.autocast(device_type="cuda", dtype=torch.bfloat16): + reference_pixels = vae.decode_to_pixel(reference_latent, use_cache=False) + reference_u8 = pixels_to_u8(reference_pixels) + save_mp4(reference_u8, run_dir / "ffff.mp4") + del reference_latent, reference_pixels + if hasattr(vae.model, "clear_cache"): + vae.model.clear_cache() + torch.cuda.empty_cache() + + results = {} + for name, predictor in predictors.items(): + print(f"[run] prompt={prompt_id} latent={latent_length} {name}", flush=True) + latent, counts = generate_rollout( + pipeline=pipeline, dataset_root=args.dataset_root, + prompt_id=prompt_id, latent_length=latent_length, + generation_seed=args.generation_seed, device=device, + predictor=predictor, source_layer=17, 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) + save_mp4(prediction_u8, run_dir / f"{name}.mp4") + metrics = frame_metrics( + reference_u8=reference_u8, + prediction_u8=prediction_u8, + lpips_model=lpips_model, + batch_size=args.metric_batch_size, + device=device, + ) + results[name] = {"fppf": counts, **metrics} + print( + f"[result] {name} prompt={prompt_id} latent={latent_length} " + f"psnr={metrics['psnr']:.4f} ssim={metrics['ssim']:.6f} " + f"lpips={metrics['lpips']:.6f}", flush=True, + ) + del latent, pixels, prediction_u8 + if hasattr(vae.model, "clear_cache"): + vae.model.clear_cache() + torch.cuda.empty_cache() + + atomic_json(result_path, { + "status": "complete", "prompt_id": prompt_id, + "prompt": prompt, "latent_length": latent_length, + "decoded_frames": next(iter(results.values()))["num_frames"], + "reference": "FFFF same prompt/seed/noise", + "ffff": ffff_counts, "predictors": results, + }) + del reference_u8 + + +if __name__ == "__main__": + main() diff --git a/scripts/evaluate_two_block_fppf.py b/scripts/evaluate_two_block_fppf.py new file mode 100644 index 0000000000000000000000000000000000000000..6480d134ed8d1a38b2a1432ce3e933be9b337de4 --- /dev/null +++ b/scripts/evaluate_two_block_fppf.py @@ -0,0 +1,640 @@ +#!/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() diff --git a/scripts/generate_ode_pairs.py b/scripts/generate_ode_pairs.py new file mode 100644 index 0000000000000000000000000000000000000000..22492ad4f38edfaab75f4438945b0a7f2cfc5c9a --- /dev/null +++ b/scripts/generate_ode_pairs.py @@ -0,0 +1,120 @@ +from utils.distributed import launch_distributed_job +from utils.scheduler import FlowMatchScheduler +from utils.wan_wrapper import WanDiffusionWrapper, WanTextEncoder +from utils.dataset import TextDataset +import torch.distributed as dist +from tqdm import tqdm +import argparse +import torch +import math +import os + + +def init_model(device): + model = WanDiffusionWrapper().to(device).to(torch.float32) + encoder = WanTextEncoder().to(device).to(torch.float32) + model.model.requires_grad_(False) + + scheduler = FlowMatchScheduler( + shift=8.0, sigma_min=0.0, extra_one_step=True) + scheduler.set_timesteps(num_inference_steps=48, denoising_strength=1.0) + scheduler.sigmas = scheduler.sigmas.to(device) + + sample_neg_prompt = '色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走' + + unconditional_dict = encoder( + text_prompts=[sample_neg_prompt] + ) + + return model, encoder, scheduler, unconditional_dict + + +def main(): + parser = argparse.ArgumentParser() + parser.add_argument("--local_rank", type=int, default=-1) + parser.add_argument("--output_folder", type=str) + parser.add_argument("--caption_path", type=str) + parser.add_argument("--guidance_scale", type=float, default=6.0) + + args = parser.parse_args() + + # launch_distributed_job() + launch_distributed_job() + + device = torch.cuda.current_device() + + torch.set_grad_enabled(False) + torch.backends.cuda.matmul.allow_tf32 = True + torch.backends.cudnn.allow_tf32 = True + + model, encoder, scheduler, unconditional_dict = init_model(device=device) + + dataset = TextDataset(args.caption_path) + + # if global_rank == 0: + os.makedirs(args.output_folder, exist_ok=True) + + for index in tqdm(range(int(math.ceil(len(dataset) / dist.get_world_size()))), disable=dist.get_rank() != 0): + prompt_index = index * dist.get_world_size() + dist.get_rank() + if prompt_index >= len(dataset): + continue + prompt = dataset[prompt_index] + + conditional_dict = encoder(text_prompts=prompt) + + latents = torch.randn( + [1, 21, 16, 60, 104], dtype=torch.float32, device=device + ) + + noisy_input = [] + + for progress_id, t in enumerate(tqdm(scheduler.timesteps)): + timestep = t * \ + torch.ones([1, 21], device=device, dtype=torch.float32) + + noisy_input.append(latents) + + _, x0_pred_cond = model( + latents, conditional_dict, timestep + ) + + _, x0_pred_uncond = model( + latents, unconditional_dict, timestep + ) + + x0_pred = x0_pred_uncond + args.guidance_scale * ( + x0_pred_cond - x0_pred_uncond + ) + + flow_pred = model._convert_x0_to_flow_pred( + scheduler=scheduler, + x0_pred=x0_pred.flatten(0, 1), + xt=latents.flatten(0, 1), + timestep=timestep.flatten(0, 1) + ).unflatten(0, x0_pred.shape[:2]) + + latents = scheduler.step( + flow_pred.flatten(0, 1), + scheduler.timesteps[progress_id] * torch.ones( + [1, 21], device=device, dtype=torch.long).flatten(0, 1), + latents.flatten(0, 1) + ).unflatten(dim=0, sizes=flow_pred.shape[:2]) + + noisy_input.append(latents) + + noisy_inputs = torch.stack(noisy_input, dim=1) + + noisy_inputs = noisy_inputs[:, [0, 12, 24, 36, -1]] + + stored_data = noisy_inputs + + torch.save( + {prompt: stored_data.cpu().detach()}, + os.path.join(args.output_folder, f"{prompt_index:05d}.pt") + ) + + dist.barrier() + + +if __name__ == "__main__": + main() diff --git a/scripts/generate_vbench8_extended_atc.py b/scripts/generate_vbench8_extended_atc.py new file mode 100644 index 0000000000000000000000000000000000000000..85851e2c9ad0ab0fbc25273adf9782e30999b41c --- /dev/null +++ b/scripts/generate_vbench8_extended_atc.py @@ -0,0 +1,292 @@ +#!/usr/bin/env python3 +"""Generate ATC FPPF/FPPP videos against the existing Extended-251 FFFF run.""" + +from __future__ import annotations + +import argparse +import json +import sys +from pathlib import Path +from typing import Any + +REPO_ROOT = Path(__file__).resolve().parents[1] +if str(REPO_ROOT) not in sys.path: + sys.path.insert(0, str(REPO_ROOT)) + +from scripts import generate_vbench8_extended_strategies as extended + +import lpips +import torch +from omegaconf import OmegaConf +from safetensors import safe_open + +from scripts import evaluate_single_block_fppf as base +from scripts.generate_vbench8_extended_disca import artifact_paths, read_u8_video +from utils.wan_wrapper import WanTextEncoder, WanVAEWrapper + + +MAPPING_DEFAULT = REPO_ROOT / "assets/vbench8_extended_subset_mapping.json" +REFERENCE_DEFAULT = ( + REPO_ROOT / "evaluation_runs/vbench8_extended_stage1_step2000_20260901" +) +OUTPUT_DEFAULT = ( + REPO_ROOT / "evaluation_runs/vbench8_extended_atc_stage1_step2000_20260901" +) +PATTERN_TO_STEPS = {"FPPF": [1, 2], "FPPP": [1, 2, 3]} + + +def read_predictor_config(weights: Path) -> dict[str, Any]: + with safe_open(weights, framework="pt", device="cpu") as handle: + metadata = handle.metadata() or {} + raw = metadata.get("predictor_config") + if raw is None: + raise ValueError(f"ATC checkpoint has no predictor_config metadata: {weights}") + config = json.loads(raw) + if config.get("input_variant") != "atc": + raise ValueError(f"Checkpoint is not ATC: {config}") + return config + + +def strategy_configs( + scope: str, patterns: list[str] | tuple[str, ...] +) -> list[dict[str, Any]]: + return [ + { + "name": f"atc_{scope}_{pattern.lower()}", + "policy": "fppf", + "pattern": pattern, + "candidate_steps": PATTERN_TO_STEPS[pattern], + "allow_chunk0_predictor": False, + "beta": None, + "threshold": None, + "target_accepts": 6 * len(PATTERN_TO_STEPS[pattern]), + "head": None, + "input_variant": "atc", + "atc_previous_scope": scope, + } + for pattern in patterns + ] + + +def load_models( + weights: Path, scope: str, device: torch.device +) -> tuple[Any, Any, Any, Any, None, None, Any]: + predictor_config = read_predictor_config(weights) + if predictor_config.get("atc_previous_scope") != scope: + raise ValueError( + f"Checkpoint scope={predictor_config.get('atc_previous_scope')} " + f"does not match requested scope={scope}" + ) + 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(REPO_ROOT / "configs/self_forcing_sid.yaml"), + ) + print(f"[setup] loading Teacher and ATC scope={scope}", flush=True) + pipeline = base.build_pipeline( + config, REPO_ROOT / "checkpoints/self_forcing_dmd.pt", vae, device + ) + predictor = base.load_predictor( + pipeline.generator.model, + { + **predictor_config, + "weights": weights, + "atc_collect_diagnostics": False, + }, + device, + ) + text_encoder = WanTextEncoder().to(device=device, dtype=torch.bfloat16).eval() + text_encoder.requires_grad_(False) + print("[setup] loading LPIPS", flush=True) + lpips_model = lpips.LPIPS(net="alex", verbose=False).to(device).eval() + lpips_model.requires_grad_(False) + return vae, pipeline, predictor, text_encoder, None, None, lpips_model + + +def is_complete( + output_root: Path, + row: dict[str, Any], + strategies: list[dict[str, Any]], +) -> bool: + return all( + all(path.is_file() for path in artifact_paths(output_root, row, config["name"])) + for config in strategies + ) + + +def run_prompt( + *, + row: dict[str, Any], + strategies: list[dict[str, Any]], + reference_root: Path, + output_root: Path, + seed: int, + models: tuple[Any, ...], + device: torch.device, +) -> None: + global_index = int(row["global_index"]) + suite = str(row["prompt_suite"]) + suite_index = int(row["suite_index"]) + prompt = str(row["extended_prompt"]) + reference_path = ( + reference_root / "generated_videos/ffff" / suite / f"{suite_index:03d}.mp4" + ) + print(f"[prompt] global={global_index} suite={suite}/{suite_index}", flush=True) + conditional = models[3](text_prompts=[prompt]) + for strategy in strategies: + name = str(strategy["name"]) + video_path, record_path = artifact_paths(output_root, row, name) + if video_path.is_file() and record_path.is_file(): + print(f"[cached-strategy] global={global_index} strategy={name}", flush=True) + continue + latent, diagnostic = extended.generate_rollout( + pipeline=models[1], + conditional_dict=conditional, + seed=seed, + device=device, + predictor=models[2], + head=None, + config=strategy, + ) + with torch.autocast(device_type="cuda", dtype=torch.bfloat16): + pixels = models[0].decode_to_pixel(latent, use_cache=False) + prediction_u8 = base.pixels_to_u8(pixels) + extended.atomic_video(prediction_u8, video_path) + reference_decoded = read_u8_video(reference_path) + prediction_decoded = read_u8_video(video_path) + metrics = base.frame_metrics( + reference_u8=reference_decoded, + prediction_u8=prediction_decoded, + lpips_model=models[6], + batch_size=4, + device=device, + ) + extended.atomic_json( + record_path, + { + "status": "complete", + "strategy": name, + "policy": "fppf", + "pattern": strategy["pattern"], + "input_variant": "atc", + "atc_previous_scope": strategy["atc_previous_scope"], + "candidate_steps": strategy["candidate_steps"], + "allow_chunk0_predictor": False, + "global_index": global_index, + "prompt_suite": suite, + "suite_index": suite_index, + "prompt": prompt, + "seed": seed, + "generation": { + key: value + for key, value in diagnostic.items() + if key != "decisions" + }, + "decisions": diagnostic["decisions"], + "pixel_metrics_vs_ffff": extended.compact_metrics(metrics), + "pixel_metric_input": ( + "MP4-decoded uint8 RGB, all 81 frames on both sides" + ), + "reference_video": str(reference_path), + "video": str(video_path.relative_to(output_root)), + }, + ) + print( + f"[result] global={global_index} strategy={name} " + f"accept={diagnostic['accepted_predictor_calls']} " + f"psnr={metrics['psnr']:.4f} ssim={metrics['ssim']:.6f} " + f"lpips={metrics['lpips']:.6f}", + flush=True, + ) + if hasattr(models[0].model, "clear_cache"): + models[0].model.clear_cache() + del latent, pixels, prediction_u8, reference_decoded, prediction_decoded + torch.cuda.empty_cache() + del conditional + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--gpu", default=extended.PHYSICAL_GPU) + parser.add_argument("--scope", choices=("chunk", "last_frame"), required=True) + parser.add_argument("--predictor-weights", type=Path, required=True) + parser.add_argument("--mapping", type=Path, default=MAPPING_DEFAULT) + parser.add_argument("--reference-root", type=Path, default=REFERENCE_DEFAULT) + parser.add_argument("--output-root", type=Path, default=OUTPUT_DEFAULT) + parser.add_argument("--pattern", action="append", choices=tuple(PATTERN_TO_STEPS)) + parser.add_argument("--shard-index", type=int, default=0) + parser.add_argument("--num-shards", type=int, default=1) + parser.add_argument("--global-index", action="append", type=int, default=None) + parser.add_argument("--seed", type=int, default=0) + return parser.parse_args() + + +def main() -> None: + args = parse_args() + if args.num_shards < 1 or not 0 <= args.shard_index < args.num_shards: + raise ValueError("Invalid shard index/number") + weights = args.predictor_weights.resolve() + mapping = extended.read_mapping(args.mapping.resolve()) + if args.global_index is None: + rows = [ + row for index, row in enumerate(mapping) + if index % args.num_shards == args.shard_index + ] + else: + requested = set(args.global_index) + rows = [row for row in mapping if int(row["global_index"]) in requested] + if len(rows) != len(requested): + raise ValueError("Unknown global index") + patterns = args.pattern or ["FPPF", "FPPP"] + strategies = strategy_configs(args.scope, patterns) + reference_root = args.reference_root.resolve() + if len(list((reference_root / "generated_videos/ffff").rglob("*.mp4"))) != 251: + raise ValueError("Incomplete FFFF reference") + output_root = args.output_root.resolve() + output_root.mkdir(parents=True, exist_ok=True) + extended.write_shard_manifest( + output_root, + str(args.gpu), + args.shard_index, + args.num_shards, + rows, + strategies, + "running", + ) + device = torch.device("cuda") + torch.set_grad_enabled(False) + models = load_models(weights, args.scope, device) + completed = 0 + try: + for row in rows: + if is_complete(output_root, row, strategies): + completed += 1 + print(f"[cached] {completed}/{len(rows)} global={row['global_index']}", flush=True) + continue + run_prompt( + row=row, + strategies=strategies, + reference_root=reference_root, + output_root=output_root, + seed=args.seed, + models=models, + device=device, + ) + completed += 1 + print(f"[progress] {completed}/{len(rows)} prompts", flush=True) + finally: + extended.write_shard_manifest( + output_root, + str(args.gpu), + args.shard_index, + args.num_shards, + rows, + strategies, + "complete" if completed == len(rows) else "failed", + ) + print(f"[complete] gpu={args.gpu} scope={args.scope} prompts={completed}", flush=True) + + +if __name__ == "__main__": + main() diff --git a/scripts/generate_vbench8_extended_atc_confidence_token.py b/scripts/generate_vbench8_extended_atc_confidence_token.py new file mode 100644 index 0000000000000000000000000000000000000000..7793d32edc9ba47306f69c8db7d256c0a18625f6 --- /dev/null +++ b/scripts/generate_vbench8_extended_atc_confidence_token.py @@ -0,0 +1,439 @@ +#!/usr/bin/env python3 +"""Generate one Confidence-token strategy for Extended-251.""" + +from __future__ import annotations + +import argparse +import json +import os +import sys +import tempfile +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) + os.environ.setdefault( + "MPLCONFIGDIR", tempfile.mkdtemp(prefix="atc_confidence_token_vbench_") + ) + return str(args.gpu) + + +PHYSICAL_GPU = preparse_gpu() + +import lpips +import torch +from omegaconf import OmegaConf +from safetensors import safe_open +from safetensors.torch import load_file + +ROOT = Path(__file__).resolve().parents[1] +if str(ROOT) not in sys.path: + sys.path.insert(0, str(ROOT)) + +from predictor_training.confidence import ConfidenceTokenHead +from scripts import evaluate_single_block_fppf as base +from scripts import generate_vbench8_extended_strategies as extended +from scripts.generate_vbench8_extended_disca import artifact_paths, read_u8_video +from utils.wan_wrapper import WanTextEncoder, WanVAEWrapper + + +MAPPING_DEFAULT = ROOT / "assets/vbench8_extended_subset_mapping.json" +REFERENCE_DEFAULT = ROOT / "evaluation_runs/vbench8_extended_stage1_step2000_20260901" +SELECTION_DEFAULT = ( + ROOT + / "confidence_experiments" + / "layer17_atc_chunk_confidence_token_beta2_vbench3_20260902" + / "threshold_summary.json" +) +PREDICTOR_DEFAULT = ( + ROOT + / "training_runs/layer17_atc_chunk_stage1_1000p_4gpu_b16_2000steps" + / "checkpoint_step_2000/predictor.safetensors" +) +HEAD_DEFAULT = ( + ROOT + / "training_runs" + / "layer17_atc_chunk_confidence_token_900train_100val_4gpu_b64_warmup_cosine" + / "confidence_latest.safetensors" +) +OUTPUT_DEFAULT = ( + ROOT + / "evaluation_runs" + / "vbench8_extended_atc_chunk_confidence_token_beta2_20260902" +) +TARGETS = (6, 9, 12, 15) + + +def read_head_config(path: Path) -> dict[str, Any]: + with safe_open(path, framework="pt", device="cpu") as handle: + metadata = handle.metadata() or {} + raw = metadata.get("head_config") + if raw is None: + raise ValueError(f"Missing head_config metadata: {path}") + config = json.loads(raw) + expected = { + "architecture": "ConfidenceTokenHead", + "dim": 1536, + "token_dim": 512, + "context_dim": 64, + "num_heads": 8, + "ffn_dim": 2048, + "num_steps": 3, + } + for key, value in expected.items(): + if config.get(key) != value: + raise ValueError(f"Unexpected head config {key}={config.get(key)}") + return config + + +def read_predictor_config( + path: Path, input_variant: str, atc_previous_scope: str +) -> dict[str, Any]: + with safe_open(path, framework="pt", device="cpu") as handle: + metadata = handle.metadata() or {} + raw = metadata.get("predictor_config") + if raw is None: + if input_variant != "self_forcing": + raise ValueError( + f"Missing predictor_config metadata: {path}; only an explicitly " + "selected legacy self_forcing/concat checkpoint may omit it" + ) + return { + "source_layer": 17, + "input_variant": "self_forcing", + "gate_mode": "baseline", + } + config = json.loads(raw) + actual_variant = str(config.get("input_variant", "self_forcing")) + if actual_variant != input_variant: + raise ValueError( + f"Predictor variant mismatch: expected={input_variant} " + f"actual={actual_variant}" + ) + if input_variant == "atc": + actual_scope = str(config.get("atc_previous_scope", "chunk")) + if actual_scope != atc_previous_scope: + raise ValueError( + f"ATC scope mismatch: expected={atc_previous_scope} " + f"actual={actual_scope}" + ) + return config + + +def strategy_prefix(input_variant: str, atc_previous_scope: str) -> str: + if input_variant == "self_forcing": + return "concat" + return f"atc_{atc_previous_scope}" + + +def load_strategy( + path: Path, + target_k: int, + input_variant: str, + atc_previous_scope: str, +) -> dict[str, Any]: + payload = json.loads(path.read_text(encoding="utf-8")) + if payload.get("status") != "complete" or float(payload.get("beta")) != 2.0: + raise ValueError(f"Selection is not complete fixed-beta=2: {path}") + if [int(value) for value in payload.get("candidate_steps", [])] != [1, 2, 3]: + raise ValueError("Selection is not Step123") + if payload.get("predictor_input_variant") != input_variant: + raise ValueError( + "Threshold selection Predictor variant does not match evaluation" + ) + expected_scope = atc_previous_scope if input_variant == "atc" else None + if payload.get("atc_previous_scope") != expected_scope: + raise ValueError("Threshold selection ATC scope does not match evaluation") + matches = [ + row for row in payload.get("selected", []) + if int(row["target_accepts"]) == target_k + ] + if len(matches) != 1: + raise ValueError(f"Expected one selected K{target_k:02d}, got {len(matches)}") + row = matches[0] + return { + "name": ( + f"{strategy_prefix(input_variant, atc_previous_scope)}" + f"_confidence_token_beta2_k{target_k:02d}" + ), + "policy": "dynamic", + "candidate_steps": [1, 2, 3], + "beta": 2.0, + "threshold": float(row["threshold"]), + "target_accepts": target_k, + "head": "confidence_token", + "risk_mode": "outgoing_span", + "allow_chunk0_predictor": False, + "source_config_name": str(row["name"]), + "calibration_mean_accepted_predictor_calls": float( + row["mean_accepted_predictor_calls"] + ), + } + + +def load_models( + predictor_weights: Path, + head_weights: Path, + input_variant: str, + atc_previous_scope: str, + device: torch.device, +) -> tuple[Any, Any, Any, ConfidenceTokenHead, Any]: + predictor_config = read_predictor_config( + predictor_weights, input_variant, atc_previous_scope + ) + + print("[setup] loading VAE", flush=True) + vae = WanVAEWrapper().to(device=device, dtype=torch.bfloat16).eval() + config = OmegaConf.merge( + OmegaConf.load(ROOT / "configs/default_config.yaml"), + OmegaConf.load(ROOT / "configs/self_forcing_sid.yaml"), + ) + print( + f"[setup] loading frozen Teacher and {input_variant} " + f"scope={atc_previous_scope if input_variant == 'atc' else 'n/a'} Predictor", + flush=True, + ) + pipeline = base.build_pipeline( + config, ROOT / "checkpoints/self_forcing_dmd.pt", vae, device + ) + predictor = base.load_predictor( + pipeline.generator.model, + { + **predictor_config, + "weights": predictor_weights, + "atc_collect_diagnostics": False, + }, + device, + ) + + print("[setup] loading text encoder and Confidence-token Head", flush=True) + text_encoder = WanTextEncoder().to(device=device, dtype=torch.bfloat16).eval() + text_encoder.requires_grad_(False) + head_config = read_head_config(head_weights) + head = ConfidenceTokenHead( + dim=int(head_config["dim"]), + token_dim=int(head_config["token_dim"]), + context_dim=int(head_config["context_dim"]), + num_heads=int(head_config["num_heads"]), + ffn_dim=int(head_config["ffn_dim"]), + dropout=float(head_config.get("dropout", 0.1)), + num_steps=int(head_config["num_steps"]), + ).to(device=device).eval() + head.load_state_dict(load_file(str(head_weights), device="cpu"), strict=True) + head.requires_grad_(False) + + print("[setup] loading LPIPS", flush=True) + lpips_model = lpips.LPIPS(net="alex", verbose=False).to(device).eval() + lpips_model.requires_grad_(False) + return vae, pipeline, predictor, text_encoder, head, lpips_model + + +def run_prompt( + *, + row: dict[str, Any], + strategy: dict[str, Any], + reference_root: Path, + output_root: Path, + seed: int, + models: tuple[Any, ...], + device: torch.device, +) -> None: + global_index = int(row["global_index"]) + suite = str(row["prompt_suite"]) + suite_index = int(row["suite_index"]) + prompt = str(row["extended_prompt"]) + name = str(strategy["name"]) + reference_path = ( + reference_root / "generated_videos/ffff" / suite / f"{suite_index:03d}.mp4" + ) + print(f"[prompt] global={global_index} suite={suite}/{suite_index}", flush=True) + conditional = models[3](text_prompts=[prompt]) + latent, diagnostic = extended.generate_rollout( + pipeline=models[1], + conditional_dict=conditional, + seed=seed, + device=device, + predictor=models[2], + head=models[4], + config=strategy, + ) + with torch.autocast(device_type="cuda", dtype=torch.bfloat16): + pixels = models[0].decode_to_pixel(latent, use_cache=False) + prediction_u8 = base.pixels_to_u8(pixels) + video_path, record_path = artifact_paths(output_root, row, name) + extended.atomic_video(prediction_u8, video_path) + + reference_decoded = read_u8_video(reference_path) + prediction_decoded = read_u8_video(video_path) + metrics = base.frame_metrics( + reference_u8=reference_decoded, + prediction_u8=prediction_decoded, + lpips_model=models[5], + batch_size=4, + device=device, + ) + extended.atomic_json( + record_path, + { + "status": "complete", + "strategy": name, + "policy": "dynamic", + "input_variant": strategy["input_variant"], + "atc_previous_scope": strategy["atc_previous_scope"], + "candidate_steps": [1, 2, 3], + "beta": 2.0, + "threshold": float(strategy["threshold"]), + "target_accepts": int(strategy["target_accepts"]), + "risk_mode": "outgoing_span", + "allow_chunk0_predictor": False, + "source_config_name": strategy["source_config_name"], + "global_index": global_index, + "prompt_suite": suite, + "suite_index": suite_index, + "prompt": prompt, + "seed": seed, + "generation": { + key: value for key, value in diagnostic.items() if key != "decisions" + }, + "decisions": diagnostic["decisions"], + "pixel_metrics_vs_ffff": extended.compact_metrics(metrics), + "pixel_metric_input": "MP4-decoded uint8 RGB, all 81 frames on both sides", + "reference_video": str(reference_path), + "video": str(video_path.relative_to(output_root)), + }, + ) + print( + f"[result] global={global_index} strategy={name} " + f"accept={diagnostic['accepted_predictor_calls']} " + f"full={diagnostic['full_calls']} " + f"latency={diagnostic['policy_latency_ms']:.2f}ms " + f"psnr={metrics['psnr']:.4f} ssim={metrics['ssim']:.6f} " + f"lpips={metrics['lpips']:.6f}", + flush=True, + ) + if hasattr(models[0].model, "clear_cache"): + models[0].model.clear_cache() + del conditional, latent, pixels, prediction_u8, reference_decoded, prediction_decoded + torch.cuda.empty_cache() + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--gpu", default=PHYSICAL_GPU) + parser.add_argument("--target-k", type=int, choices=TARGETS, required=True) + parser.add_argument("--mapping", type=Path, default=MAPPING_DEFAULT) + parser.add_argument("--selection", type=Path, default=SELECTION_DEFAULT) + parser.add_argument("--predictor-weights", type=Path, default=PREDICTOR_DEFAULT) + parser.add_argument("--head-weights", type=Path, default=HEAD_DEFAULT) + parser.add_argument( + "--predictor-input-variant", + choices=("self_forcing", "atc"), + default="atc", + ) + parser.add_argument( + "--atc-previous-scope", + choices=("chunk", "last_frame"), + default="chunk", + ) + parser.add_argument("--reference-root", type=Path, default=REFERENCE_DEFAULT) + parser.add_argument("--output-root", type=Path, default=OUTPUT_DEFAULT) + parser.add_argument("--global-index", action="append", type=int, default=None) + parser.add_argument("--seed", type=int, default=0) + parser.add_argument("--overwrite", action="store_true") + args = parser.parse_args() + for name in ( + "mapping", "selection", "predictor_weights", "head_weights", + "reference_root", "output_root", + ): + setattr(args, name, getattr(args, name).expanduser().resolve()) + return args + + +def main() -> None: + args = parse_args() + strategy = load_strategy( + args.selection, + args.target_k, + args.predictor_input_variant, + args.atc_previous_scope, + ) + strategy["input_variant"] = args.predictor_input_variant + strategy["atc_previous_scope"] = ( + args.atc_previous_scope + if args.predictor_input_variant == "atc" + else None + ) + mapping = extended.read_mapping(args.mapping) + if args.global_index is not None: + requested = set(args.global_index) + mapping = [row for row in mapping if int(row["global_index"]) in requested] + if len(mapping) != len(requested): + raise ValueError("At least one requested global index is unknown") + if len(list((args.reference_root / "generated_videos/ffff").rglob("*.mp4"))) != 251: + raise ValueError("Incomplete FFFF reference videos") + if len(list((args.reference_root / "generation_metrics/per_prompt/ffff").glob("global_*.json"))) != 251: + raise ValueError("Incomplete FFFF reference records") + + args.output_root.mkdir(parents=True, exist_ok=True) + extended.write_shard_manifest( + args.output_root, PHYSICAL_GPU, 0, 1, mapping, [strategy], "running" + ) + print( + f"[setup] physical_gpu={PHYSICAL_GPU} strategy={strategy['name']} " + f"threshold={strategy['threshold']:.10f} prompts={len(mapping)} " + "candidate_steps=1,2,3 chunk0=FFFF", + flush=True, + ) + device = torch.device("cuda") + torch.set_grad_enabled(False) + models = load_models( + args.predictor_weights, + args.head_weights, + args.predictor_input_variant, + args.atc_previous_scope, + device, + ) + completed = 0 + try: + for row in mapping: + video_path, record_path = artifact_paths( + args.output_root, row, strategy["name"] + ) + if not args.overwrite and video_path.is_file() and record_path.is_file(): + completed += 1 + print(f"[cached] {completed}/{len(mapping)} global={row['global_index']}", flush=True) + continue + run_prompt( + row=row, + strategy=strategy, + reference_root=args.reference_root, + output_root=args.output_root, + seed=args.seed, + models=models, + device=device, + ) + completed += 1 + print(f"[progress] {completed}/{len(mapping)} prompts", flush=True) + finally: + extended.write_shard_manifest( + args.output_root, + PHYSICAL_GPU, + 0, + 1, + mapping, + [strategy], + "complete" if completed == len(mapping) else "failed", + ) + print( + f"[complete] gpu={PHYSICAL_GPU} strategy={strategy['name']} prompts={completed}", + flush=True, + ) + + +if __name__ == "__main__": + main() diff --git a/scripts/generate_vbench8_extended_concat.py b/scripts/generate_vbench8_extended_concat.py new file mode 100644 index 0000000000000000000000000000000000000000..1326a74a0e00b11f2b20869c83471e2d442a5d11 --- /dev/null +++ b/scripts/generate_vbench8_extended_concat.py @@ -0,0 +1,96 @@ +#!/usr/bin/env python3 +"""Generate original-concat FPPP videos against the existing FFFF reference.""" + +from __future__ import annotations + +import sys +from pathlib import Path +from typing import Any + +REPO_ROOT = Path(__file__).resolve().parents[1] +if str(REPO_ROOT) not in sys.path: + sys.path.insert(0, str(REPO_ROOT)) + +from scripts import evaluate_single_block_fppf as base +from scripts import generate_vbench8_extended_disca as runner + +import lpips +import torch +from omegaconf import OmegaConf + +from utils.wan_wrapper import WanTextEncoder, WanVAEWrapper + + +WEIGHTS_DEFAULT = ( + REPO_ROOT + / "training_runs/layer17_stage1_1000p_4gpu_b16_2000steps" + / "checkpoint_step_2000/predictor.safetensors" +) +OUTPUT_DEFAULT = ( + REPO_ROOT + / "evaluation_runs/vbench8_extended_concat_fppp_stage1_step2000_20260902" +) + + +def strategy_configs(patterns: list[str] | tuple[str, ...]) -> list[dict[str, Any]]: + if patterns != ["FPPP"] and tuple(patterns) != ("FPPP",): + raise ValueError("Original-concat evaluation supports exactly --pattern FPPP") + return [ + { + "name": "concat_fppp", + "policy": "fppf", + "pattern": "FPPP", + "candidate_steps": [1, 2, 3], + "allow_chunk0_predictor": False, + "beta": None, + "threshold": None, + "target_accepts": 18, + "head": None, + "input_variant": "self_forcing", + } + ] + + +def load_models( + weights: Path, device: torch.device +) -> tuple[Any, Any, Any, Any, None, None, 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(REPO_ROOT / "configs/self_forcing_sid.yaml"), + ) + print("[setup] loading frozen Teacher and original concat Predictor", flush=True) + pipeline = base.build_pipeline( + config, REPO_ROOT / "checkpoints/self_forcing_dmd.pt", vae, device + ) + predictor = base.load_predictor( + pipeline.generator.model, + { + "source_layer": 17, + "weights": weights, + "input_variant": "self_forcing", + "gate_mode": "baseline", + }, + device, + ) + if predictor.input_variant != "self_forcing": + raise RuntimeError("Loaded Predictor is not the original concat variant") + text_encoder = WanTextEncoder().to(device=device, dtype=torch.bfloat16).eval() + text_encoder.requires_grad_(False) + print("[setup] loading LPIPS", flush=True) + lpips_model = lpips.LPIPS(net="alex", verbose=False).to(device).eval() + lpips_model.requires_grad_(False) + return vae, pipeline, predictor, text_encoder, None, None, lpips_model + + +def main() -> None: + runner.WEIGHTS_DEFAULT = WEIGHTS_DEFAULT + runner.OUTPUT_DEFAULT = OUTPUT_DEFAULT + runner.strategy_configs = strategy_configs + runner.load_models = load_models + runner.main() + + +if __name__ == "__main__": + main() diff --git a/scripts/generate_vbench8_extended_disca.py b/scripts/generate_vbench8_extended_disca.py new file mode 100644 index 0000000000000000000000000000000000000000..5529aa423c25f9e81345313e391fd02cd275fea6 --- /dev/null +++ b/scripts/generate_vbench8_extended_disca.py @@ -0,0 +1,374 @@ +#!/usr/bin/env python3 +"""Generate DISCA pattern videos using an existing Extended-251 FFFF reference.""" + +from __future__ import annotations + +import argparse +import sys +from pathlib import Path +from typing import Any + +REPO_ROOT = Path(__file__).resolve().parents[1] +if str(REPO_ROOT) not in sys.path: + sys.path.insert(0, str(REPO_ROOT)) + +from scripts import generate_vbench8_extended_strategies as extended + +import lpips +import torch +from omegaconf import OmegaConf +from torchvision.io import read_video + +from scripts import evaluate_single_block_fppf as base +from utils.wan_wrapper import WanTextEncoder, WanVAEWrapper + + +MAPPING_DEFAULT = REPO_ROOT / "assets/vbench8_extended_subset_mapping.json" +WEIGHTS_DEFAULT = ( + REPO_ROOT + / "training_runs/disca_layer17_stage1_1000p_4gpu_b16_2000steps" + / "checkpoint_step_2000/predictor.safetensors" +) +OUTPUT_DEFAULT = ( + REPO_ROOT + / "evaluation_runs/vbench8_extended_disca_chunk0_fppf_stage1_step2000" +) +REFERENCE_DEFAULT = ( + REPO_ROOT / "evaluation_runs/vbench8_extended_stage1_step2000_20260901" +) +PATTERN_TO_CANDIDATE_STEPS = { + "FPFF": [1], + "FPPF": [1, 2], + "FPPP": [1, 2, 3], +} + + +def strategy_configs(patterns: list[str] | tuple[str, ...]) -> list[dict[str, Any]]: + configs: list[dict[str, Any]] = [] + for pattern in patterns: + candidate_steps = PATTERN_TO_CANDIDATE_STEPS[pattern] + configs.append( + { + "name": f"disca_{pattern.lower()}", + "policy": "fppf", + "pattern": pattern, + "candidate_steps": candidate_steps, + "allow_chunk0_predictor": True, + "beta": None, + "threshold": None, + "target_accepts": 7 * len(candidate_steps), + "head": None, + "input_variant": "disca", + } + ) + return configs + + +def load_models( + weights: Path, device: torch.device +) -> tuple[Any, Any, Any, Any, None, None, 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(REPO_ROOT / "configs/self_forcing_sid.yaml"), + ) + print("[setup] loading frozen Teacher and DISCA Predictor", flush=True) + pipeline = base.build_pipeline( + config, REPO_ROOT / "checkpoints/self_forcing_dmd.pt", vae, device + ) + predictor = base.load_predictor( + pipeline.generator.model, + { + "source_layer": 17, + "weights": weights, + "input_variant": "disca", + "gate_mode": "baseline", + }, + device, + ) + if predictor.input_variant != "disca": + raise RuntimeError("Loaded Predictor is not the DISCA input variant") + text_encoder = WanTextEncoder().to(device=device, dtype=torch.bfloat16).eval() + text_encoder.requires_grad_(False) + print("[setup] loading LPIPS", flush=True) + lpips_model = lpips.LPIPS(net="alex", verbose=False).to(device).eval() + lpips_model.requires_grad_(False) + return vae, pipeline, predictor, text_encoder, None, None, lpips_model + + +def read_u8_video(path: Path) -> torch.Tensor: + if not path.is_file(): + raise FileNotFoundError(path) + frames, _, _ = read_video(str(path), pts_unit="sec", output_format="TCHW") + frames = frames.to(device="cpu", dtype=torch.uint8).contiguous() + if frames.ndim != 4 or frames.shape[0] != 81 or frames.shape[1] != 3: + raise RuntimeError(f"Unexpected decoded video shape {tuple(frames.shape)}: {path}") + return frames + + +def artifact_paths( + output_root: Path, mapping_row: dict[str, Any], strategy_name: str +) -> tuple[Path, Path]: + suite = str(mapping_row["prompt_suite"]) + suite_index = int(mapping_row["suite_index"]) + global_index = int(mapping_row["global_index"]) + return ( + output_root + / "generated_videos" + / strategy_name + / suite + / f"{suite_index:03d}.mp4", + output_root + / "generation_metrics/per_prompt" + / strategy_name + / f"global_{global_index:04d}.json", + ) + + +def is_prompt_complete( + output_root: Path, + mapping_row: dict[str, Any], + strategies: list[dict[str, Any]], +) -> bool: + return all( + all( + path.is_file() + for path in artifact_paths( + output_root, mapping_row, str(strategy["name"]) + ) + ) + for strategy in strategies + ) + + +def run_prompt( + *, + mapping_row: dict[str, Any], + strategy: dict[str, Any], + reference_root: Path, + output_root: Path, + seed: int, + vae: Any, + pipeline: Any, + predictor: Any, + text_encoder: Any, + lpips_model: Any, + device: torch.device, +) -> None: + global_index = int(mapping_row["global_index"]) + suite = str(mapping_row["prompt_suite"]) + suite_index = int(mapping_row["suite_index"]) + prompt = str(mapping_row["extended_prompt"]) + strategy_name = str(strategy["name"]) + reference_path = ( + reference_root + / "generated_videos/ffff" + / suite + / f"{suite_index:03d}.mp4" + ) + print(f"[prompt] global={global_index} suite={suite}/{suite_index}", flush=True) + conditional = text_encoder(text_prompts=[prompt]) + latent, diagnostic = extended.generate_rollout( + pipeline=pipeline, + conditional_dict=conditional, + seed=seed, + device=device, + predictor=predictor, + head=None, + config=strategy, + ) + 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) + video_path, record_path = artifact_paths(output_root, mapping_row, strategy_name) + extended.atomic_video(prediction_u8, video_path) + + # The original run retained FFFF MP4s, not its pre-encoding uint8 frames. + # Decode both sides so the reused reference and DISCA prediction receive + # the same H.264/CRF-18 treatment before full-81-frame pixel metrics. + reference_decoded = read_u8_video(reference_path) + prediction_decoded = read_u8_video(video_path) + metrics = base.frame_metrics( + reference_u8=reference_decoded, + prediction_u8=prediction_decoded, + lpips_model=lpips_model, + batch_size=4, + device=device, + ) + record = { + "status": "complete", + "strategy": strategy_name, + "policy": "fppf", + "pattern": strategy["pattern"], + "input_variant": strategy["input_variant"], + "candidate_steps": strategy["candidate_steps"], + "allow_chunk0_predictor": strategy["allow_chunk0_predictor"], + "global_index": global_index, + "prompt_suite": suite, + "suite_index": suite_index, + "prompt": prompt, + "seed": seed, + "generation": { + key: value + for key, value in diagnostic.items() + if key != "decisions" + }, + "decisions": diagnostic["decisions"], + "pixel_metrics_vs_ffff": extended.compact_metrics(metrics), + "pixel_metric_input": "MP4-decoded uint8 RGB, all 81 frames on both sides", + "reference_video": str(reference_path), + "video": str(video_path.relative_to(output_root)), + } + extended.atomic_json(record_path, record) + print( + f"[result] global={global_index} strategy={strategy_name} " + f"accept={diagnostic['accepted_predictor_calls']} " + f"psnr={metrics['psnr']:.4f} ssim={metrics['ssim']:.6f} " + f"lpips={metrics['lpips']:.6f}", + flush=True, + ) + if hasattr(vae.model, "clear_cache"): + vae.model.clear_cache() + del ( + conditional, + latent, + pixels, + prediction_u8, + reference_decoded, + prediction_decoded, + metrics, + ) + torch.cuda.empty_cache() + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--gpu", default=extended.PHYSICAL_GPU) + parser.add_argument("--mapping", type=Path, default=MAPPING_DEFAULT) + parser.add_argument("--predictor-weights", type=Path, default=WEIGHTS_DEFAULT) + parser.add_argument("--output-root", type=Path, default=OUTPUT_DEFAULT) + parser.add_argument("--reference-root", type=Path, default=REFERENCE_DEFAULT) + parser.add_argument( + "--pattern", + action="append", + choices=tuple(PATTERN_TO_CANDIDATE_STEPS), + default=None, + help="DISCA denoise pattern; repeat to generate several. Default: FPPF.", + ) + parser.add_argument("--shard-index", type=int, default=0) + parser.add_argument("--num-shards", type=int, default=1) + parser.add_argument( + "--global-index", + dest="global_indices", + action="append", + type=int, + default=None, + help="Process only this global mapping index; may be repeated.", + ) + parser.add_argument("--seed", type=int, default=0) + parser.add_argument("--overwrite", action="store_true") + return parser.parse_args() + + +def main() -> None: + args = parse_args() + if args.num_shards < 1 or not 0 <= args.shard_index < args.num_shards: + raise ValueError("Invalid shard index/number of shards") + weights = args.predictor_weights.resolve() + if not weights.is_file(): + raise FileNotFoundError(weights) + mapping = extended.read_mapping(args.mapping.resolve()) + patterns = args.pattern or ["FPPF"] + if len(set(patterns)) != len(patterns): + raise ValueError("--pattern values must be unique") + strategies = strategy_configs(patterns) + reference_root = args.reference_root.resolve() + reference_records = reference_root / "generation_metrics/per_prompt/ffff" + reference_videos = reference_root / "generated_videos/ffff" + if len(list(reference_records.glob("global_*.json"))) != 251: + raise ValueError(f"Incomplete FFFF reference records: {reference_records}") + if len(list(reference_videos.rglob("*.mp4"))) != 251: + raise ValueError(f"Incomplete FFFF reference videos: {reference_videos}") + if args.global_indices is None: + rows = [ + row + for position, row in enumerate(mapping) + if position % args.num_shards == args.shard_index + ] + else: + requested = set(args.global_indices) + mapping_by_global = {int(row["global_index"]): row for row in mapping} + unknown = sorted(requested - set(mapping_by_global)) + if unknown: + raise ValueError(f"Unknown global indices: {unknown}") + rows = [row for row in mapping if int(row["global_index"]) in requested] + + output_root = args.output_root.resolve() + output_root.mkdir(parents=True, exist_ok=True) + extended.write_shard_manifest( + output_root, + str(args.gpu), + args.shard_index, + args.num_shards, + rows, + strategies, + "running", + ) + print( + f"[setup] physical_gpu={args.gpu} shard={args.shard_index}/{args.num_shards} " + f"prompts={len(rows)} strategies={','.join(strategy['name'] for strategy in strategies)} " + "reference=existing_ffff", + flush=True, + ) + device = torch.device("cuda") + torch.set_grad_enabled(False) + models = load_models(weights, device) + completed = 0 + try: + for row in rows: + if not args.overwrite and is_prompt_complete( + output_root, row, strategies + ): + completed += 1 + print( + f"[cached] {completed}/{len(rows)} global={row['global_index']}", + flush=True, + ) + continue + for strategy in strategies: + video, record = artifact_paths( + output_root, row, str(strategy["name"]) + ) + if not args.overwrite and video.is_file() and record.is_file(): + continue + run_prompt( + mapping_row=row, + strategy=strategy, + reference_root=reference_root, + output_root=output_root, + seed=args.seed, + vae=models[0], + pipeline=models[1], + predictor=models[2], + text_encoder=models[3], + lpips_model=models[6], + device=device, + ) + completed += 1 + print(f"[progress] {completed}/{len(rows)} prompts", flush=True) + finally: + extended.write_shard_manifest( + output_root, + str(args.gpu), + args.shard_index, + args.num_shards, + rows, + strategies, + "complete" if completed == len(rows) else "failed", + ) + print(f"[complete] gpu={args.gpu} prompts={completed}/{len(rows)}", flush=True) + + +if __name__ == "__main__": + main() diff --git a/scripts/generate_vbench8_extended_naive_baselines.py b/scripts/generate_vbench8_extended_naive_baselines.py new file mode 100644 index 0000000000000000000000000000000000000000..ba17e5f7a71e09989a6f289d9f841896996c2df7 --- /dev/null +++ b/scripts/generate_vbench8_extended_naive_baselines.py @@ -0,0 +1,630 @@ +#!/usr/bin/env python3 +"""Generate Extended-251 videos for timestep-skipping and velocity-cache baselines. + +For every prompt, an original four-step FFFF rollout is generated first with +the same seed and used as the PSNR/SSIM/LPIPS reference. Timestep-skipping +policies shorten the denoising schedule itself. Velocity-cache policies retain +the original four timesteps and, at each R, reuse the velocity from the most +recent F while converting it to x0 using the current noisy latent and timestep. +The cache pattern is applied independently to every autoregressive chunk. +""" + +from __future__ import annotations + +import argparse +import hashlib +import json +import os +import sys +import tempfile +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) + os.environ.setdefault("MPLCONFIGDIR", tempfile.mkdtemp(prefix="self_forcing_mpl_")) + return str(args.gpu) + + +PHYSICAL_GPU = preparse_gpu() + +import lpips +import torch +from omegaconf import OmegaConf + +REPO_ROOT = Path(__file__).resolve().parents[1] +if str(REPO_ROOT) not in sys.path: + sys.path.insert(0, str(REPO_ROOT)) + +from scripts import evaluate_single_block_fppf as base +from scripts.naive_vbench_policies import ( + EVALUATION_STRATEGY_NAMES, + POLICIES_BY_NAME, + REFERENCE_POLICY, + NaivePolicy, +) +from scripts.vbench8_protocol import SUITE_COUNTS +from utils.misc import set_seed +from utils.wan_wrapper import WanTextEncoder, WanVAEWrapper + + +NUM_CHUNKS = base.NUM_CHUNKS +FRAMES_PER_CHUNK = base.FRAMES_PER_CHUNK +MAPPING_DEFAULT = REPO_ROOT / "assets/vbench8_extended_subset_mapping.json" +OUTPUT_DEFAULT = REPO_ROOT / "evaluation_runs/vbench8_extended_naive_baselines" + + +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, ensure_ascii=False, indent=2, allow_nan=True) + "\n", + encoding="utf-8", + ) + os.replace(temporary, path) + + +def atomic_video(frames: torch.Tensor, path: Path) -> None: + path.parent.mkdir(parents=True, exist_ok=True) + temporary = path.with_name(path.name + ".partial.mp4") + if temporary.exists(): + temporary.unlink() + base.save_mp4(frames, temporary) + os.replace(temporary, path) + + +def read_mapping(path: Path) -> list[dict[str, Any]]: + value = json.loads(path.read_text(encoding="utf-8")) + if not isinstance(value, list) or len(value) != 251: + raise ValueError(f"Expected a 251-row mapping list: {path}") + counts = {suite: 0 for suite in SUITE_COUNTS} + globals_seen: set[int] = set() + suite_seen: dict[str, set[int]] = {suite: set() for suite in SUITE_COUNTS} + for row in value: + suite = str(row["prompt_suite"]) + global_index = int(row["global_index"]) + suite_index = int(row["suite_index"]) + if suite not in counts: + raise ValueError(f"Unknown prompt suite: {suite}") + if global_index in globals_seen or suite_index in suite_seen[suite]: + raise ValueError( + f"Duplicate mapping index: {global_index} / {suite}:{suite_index}" + ) + if not str(row["extended_prompt"]).strip(): + raise ValueError(f"Empty prompt at global index {global_index}") + counts[suite] += 1 + globals_seen.add(global_index) + suite_seen[suite].add(suite_index) + if counts != SUITE_COUNTS: + raise ValueError(f"Unexpected prompt-suite counts: {counts}") + return sorted(value, key=lambda row: int(row["global_index"])) + + +def reset_runtime_caches(pipeline: Any, 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_() + for cache in pipeline.crossattn_cache: + cache["is_init"] = False + + +def warped_timesteps(pipeline: Any, raw_timesteps: tuple[int, ...]) -> torch.Tensor: + raw = torch.tensor(raw_timesteps, dtype=torch.long) + if not bool(pipeline.args.warp_denoising_step): + return raw + scheduler_timesteps = torch.cat( + (pipeline.scheduler.timesteps.cpu(), torch.tensor([0], dtype=torch.float32)) + ) + return scheduler_timesteps[1000 - raw] + + +def policy_timesteps(pipeline: Any, policy: NaivePolicy) -> torch.Tensor: + if policy.family in {"full", "velocity_reuse"}: + actual = pipeline.denoising_step_list.detach().cpu() + else: + actual = warped_timesteps(pipeline, policy.raw_timesteps) + if len(actual) != policy.denoise_steps: + raise ValueError( + f"{policy.name}: expected {policy.denoise_steps} timesteps, got {len(actual)}" + ) + return actual + + +def start_timing() -> tuple[torch.cuda.Event, torch.cuda.Event]: + start = torch.cuda.Event(enable_timing=True) + end = torch.cuda.Event(enable_timing=True) + start.record() + return start, end + + +def finish_timing( + events_by_category: dict[str, list[tuple[torch.cuda.Event, torch.cuda.Event]]], + category: str, + events: tuple[torch.cuda.Event, torch.cuda.Event], +) -> None: + events[1].record() + events_by_category[category].append(events) + + +@torch.inference_mode() +def generate_rollout( + *, + pipeline: Any, + conditional_dict: dict[str, torch.Tensor], + seed: int, + device: torch.device, + policy: NaivePolicy, +) -> tuple[torch.Tensor, dict[str, Any]]: + reset_runtime_caches(pipeline, device) + set_seed(seed) + noise = torch.randn( + 1, + NUM_CHUNKS * FRAMES_PER_CHUNK, + base.LATENT_CHANNELS, + base.LATENT_HEIGHT, + base.LATENT_WIDTH, + dtype=torch.bfloat16, + device=device, + ) + timesteps_cpu = policy_timesteps(pipeline, policy) + actual_timestep_values = [float(value) for value in timesteps_cpu] + timesteps = timesteps_cpu.to(device=device) + output_chunks: list[torch.Tensor] = [] + decisions: list[dict[str, Any]] = [] + full_calls = 0 + reuse_calls = 0 + timing_events: dict[str, list[tuple[torch.cuda.Event, torch.cuda.Event]]] = { + "full_dit": [], + "context_dit": [], + } + 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=device + ) * current_timestep + action = ( + policy.cache_pattern[step] + if policy.family == "velocity_reuse" and policy.cache_pattern is not None + else "F" + ) + if action == "R": + if cached_flow is None: + raise RuntimeError( + f"{policy.name}: reuse requested before a full step in chunk {chunk}" + ) + denoised_pred = pipeline.generator._convert_flow_pred_to_x0( + flow_pred=cached_flow.flatten(0, 1), + xt=noisy_input.flatten(0, 1), + timestep=timestep.flatten(0, 1), + ).unflatten(0, cached_flow.shape[:2]) + reuse_calls += 1 + else: + full_events = start_timing() + 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 * base.TOKENS_PER_CHUNK, + ) + finish_timing(timing_events, "full_dit", full_events) + cached_flow = flow.detach() + full_calls += 1 + decisions.append( + { + "chunk": chunk, + "step": step, + "action": action, + "timestep": actual_timestep_values[step], + } + ) + if step < len(timesteps) - 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=device + ), + ).unflatten(0, denoised_pred.shape[:2]) + + if denoised_pred is None or timestep is None: + raise RuntimeError(f"{policy.name}: chunk {chunk} produced no clean latent") + output_chunks.append(denoised_pred) + context_events = start_timing() + 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 * base.TOKENS_PER_CHUNK, + ) + finish_timing(timing_events, "context_dit", context_events) + + torch.cuda.synchronize() + elapsed = { + category: sum(start.elapsed_time(end) for start, end in events) + for category, events in timing_events.items() + } + expected_full = NUM_CHUNKS * policy.full_calls_per_chunk + expected_reuse = NUM_CHUNKS * policy.reuse_calls_per_chunk + if (full_calls, reuse_calls) != (expected_full, expected_reuse): + raise RuntimeError( + f"{policy.name}: got calls F={full_calls}, R={reuse_calls}; " + f"expected F={expected_full}, R={expected_reuse}" + ) + return torch.cat(output_chunks, dim=1), { + "generation_time_s": time.perf_counter() - started, + "full_calls": full_calls, + "reuse_calls": reuse_calls, + "predictor_calls": 0, + "accepted_predictor_calls": 0, + "rejected_predictor_calls": 0, + "full_dit_time_ms": elapsed["full_dit"], + "predictor_time_ms": 0.0, + "confidence_head_time_ms": 0.0, + "context_dit_time_ms": elapsed["context_dit"], + "actual_dit_time_ms": elapsed["full_dit"] + elapsed["context_dit"], + "model_path_time_ms": elapsed["full_dit"], + "raw_timesteps": list(policy.raw_timesteps), + "actual_timesteps": actual_timestep_values, + "decisions": decisions, + } + + +def compact_metrics(value: dict[str, Any]) -> dict[str, Any]: + keys = ( + "pixel_mse", + "psnr", + "psnr_frame_mean", + "ssim", + "lpips", + "rollout_start_frame", + "rollout_pixel_mse", + "rollout_psnr", + "rollout_ssim", + "rollout_lpips", + "num_frames", + ) + return {key: value[key] for key in keys} + + +def identity_metrics(num_frames: int) -> dict[str, Any]: + return { + "pixel_mse": 0.0, + "psnr": 120.0, + "psnr_frame_mean": 120.0, + "ssim": 1.0, + "lpips": 0.0, + "rollout_start_frame": base.PIXEL_FRAMES_FIRST_CHUNK, + "rollout_pixel_mse": 0.0, + "rollout_psnr": 120.0, + "rollout_ssim": 1.0, + "rollout_lpips": 0.0, + "num_frames": num_frames, + } + + +def load_models(device: torch.device) -> tuple[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(REPO_ROOT / "configs/self_forcing_sid.yaml"), + ) + print("[setup] loading frozen Self-Forcing Teacher", flush=True) + pipeline = base.build_pipeline( + config, REPO_ROOT / "checkpoints/self_forcing_dmd.pt", vae, device + ) + text_encoder = WanTextEncoder().to(device=device, dtype=torch.bfloat16).eval() + text_encoder.requires_grad_(False) + print("[setup] loading LPIPS", flush=True) + lpips_model = lpips.LPIPS(net="alex", verbose=False).to(device).eval() + lpips_model.requires_grad_(False) + return vae, pipeline, text_encoder, lpips_model + + +def artifact_paths( + output_root: Path, mapping_row: dict[str, Any], strategy: str +) -> tuple[Path, Path]: + suite = str(mapping_row["prompt_suite"]) + suite_index = int(mapping_row["suite_index"]) + global_index = int(mapping_row["global_index"]) + return ( + output_root / "generated_videos" / strategy / suite / f"{suite_index:03d}.mp4", + output_root + / "generation_metrics/per_prompt" + / strategy + / f"global_{global_index:04d}.json", + ) + + +def is_prompt_complete( + output_root: Path, mapping_row: dict[str, Any], policies: tuple[NaivePolicy, ...] +) -> bool: + for policy in (REFERENCE_POLICY, *policies): + video, record = artifact_paths(output_root, mapping_row, policy.name) + if not video.is_file() or not record.is_file(): + return False + return True + + +def write_result( + *, + output_root: Path, + mapping_row: dict[str, Any], + policy: NaivePolicy, + seed: int, + diagnostic: dict[str, Any], + pixel_metrics: dict[str, Any], + frames: torch.Tensor, +) -> None: + video_path, record_path = artifact_paths(output_root, mapping_row, policy.name) + atomic_video(frames, video_path) + diagnostic_without_decisions = { + key: value for key, value in diagnostic.items() if key != "decisions" + } + record = { + "status": "complete", + "strategy": policy.name, + "label": policy.label, + "policy": policy.to_dict(), + "global_index": int(mapping_row["global_index"]), + "prompt_suite": str(mapping_row["prompt_suite"]), + "suite_index": int(mapping_row["suite_index"]), + "prompt": str(mapping_row["extended_prompt"]), + "seed": seed, + "generation": diagnostic_without_decisions, + "decisions": diagnostic["decisions"], + "pixel_metrics_vs_ffff": pixel_metrics, + "video": str(video_path.relative_to(output_root)), + } + atomic_json(record_path, record) + + +def run_prompt( + *, + mapping_row: dict[str, Any], + policies: tuple[NaivePolicy, ...], + output_root: Path, + seed: int, + vae: Any, + pipeline: Any, + text_encoder: Any, + lpips_model: Any, + device: torch.device, +) -> None: + global_index = int(mapping_row["global_index"]) + suite = str(mapping_row["prompt_suite"]) + suite_index = int(mapping_row["suite_index"]) + prompt = str(mapping_row["extended_prompt"]) + print(f"[prompt] global={global_index} suite={suite}/{suite_index}", flush=True) + conditional = text_encoder(text_prompts=[prompt]) + + reference_latent, reference_diagnostic = generate_rollout( + pipeline=pipeline, + conditional_dict=conditional, + seed=seed, + device=device, + policy=REFERENCE_POLICY, + ) + with torch.autocast(device_type="cuda", dtype=torch.bfloat16): + reference_pixels = vae.decode_to_pixel(reference_latent, use_cache=False) + reference_u8 = base.pixels_to_u8(reference_pixels) + if reference_u8.ndim != 4 or reference_u8.shape[0] != 81: + raise RuntimeError(f"Unexpected reference shape: {tuple(reference_u8.shape)}") + write_result( + output_root=output_root, + mapping_row=mapping_row, + policy=REFERENCE_POLICY, + seed=seed, + diagnostic=reference_diagnostic, + pixel_metrics=identity_metrics(int(reference_u8.shape[0])), + frames=reference_u8, + ) + + for policy in policies: + latent, diagnostic = generate_rollout( + pipeline=pipeline, + conditional_dict=conditional, + seed=seed, + device=device, + policy=policy, + ) + 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) + metrics = base.frame_metrics( + reference_u8=reference_u8, + prediction_u8=prediction_u8, + lpips_model=lpips_model, + batch_size=4, + device=device, + ) + write_result( + output_root=output_root, + mapping_row=mapping_row, + policy=policy, + seed=seed, + diagnostic=diagnostic, + pixel_metrics=compact_metrics(metrics), + frames=prediction_u8, + ) + print( + f"[result] global={global_index} strategy={policy.name} " + f"F={diagnostic['full_calls']} R={diagnostic['reuse_calls']} " + f"psnr={metrics['psnr']:.4f} ssim={metrics['ssim']:.6f} " + f"lpips={metrics['lpips']:.6f}", + flush=True, + ) + if hasattr(vae.model, "clear_cache"): + vae.model.clear_cache() + del latent, pixels, prediction_u8, metrics + torch.cuda.empty_cache() + + if hasattr(vae.model, "clear_cache"): + vae.model.clear_cache() + del conditional, reference_latent, reference_pixels, reference_u8 + torch.cuda.empty_cache() + + +def write_manifest( + *, + output_root: Path, + gpu: str, + shard_index: int, + num_shards: int, + rows: list[dict[str, Any]], + policies: tuple[NaivePolicy, ...], + seed: int, + status: str, +) -> None: + selected = "_".join(policy.name for policy in policies) + selection_id = hashlib.sha256(selected.encode("utf-8")).hexdigest()[:8] + manifest_name = f"naive_gpu{gpu}_shard{shard_index:02d}_{selection_id}.json" + atomic_json( + output_root / "generation_metrics/manifests" / manifest_name, + { + "status": status, + "physical_gpu": gpu, + "shard_index": shard_index, + "num_shards": num_shards, + "global_indices": [int(row["global_index"]) for row in rows], + "num_prompts": len(rows), + "strategies": [REFERENCE_POLICY.name, *(policy.name for policy in policies)], + "seed": seed, + }, + ) + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--gpu", default=PHYSICAL_GPU) + parser.add_argument("--mapping", type=Path, default=MAPPING_DEFAULT) + parser.add_argument("--output-root", type=Path, default=OUTPUT_DEFAULT) + parser.add_argument("--shard-index", type=int, default=0) + parser.add_argument("--num-shards", type=int, default=1) + parser.add_argument( + "--strategy", + action="append", + choices=EVALUATION_STRATEGY_NAMES, + default=None, + help="Baseline to generate; repeat to select several. Default: all six.", + ) + parser.add_argument( + "--global-index", + dest="global_indices", + action="append", + type=int, + default=None, + help="Process only this global mapping index; may be repeated.", + ) + parser.add_argument("--seed", type=int, default=0) + parser.add_argument("--overwrite", action="store_true") + return parser.parse_args() + + +def main() -> None: + args = parse_args() + if args.num_shards < 1 or not 0 <= args.shard_index < args.num_shards: + raise ValueError("Invalid shard index/number of shards") + strategy_names = args.strategy or list(EVALUATION_STRATEGY_NAMES) + if len(set(strategy_names)) != len(strategy_names): + raise ValueError("--strategy values must be unique") + policies = tuple(POLICIES_BY_NAME[name] for name in strategy_names) + mapping = read_mapping(args.mapping.resolve()) + if args.global_indices is None: + rows = [ + row + for position, row in enumerate(mapping) + if position % args.num_shards == args.shard_index + ] + else: + requested = set(args.global_indices) + mapping_by_global = {int(row["global_index"]): row for row in mapping} + unknown = sorted(requested - set(mapping_by_global)) + if unknown: + raise ValueError(f"Unknown global indices: {unknown}") + rows = [row for row in mapping if int(row["global_index"]) in requested] + + output_root = args.output_root.resolve() + output_root.mkdir(parents=True, exist_ok=True) + write_manifest( + output_root=output_root, + gpu=str(args.gpu), + shard_index=args.shard_index, + num_shards=args.num_shards, + rows=rows, + policies=policies, + seed=args.seed, + status="running", + ) + print( + f"[setup] physical_gpu={args.gpu} shard={args.shard_index}/{args.num_shards} " + f"prompts={len(rows)} strategies={','.join(strategy_names)}", + flush=True, + ) + device = torch.device("cuda") + torch.set_grad_enabled(False) + vae, pipeline, text_encoder, lpips_model = load_models(device) + completed = 0 + try: + for row in rows: + if not args.overwrite and is_prompt_complete(output_root, row, policies): + completed += 1 + print( + f"[cached] {completed}/{len(rows)} global={row['global_index']}", + flush=True, + ) + continue + run_prompt( + mapping_row=row, + policies=policies, + output_root=output_root, + seed=args.seed, + vae=vae, + pipeline=pipeline, + text_encoder=text_encoder, + lpips_model=lpips_model, + device=device, + ) + completed += 1 + print(f"[progress] {completed}/{len(rows)} prompts", flush=True) + finally: + write_manifest( + output_root=output_root, + gpu=str(args.gpu), + shard_index=args.shard_index, + num_shards=args.num_shards, + rows=rows, + policies=policies, + seed=args.seed, + status="complete" if completed == len(rows) else "failed", + ) + print(f"[complete] gpu={args.gpu} prompts={completed}/{len(rows)}", flush=True) + + +if __name__ == "__main__": + main() diff --git a/scripts/generate_vbench8_extended_spanrisk.py b/scripts/generate_vbench8_extended_spanrisk.py new file mode 100644 index 0000000000000000000000000000000000000000..8cc528fb3bfca1a4afd90736f4e6f0ad352f492b --- /dev/null +++ b/scripts/generate_vbench8_extended_spanrisk.py @@ -0,0 +1,268 @@ +#!/usr/bin/env python3 +"""Generate one outgoing-span Confidence-Head strategy for Extended-251. + +The FFFF baseline is reused from the established Extended-251 run. Pixel +metrics decode both MP4s, so reference and prediction receive identical video +encoding/decoding treatment. +""" + +from __future__ import annotations + +import argparse +import json +import sys +from pathlib import Path +from typing import Any + +REPO_ROOT = Path(__file__).resolve().parents[1] +if str(REPO_ROOT) not in sys.path: + sys.path.insert(0, str(REPO_ROOT)) + +from scripts import generate_vbench8_extended_strategies as extended + +import torch +from torchvision.io import read_video + +from scripts import evaluate_single_block_fppf as base + + +MAPPING_DEFAULT = REPO_ROOT / "assets/vbench8_extended_subset_mapping.json" +EXPERIMENT_DEFAULT = ( + REPO_ROOT / "confidence_experiments/layer17_stage1_step2000_20260831" +) +SELECTION_DEFAULT = ( + REPO_ROOT + / "confidence_experiments" + / "layer17_stage1_step2000_spanrisk_beta2_threshold_vbench3_20260901" + / "threshold_summary.json" +) +REFERENCE_DEFAULT = ( + REPO_ROOT / "evaluation_runs/vbench8_extended_stage1_step2000_20260901" +) +OUTPUT_DEFAULT = ( + REPO_ROOT + / "evaluation_runs/vbench8_extended_stage1_step2000_spanrisk_beta2_20260901" +) +STRATEGIES = ( + "step12_span_k06", + "step12_span_k08", + "step12_span_k10", + "step123_span_k06", + "step123_span_k09", + "step123_span_k12", + "step123_span_k15", +) + + +def load_strategy_configs(path: Path) -> dict[str, dict[str, Any]]: + payload = json.loads(path.read_text(encoding="utf-8")) + if payload.get("status") != "complete" or float(payload["beta"]) != 2.0: + raise ValueError(f"Threshold selection is not completed fixed-beta=2: {path}") + configs: dict[str, dict[str, Any]] = {} + for row in payload["selected"]: + head = str(row["head"]) + target = int(row["target_accepts"]) + name = f"{head}_span_k{target:02d}" + configs[name] = { + "name": name, + "policy": "dynamic", + "candidate_steps": [int(step) for step in row["candidate_steps"]], + "beta": float(row["beta"]), + "threshold": float(row["threshold"]), + "target_accepts": target, + "head": head, + "risk_mode": "outgoing_span", + "source_config_name": str(row["name"]), + } + if tuple(configs) != STRATEGIES: + raise ValueError(f"Unexpected selected strategies: {tuple(configs)}") + return configs + + +def read_u8_video(path: Path) -> torch.Tensor: + if not path.is_file(): + raise FileNotFoundError(path) + frames, _, _ = read_video(str(path), pts_unit="sec", output_format="TCHW") + frames = frames.to(device="cpu", dtype=torch.uint8).contiguous() + if frames.ndim != 4 or frames.shape[0] != 81 or frames.shape[1] != 3: + raise RuntimeError(f"Unexpected decoded video shape {tuple(frames.shape)}: {path}") + return frames + + +def artifact_paths( + output_root: Path, mapping_row: dict[str, Any], strategy_name: str +) -> tuple[Path, Path]: + suite = str(mapping_row["prompt_suite"]) + suite_index = int(mapping_row["suite_index"]) + global_index = int(mapping_row["global_index"]) + return ( + output_root / "generated_videos" / strategy_name / suite / f"{suite_index:03d}.mp4", + output_root + / "generation_metrics/per_prompt" + / strategy_name + / f"global_{global_index:04d}.json", + ) + + +def run_prompt( + *, + mapping_row: dict[str, Any], + strategy: dict[str, Any], + reference_root: Path, + output_root: Path, + seed: int, + models: tuple[Any, ...], + device: torch.device, +) -> None: + global_index = int(mapping_row["global_index"]) + suite = str(mapping_row["prompt_suite"]) + suite_index = int(mapping_row["suite_index"]) + prompt = str(mapping_row["extended_prompt"]) + name = str(strategy["name"]) + reference_path = ( + reference_root / "generated_videos/ffff" / suite / f"{suite_index:03d}.mp4" + ) + print(f"[prompt] global={global_index} suite={suite}/{suite_index}", flush=True) + conditional = models[3](text_prompts=[prompt]) + head = models[4] if strategy["head"] == "step12" else models[5] + latent, diagnostic = extended.generate_rollout( + pipeline=models[1], + conditional_dict=conditional, + seed=seed, + device=device, + predictor=models[2], + head=head, + config=strategy, + ) + with torch.autocast(device_type="cuda", dtype=torch.bfloat16): + pixels = models[0].decode_to_pixel(latent, use_cache=False) + prediction_u8 = base.pixels_to_u8(pixels) + video_path, record_path = artifact_paths(output_root, mapping_row, name) + extended.atomic_video(prediction_u8, video_path) + + reference_decoded = read_u8_video(reference_path) + prediction_decoded = read_u8_video(video_path) + metrics = base.frame_metrics( + reference_u8=reference_decoded, + prediction_u8=prediction_decoded, + lpips_model=models[6], + batch_size=4, + device=device, + ) + record = { + "status": "complete", + "strategy": name, + "policy": "dynamic", + "candidate_steps": strategy["candidate_steps"], + "beta": strategy["beta"], + "threshold": strategy["threshold"], + "target_accepts": strategy["target_accepts"], + "risk_mode": "outgoing_span", + "allow_chunk0_predictor": False, + "source_config_name": strategy["source_config_name"], + "global_index": global_index, + "prompt_suite": suite, + "suite_index": suite_index, + "prompt": prompt, + "seed": seed, + "generation": { + key: value for key, value in diagnostic.items() if key != "decisions" + }, + "decisions": diagnostic["decisions"], + "pixel_metrics_vs_ffff": extended.compact_metrics(metrics), + "pixel_metric_input": "MP4-decoded uint8 RGB, all 81 frames on both sides", + "reference_video": str(reference_path), + "video": str(video_path.relative_to(output_root)), + } + extended.atomic_json(record_path, record) + print( + f"[result] global={global_index} strategy={name} " + f"accept={diagnostic['accepted_predictor_calls']} " + f"psnr={metrics['psnr']:.4f} ssim={metrics['ssim']:.6f} " + f"lpips={metrics['lpips']:.6f}", + flush=True, + ) + if hasattr(models[0].model, "clear_cache"): + models[0].model.clear_cache() + del conditional, latent, pixels, prediction_u8, reference_decoded, prediction_decoded + torch.cuda.empty_cache() + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--gpu", default=extended.PHYSICAL_GPU) + parser.add_argument("--strategy", required=True, choices=STRATEGIES) + parser.add_argument("--mapping", type=Path, default=MAPPING_DEFAULT) + parser.add_argument("--experiment-root", type=Path, default=EXPERIMENT_DEFAULT) + parser.add_argument("--selection", type=Path, default=SELECTION_DEFAULT) + parser.add_argument("--reference-root", type=Path, default=REFERENCE_DEFAULT) + parser.add_argument("--output-root", type=Path, default=OUTPUT_DEFAULT) + parser.add_argument("--global-index", action="append", type=int, default=None) + parser.add_argument("--seed", type=int, default=0) + parser.add_argument("--overwrite", action="store_true") + return parser.parse_args() + + +def main() -> None: + args = parse_args() + mapping = extended.read_mapping(args.mapping.resolve()) + configs = load_strategy_configs(args.selection.resolve()) + strategy = configs[args.strategy] + reference_root = args.reference_root.resolve() + if len(list((reference_root / "generation_metrics/per_prompt/ffff").glob("global_*.json"))) != 251: + raise ValueError("Incomplete FFFF reference records") + if len(list((reference_root / "generated_videos/ffff").rglob("*.mp4"))) != 251: + raise ValueError("Incomplete FFFF reference videos") + if args.global_index is not None: + requested = set(args.global_index) + mapping = [row for row in mapping if int(row["global_index"]) in requested] + if len(mapping) != len(requested): + raise ValueError("At least one requested global index is unknown") + + output_root = args.output_root.resolve() + output_root.mkdir(parents=True, exist_ok=True) + extended.write_shard_manifest( + output_root, str(args.gpu), 0, 1, mapping, [strategy], "running" + ) + print( + f"[setup] physical_gpu={args.gpu} strategy={args.strategy} " + f"prompts={len(mapping)} reference=existing_ffff chunk0=ffff", + flush=True, + ) + device = torch.device("cuda") + torch.set_grad_enabled(False) + models = extended.load_models(args.experiment_root.resolve(), device) + completed = 0 + try: + for row in mapping: + video_path, record_path = artifact_paths(output_root, row, args.strategy) + if not args.overwrite and video_path.is_file() and record_path.is_file(): + completed += 1 + print(f"[cached] {completed}/{len(mapping)} global={row['global_index']}", flush=True) + continue + run_prompt( + mapping_row=row, + strategy=strategy, + reference_root=reference_root, + output_root=output_root, + seed=args.seed, + models=models, + device=device, + ) + completed += 1 + print(f"[progress] {completed}/{len(mapping)} prompts", flush=True) + finally: + extended.write_shard_manifest( + output_root, + str(args.gpu), + 0, + 1, + mapping, + [strategy], + "complete" if completed == len(mapping) else "failed", + ) + print(f"[complete] gpu={args.gpu} strategy={args.strategy} prompts={completed}", flush=True) + + +if __name__ == "__main__": + main() diff --git a/scripts/generate_vbench8_extended_strategies.py b/scripts/generate_vbench8_extended_strategies.py new file mode 100644 index 0000000000000000000000000000000000000000..7a5260a0b714899676eb37899fef77a8ff5217d9 --- /dev/null +++ b/scripts/generate_vbench8_extended_strategies.py @@ -0,0 +1,781 @@ +#!/usr/bin/env python3 +"""Generate all Layer-17 strategies for the VBench-8 extended subset. + +One process handles a prompt shard and keeps the frozen Teacher, Layer-17 +Predictor, both confidence heads, VAE, text encoder, and LPIPS model resident +on one GPU. For each prompt, FFFF is generated first and all other strategies +are compared with that same-prompt, same-seed reference before their videos +are written. +""" + +from __future__ import annotations + +import argparse +import csv +import json +import math +import os +import sys +import tempfile +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) + os.environ.setdefault("MPLCONFIGDIR", tempfile.mkdtemp(prefix="self_forcing_mpl_")) + 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_single_block_fppf as base +from scripts.evaluate_layer17_dynamic_gate import predictor_with_features +from scripts.vbench8_protocol import DIMENSIONS, SUITE_COUNTS +from utils.misc import set_seed +from utils.wan_wrapper import WanTextEncoder, WanVAEWrapper + + +NUM_CHUNKS = base.NUM_CHUNKS +FRAMES_PER_CHUNK = base.FRAMES_PER_CHUNK +NUM_DENOISING_STEPS = base.NUM_DENOISING_STEPS + +MAPPING_DEFAULT = REPO_ROOT / "assets/vbench8_extended_subset_mapping.json" +EXPERIMENT_DEFAULT = REPO_ROOT / "confidence_experiments/layer17_stage1_step2000_20260831" +OUTPUT_DEFAULT = REPO_ROOT / "evaluation_runs/vbench8_extended_stage1_step2000_20260901" + + +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, ensure_ascii=False, indent=2, allow_nan=True) + "\n", + encoding="utf-8", + ) + os.replace(temporary, path) + + +def atomic_video(frames: torch.Tensor, path: Path) -> None: + path.parent.mkdir(parents=True, exist_ok=True) + temporary = path.with_name(path.name + ".partial.mp4") + if temporary.exists(): + temporary.unlink() + base.save_mp4(frames, temporary) + os.replace(temporary, path) + + +def read_mapping(path: Path) -> list[dict[str, Any]]: + value = json.loads(path.read_text(encoding="utf-8")) + if not isinstance(value, list): + raise ValueError(f"Mapping must be a JSON list: {path}") + if len(value) != 251: + raise ValueError(f"Expected 251 mapping rows, got {len(value)}") + counts = {suite: 0 for suite in SUITE_COUNTS} + seen_global: set[int] = set() + seen_suite: dict[str, set[int]] = {suite: set() for suite in SUITE_COUNTS} + for row in value: + suite = str(row["prompt_suite"]) + global_index = int(row["global_index"]) + suite_index = int(row["suite_index"]) + if suite not in counts: + raise ValueError(f"Unknown prompt suite in mapping: {suite}") + if global_index in seen_global: + raise ValueError(f"Duplicate global index: {global_index}") + if suite_index in seen_suite[suite]: + raise ValueError(f"Duplicate suite index: {suite}/{suite_index}") + if not str(row["extended_prompt"]).strip(): + raise ValueError(f"Empty extended prompt at global index {global_index}") + seen_global.add(global_index) + seen_suite[suite].add(suite_index) + counts[suite] += 1 + if counts != SUITE_COUNTS: + raise ValueError(f"Unexpected mapping counts: {counts}") + for suite, expected in SUITE_COUNTS.items(): + if seen_suite[suite] != set(range(expected)): + raise ValueError(f"Non-contiguous suite indices for {suite}") + return sorted(value, key=lambda row: int(row["global_index"])) + + +def find_selected(path: Path, config_name: str) -> dict[str, Any]: + value = json.loads(path.read_text(encoding="utf-8")) + matches = [row for row in value["selected_dynamic"] if row["config_name"] == config_name] + if len(matches) != 1: + raise ValueError(f"Expected one {config_name} in {path}, got {len(matches)}") + return matches[0] + + +def load_strategy_configs(experiment_root: Path) -> list[dict[str, Any]]: + step12_selected = experiment_root / "dynamic_step12/validation/selected.json" + step123_selected = experiment_root / "dynamic_step123/validation/selected.json" + configs: list[dict[str, Any]] = [ + { + "name": "ffff", + "policy": "ffff", + "candidate_steps": [], + "beta": None, + "threshold": None, + "target_accepts": 0, + "head": None, + }, + { + "name": "fppf", + "policy": "fppf", + "candidate_steps": [1, 2], + "beta": None, + "threshold": None, + "target_accepts": 12, + "head": None, + }, + ] + for source, head, names in ( + (step12_selected, "step12", ("dynamic_b1p0_k06", "dynamic_b2p0_k08", "dynamic_b0p0_k10")), + (step123_selected, "step123", ("dynamic_b1p0_k06", "dynamic_b2p0_k09", "dynamic_b1p5_k12", "dynamic_b2p0_k15")), + ): + for source_name in names: + row = find_selected(source, source_name) + target = int(row["target_accepts"]) + candidates = [1, 2] if head == "step12" else [1, 2, 3] + configs.append( + { + "name": f"{head}_k{target:02d}", + "policy": "dynamic", + "candidate_steps": candidates, + "beta": float(row["beta"]), + "threshold": float(row["threshold"]), + "target_accepts": target, + "head": head, + "source_config_name": source_name, + } + ) + if [config["name"] for config in configs] != [ + "ffff", "fppf", "step12_k06", "step12_k08", "step12_k10", + "step123_k06", "step123_k09", "step123_k12", "step123_k15", + ]: + raise AssertionError("Unexpected strategy order") + return configs + + +def reset_runtime_caches(pipeline: Any, 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_() + for cache in pipeline.crossattn_cache: + cache["is_init"] = False + + +@torch.inference_mode() +def generate_rollout( + *, + pipeline: Any, + conditional_dict: dict[str, torch.Tensor], + seed: int, + device: torch.device, + predictor: Any, + head: Any, + config: dict[str, Any], +) -> tuple[torch.Tensor, dict[str, Any]]: + reset_runtime_caches(pipeline, device) + set_seed(seed) + noise = torch.randn( + 1, + NUM_CHUNKS * FRAMES_PER_CHUNK, + base.LATENT_CHANNELS, + base.LATENT_HEIGHT, + base.LATENT_WIDTH, + dtype=torch.bfloat16, + device=device, + ) + teacher = pipeline.generator.model + timesteps = pipeline.denoising_step_list.to(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(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 + allow_chunk0_predictor = bool( + config.get("allow_chunk0_predictor", False) + ) + candidate = ( + step in config["candidate_steps"] + and (chunk > 0 or allow_chunk0_predictor) + ) + policy = str(config["policy"]) + run_predictor = candidate and policy in {"dynamic", "fppf"} + accepted = False + pred_hidden = None + pred_x0 = None + predicted_local_error = None + risk = None + alpha = None + outgoing_span = None + span_weight = None + if run_predictor: + anchor_hidden = current_hidden[step - 1] + if anchor_hidden is None: + raise RuntimeError("Predictor inputs are unavailable") + if previous_chunk_hidden is None: + if not ( + chunk == 0 + and allow_chunk0_predictor + and getattr(predictor, "input_variant", None) == "disca" + ): + raise RuntimeError("Previous chunk hidden is unavailable") + # DISCA physically has no previous-chunk channel. Its + # forward signature retains this argument only for API + # compatibility, so an anchor-shaped placeholder is safe. + previous_hidden = anchor_hidden + else: + previous_hidden = previous_chunk_hidden[step] + if previous_hidden is None: + raise RuntimeError("Previous chunk hidden is 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 == "fppf": + accepted = True + else: + 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 = (NUM_CHUNKS - 1 - chunk) / (NUM_CHUNKS - 2) + next_policy_timestep = ( + timesteps[step + 1] + if step + 1 < len(timesteps) + else timesteps.new_zeros(()) + ) + outgoing_span = abs( + float(current_timestep) - float(next_policy_timestep) + ) + risk_mode = str(config.get("risk_mode", "legacy")) + if risk_mode == "legacy": + span_weight = 1.0 + elif risk_mode == "outgoing_span": + span_weight = outgoing_span / 1000.0 + else: + raise ValueError(f"Unknown risk mode: {risk_mode}") + risk = predicted_local_error * span_weight * ( + 1.0 + float(config["beta"]) * alpha + ) + accepted = risk <= float(config["threshold"]) + + 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, + "current_timestep": float(current_timestep), + "next_policy_timestep": ( + float(timesteps[step + 1]) + if step + 1 < len(timesteps) + else 0.0 + ), + "outgoing_span": outgoing_span, + "span_weight": span_weight, + "impact_risk": risk, + } + ) + if step < NUM_DENOISING_STEPS - 1: + assert denoised_pred is not None + 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=device + ), + ).unflatten(0, denoised_pred.shape[:2]) + + assert denoised_pred is not None and timestep is not None + output_chunks.append(denoised_pred) + context_events = start_timing() + 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 * 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"] + ) + policy_latency_ms = ( + elapsed["full_dit"] + elapsed["predictor"] + elapsed["confidence"] + ) + 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"], + "policy_latency_ms": policy_latency_ms, + "decisions": decisions, + } + + +def compact_metrics(value: dict[str, Any]) -> dict[str, Any]: + keys = ( + "pixel_mse", "psnr", "psnr_frame_mean", "ssim", "lpips", + "rollout_start_frame", "rollout_pixel_mse", "rollout_psnr", + "rollout_ssim", "rollout_lpips", "num_frames", + ) + return {key: value[key] for key in keys} + + +def identity_metrics(num_frames: int) -> dict[str, Any]: + return { + "pixel_mse": 0.0, + "psnr": 120.0, + "psnr_frame_mean": 120.0, + "ssim": 1.0, + "lpips": 0.0, + "rollout_start_frame": base.PIXEL_FRAMES_FIRST_CHUNK, + "rollout_pixel_mse": 0.0, + "rollout_psnr": 120.0, + "rollout_ssim": 1.0, + "rollout_lpips": 0.0, + "num_frames": num_frames, + } + + +def load_models( + experiment_root: Path, device: torch.device +) -> tuple[Any, Any, 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(REPO_ROOT / "configs/self_forcing_sid.yaml"), + ) + print("[setup] loading frozen Teacher and Layer-17 Predictor", flush=True) + pipeline = base.build_pipeline( + config, REPO_ROOT / "checkpoints/self_forcing_dmd.pt", vae, device + ) + predictor = base.load_predictor( + pipeline.generator.model, + { + "source_layer": 17, + "weights": REPO_ROOT + / "training_runs/layer17_stage1_1000p_4gpu_b16_2000steps/checkpoint_step_2000/predictor.safetensors", + "gate_mode": "baseline", + }, + device, + ) + print("[setup] loading text encoder and confidence heads", flush=True) + text_encoder = WanTextEncoder().to(device=device, dtype=torch.bfloat16).eval() + text_encoder.requires_grad_(False) + step12_head = PredictorConfidenceHead(num_steps=2).to(device=device).eval() + step12_head.load_state_dict( + load_file( + str(experiment_root / "confidence_step12/confidence_best.safetensors"), + device="cpu", + ), + strict=True, + ) + step12_head.requires_grad_(False) + step123_head = PredictorConfidenceHead(num_steps=3).to(device=device).eval() + step123_head.load_state_dict( + load_file( + str(experiment_root / "confidence_step123/confidence_best.safetensors"), + device="cpu", + ), + strict=True, + ) + step123_head.requires_grad_(False) + print("[setup] loading LPIPS", flush=True) + lpips_model = lpips.LPIPS(net="alex", verbose=False).to(device).eval() + lpips_model.requires_grad_(False) + return ( + vae, + pipeline, + predictor, + text_encoder, + step12_head, + step123_head, + lpips_model, + ) + + +def is_prompt_complete(output_root: Path, mapping_row: dict[str, Any], strategies: list[dict[str, Any]]) -> bool: + global_index = int(mapping_row["global_index"]) + suite = str(mapping_row["prompt_suite"]) + suite_index = int(mapping_row["suite_index"]) + for strategy in strategies: + name = strategy["name"] + video = output_root / "generated_videos" / name / suite / f"{suite_index:03d}.mp4" + record = output_root / "generation_metrics/per_prompt" / name / f"global_{global_index:04d}.json" + if not video.is_file() or not record.is_file(): + return False + return True + + +def run_prompt( + *, + mapping_row: dict[str, Any], + strategies: list[dict[str, Any]], + output_root: Path, + seed: int, + vae: Any, + pipeline: Any, + predictor: Any, + text_encoder: Any, + step12_head: Any, + step123_head: Any, + lpips_model: Any, + device: torch.device, +) -> None: + global_index = int(mapping_row["global_index"]) + suite = str(mapping_row["prompt_suite"]) + suite_index = int(mapping_row["suite_index"]) + prompt = str(mapping_row["extended_prompt"]) + print(f"[prompt] global={global_index} suite={suite}/{suite_index}", flush=True) + conditional = text_encoder(text_prompts=[prompt]) + ffff_config = strategies[0] + reference_latent, ffff_diagnostic = generate_rollout( + pipeline=pipeline, + conditional_dict=conditional, + seed=seed, + device=device, + predictor=predictor, + head=None, + config=ffff_config, + ) + with torch.autocast(device_type="cuda", dtype=torch.bfloat16): + reference_pixels = vae.decode_to_pixel(reference_latent, use_cache=False) + reference_u8 = base.pixels_to_u8(reference_pixels) + if reference_u8.ndim != 4 or reference_u8.shape[0] != 81: + raise RuntimeError(f"Unexpected decoded reference shape: {tuple(reference_u8.shape)}") + reference_path = ( + output_root / "generated_videos" / "ffff" / suite / f"{suite_index:03d}.mp4" + ) + atomic_video(reference_u8, reference_path) + reference_record = { + "status": "complete", + "strategy": "ffff", + "policy": "ffff", + "global_index": global_index, + "prompt_suite": suite, + "suite_index": suite_index, + "prompt": prompt, + "seed": seed, + "generation": {key: value for key, value in ffff_diagnostic.items() if key != "decisions"}, + "pixel_metrics_vs_ffff": identity_metrics(int(reference_u8.shape[0])), + "video": str(reference_path.relative_to(output_root)), + } + atomic_json( + output_root + / "generation_metrics/per_prompt/ffff" + / f"global_{global_index:04d}.json", + reference_record, + ) + + for strategy in strategies[1:]: + name = str(strategy["name"]) + head = step12_head if strategy["head"] == "step12" else step123_head + latent, diagnostic = generate_rollout( + pipeline=pipeline, + conditional_dict=conditional, + seed=seed, + device=device, + predictor=predictor, + head=head, + config=strategy, + ) + 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) + metrics = base.frame_metrics( + reference_u8=reference_u8, + prediction_u8=prediction_u8, + lpips_model=lpips_model, + batch_size=4, + device=device, + ) + video_path = output_root / "generated_videos" / name / suite / f"{suite_index:03d}.mp4" + atomic_video(prediction_u8, video_path) + record = { + "status": "complete", + "strategy": name, + "policy": strategy["policy"], + "candidate_steps": strategy["candidate_steps"], + "beta": strategy["beta"], + "threshold": strategy["threshold"], + "target_accepts": strategy["target_accepts"], + "source_config_name": strategy.get("source_config_name"), + "global_index": global_index, + "prompt_suite": suite, + "suite_index": suite_index, + "prompt": prompt, + "seed": seed, + "generation": {key: value for key, value in diagnostic.items() if key != "decisions"}, + "decisions": diagnostic["decisions"], + "pixel_metrics_vs_ffff": compact_metrics(metrics), + "video": str(video_path.relative_to(output_root)), + } + atomic_json( + output_root + / "generation_metrics/per_prompt" + / name + / f"global_{global_index:04d}.json", + record, + ) + print( + f"[result] global={global_index} strategy={name} " + f"accept={diagnostic['accepted_predictor_calls']} " + f"psnr={metrics['psnr']:.4f} ssim={metrics['ssim']:.6f} " + f"lpips={metrics['lpips']:.6f}", + flush=True, + ) + if hasattr(vae.model, "clear_cache"): + vae.model.clear_cache() + del latent, pixels, prediction_u8, metrics + torch.cuda.empty_cache() + + if hasattr(vae.model, "clear_cache"): + vae.model.clear_cache() + del conditional, reference_latent, reference_pixels, reference_u8 + torch.cuda.empty_cache() + + +def write_shard_manifest( + output_root: Path, + gpu: str, + shard_index: int, + num_shards: int, + rows: list[dict[str, Any]], + strategies: list[dict[str, Any]], + status: str, +) -> None: + atomic_json( + output_root / "generation_metrics" / f"manifest_gpu{gpu}.json", + { + "status": status, + "physical_gpu": gpu, + "shard_index": shard_index, + "num_shards": num_shards, + "global_indices": [int(row["global_index"]) for row in rows], + "num_prompts": len(rows), + "strategies": [strategy["name"] for strategy in strategies], + "seed": 0, + }, + ) + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--gpu", default=PHYSICAL_GPU) + parser.add_argument("--mapping", type=Path, default=MAPPING_DEFAULT) + parser.add_argument("--experiment-root", type=Path, default=EXPERIMENT_DEFAULT) + parser.add_argument("--output-root", type=Path, default=OUTPUT_DEFAULT) + parser.add_argument("--shard-index", type=int, default=0) + parser.add_argument("--num-shards", type=int, default=1) + parser.add_argument( + "--global-index", + dest="global_indices", + action="append", + type=int, + default=None, + help="Process only the specified global index; may be repeated.", + ) + parser.add_argument("--seed", type=int, default=0) + parser.add_argument("--overwrite", action="store_true") + return parser.parse_args() + + +def main() -> None: + args = parse_args() + if args.num_shards < 1 or not 0 <= args.shard_index < args.num_shards: + raise ValueError("Invalid shard index/number of shards") + mapping = read_mapping(args.mapping.resolve()) + strategies = load_strategy_configs(args.experiment_root.resolve()) + if args.global_indices is None: + rows = [ + row for position, row in enumerate(mapping) + if position % args.num_shards == args.shard_index + ] + else: + requested = set(args.global_indices) + mapping_by_global = {int(row["global_index"]): row for row in mapping} + unknown = sorted(requested - set(mapping_by_global)) + if unknown: + raise ValueError(f"Unknown global indices: {unknown}") + rows = [row for row in mapping if int(row["global_index"]) in requested] + output_root = args.output_root.resolve() + output_root.mkdir(parents=True, exist_ok=True) + write_shard_manifest( + output_root, str(args.gpu), args.shard_index, args.num_shards, + rows, strategies, "running", + ) + print( + f"[setup] physical_gpu={args.gpu} shard={args.shard_index}/{args.num_shards} " + f"prompts={len(rows)} strategies={len(strategies)}", + flush=True, + ) + device = torch.device("cuda") + torch.set_grad_enabled(False) + ( + vae, + pipeline, + predictor, + text_encoder, + step12_head, + step123_head, + lpips_model, + ) = load_models(args.experiment_root.resolve(), device) + completed = 0 + try: + for row in rows: + if not args.overwrite and is_prompt_complete(output_root, row, strategies): + completed += 1 + print( + f"[cached] {completed}/{len(rows)} global={row['global_index']}", + flush=True, + ) + continue + run_prompt( + mapping_row=row, + strategies=strategies, + output_root=output_root, + seed=args.seed, + vae=vae, + pipeline=pipeline, + predictor=predictor, + text_encoder=text_encoder, + step12_head=step12_head, + step123_head=step123_head, + lpips_model=lpips_model, + device=device, + ) + completed += 1 + print(f"[progress] {completed}/{len(rows)} prompts", flush=True) + finally: + write_shard_manifest( + output_root, str(args.gpu), args.shard_index, args.num_shards, + rows, strategies, "complete" if completed == len(rows) else "failed", + ) + print(f"[complete] gpu={args.gpu} prompts={completed}/{len(rows)}", flush=True) + + +if __name__ == "__main__": + main() diff --git a/scripts/merge_three_block_sweep_shards.py b/scripts/merge_three_block_sweep_shards.py new file mode 100644 index 0000000000000000000000000000000000000000..17b2ba59c1023627440e3e8bdda3b6863870b272 --- /dev/null +++ b/scripts/merge_three_block_sweep_shards.py @@ -0,0 +1,159 @@ +#!/usr/bin/env python3 +"""Merge four independent consecutive three-block training sweeps.""" + +from __future__ import annotations + +import argparse +import csv +import json +import os +import shutil +from pathlib import Path +from typing import Any + + +EXPECTED_TRIPLES = [(index, index + 1, index + 2) for index in range(28)] +EXPECTED_PROMPTS = {"train": list(range(80)), "val": list(range(80, 100))} +GPU_SHARDS = { + "0": EXPECTED_TRIPLES[0:7], + "2": EXPECTED_TRIPLES[7:14], + "6": EXPECTED_TRIPLES[14:21], + "7": EXPECTED_TRIPLES[21:28], +} + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--shard_root", type=Path, required=True) + parser.add_argument("--output_dir", type=Path, required=True) + return parser.parse_args() + + +def experiment_name(triple: tuple[int, int, int]) -> str: + return "triple_" + "_".join(f"{layer:02d}" for layer in triple) + + +def atomic_json(path: Path, value: Any) -> None: + temporary = path.with_suffix(path.suffix + ".tmp") + temporary.write_text( + json.dumps(value, indent=2, ensure_ascii=False) + "\n", + encoding="utf-8", + ) + os.replace(temporary, path) + + +def validate_metrics( + metrics: dict[str, Any], triple: tuple[int, int, int] +) -> None: + name = experiment_name(triple) + expected = { + "status": "complete", + "name": name, + "architecture": "three_block_predictor", + "initialization_method": "teacher_full", + "source_layers": list(triple), + "seed": 0, + "max_steps": 1000, + "train_prompt_ids": EXPECTED_PROMPTS["train"], + "val_prompt_ids": EXPECTED_PROMPTS["val"], + } + for key, value in expected.items(): + if metrics.get(key) != value: + raise RuntimeError( + f"{name}: expected {key}={value!r}, got {metrics.get(key)!r}" + ) + + +def result_row(metrics: dict[str, Any]) -> dict[str, Any]: + return { + "name": metrics["name"], + "source_layer_1": metrics["source_layers"][0], + "source_layer_2": metrics["source_layers"][1], + "source_layer_3": metrics["source_layers"][2], + "final_val_flow_mse": metrics["final_val_flow_mse"], + "final_val_hidden_mse": metrics["final_val_hidden_mse"], + "final_val_total_loss": metrics["final_val_total_loss"], + "val_flow_mse_auc": metrics["val_flow_mse_auc"], + "val_hidden_mse_auc": metrics["val_hidden_mse_auc"], + "val_total_loss_auc": metrics["val_total_loss_auc"], + "training_time_s": metrics["training_time_s"], + } + + +def write_summary(output_dir: Path, rows: list[dict[str, Any]]) -> None: + rows.sort(key=lambda row: float(row["final_val_flow_mse"])) + 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() + shard_root = args.shard_root.resolve() + output_dir = args.output_dir.resolve() + output_dir.mkdir(parents=True, exist_ok=True) + + rows = [] + fingerprints = set() + for gpu, triples in GPU_SHARDS.items(): + shard_dir = shard_root / f"gpu{gpu}" + shard_manifest = json.loads( + (shard_dir / "sweep_manifest.json").read_text(encoding="utf-8") + ) + if shard_manifest.get("status") != "complete": + raise RuntimeError(f"GPU {gpu} shard is not complete") + if shard_manifest.get("triples") != [list(value) for value in triples]: + raise RuntimeError(f"GPU {gpu} shard contains unexpected triples") + fingerprints.add(shard_manifest["batch_schedule_sha256"]) + + for triple in triples: + name = experiment_name(triple) + source = shard_dir / name + metrics = json.loads( + (source / "metrics.json").read_text(encoding="utf-8") + ) + validate_metrics(metrics, triple) + destination = output_dir / name + if destination.exists(): + shutil.rmtree(destination) + shutil.copytree(source, destination) + rows.append(result_row(metrics)) + + if len(rows) != 28: + raise RuntimeError(f"Expected 28 results, found {len(rows)}") + if len(fingerprints) != 1: + raise RuntimeError(f"Batch schedules differ across shards: {fingerprints}") + write_summary(output_dir, rows) + atomic_json( + output_dir / "sweep_manifest.json", + { + "status": "complete", + "architecture": "three_block_predictor", + "initialization_method": "teacher_full", + "triples": [list(value) for value in EXPECTED_TRIPLES], + "num_triples": 28, + "max_steps": 1000, + "seed": 0, + "train_prompt_ids": EXPECTED_PROMPTS["train"], + "val_prompt_ids": EXPECTED_PROMPTS["val"], + "batch_schedule_sha256": fingerprints.pop(), + "gpu_shards": { + gpu: [list(value) for value in triples] + for gpu, triples in GPU_SHARDS.items() + }, + "source_shard_root": str(shard_root), + }, + ) + print( + f"[merge] {len(rows)} triples -> {output_dir / 'summary.csv'}", + flush=True, + ) + + +if __name__ == "__main__": + main() diff --git a/scripts/merge_two_block_fppp_shards.py b/scripts/merge_two_block_fppp_shards.py new file mode 100644 index 0000000000000000000000000000000000000000..5be1610633ff846655b93d78925768fe6e06d0e5 --- /dev/null +++ b/scripts/merge_two_block_fppp_shards.py @@ -0,0 +1,126 @@ +#!/usr/bin/env python3 +"""Merge independently evaluated two-block FPPP shards into one result set.""" + +from __future__ import annotations + +import argparse +import csv +import json +import os +import shutil +from pathlib import Path +from typing import Any + + +PROMPT_IDS = list(range(80, 100)) +SCHEDULE = "chunk0=FFFF; chunks1-6=FPPP" + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--main_dir", type=Path, required=True) + parser.add_argument("--shard_dir", type=Path, required=True) + parser.add_argument("--sweep_dir", type=Path, required=True) + return parser.parse_args() + + +def atomic_json(path: Path, value: Any) -> None: + temporary = path.with_suffix(path.suffix + ".tmp") + temporary.write_text( + json.dumps(value, indent=2, ensure_ascii=False) + "\n", + encoding="utf-8", + ) + os.replace(temporary, path) + + +def read_experiment_names(sweep_dir: Path) -> list[str]: + with (sweep_dir / "summary.csv").open( + "r", encoding="utf-8", newline="" + ) as handle: + return [row["name"] for row in csv.DictReader(handle)] + + +def merge_pair_directories(main_dir: Path, shard_dir: Path) -> None: + for source in shard_dir.glob("pair_*"): + if source.is_dir(): + shutil.copytree(source, main_dir / source.name, dirs_exist_ok=True) + + +def collect_rows(main_dir: Path, names: list[str]) -> list[dict[str, Any]]: + rows = [] + for name in names: + metrics_path = main_dir / name / "metrics.json" + if not metrics_path.exists(): + raise FileNotFoundError(f"Missing completed metrics: {metrics_path}") + metrics = json.loads(metrics_path.read_text(encoding="utf-8")) + if metrics.get("status") != "complete": + raise RuntimeError(f"Incomplete metrics for {name}") + if metrics.get("schedule") != SCHEDULE: + raise RuntimeError(f"Unexpected schedule for {name}: {metrics.get('schedule')}") + if metrics.get("prompt_ids") != PROMPT_IDS: + raise RuntimeError(f"Unexpected prompt split for {name}") + rows.append( + { + "name": metrics["name"], + "source_layer_1": metrics["source_layers"][0], + "source_layer_2": metrics["source_layers"][1], + "pair_kind": metrics["pair_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.sort(key=lambda row: float(row["lpips"])) + return rows + + +def write_summary(main_dir: Path, rows: list[dict[str, Any]]) -> None: + destination = main_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(main_dir / "summary.json", rows) + + +def main() -> None: + args = parse_args() + main_dir = args.main_dir.resolve() + shard_dir = args.shard_dir.resolve() + sweep_dir = args.sweep_dir.resolve() + names = read_experiment_names(sweep_dir) + if len(names) != 43: + raise RuntimeError(f"Expected 43 experiments, found {len(names)}") + merge_pair_directories(main_dir, shard_dir) + rows = collect_rows(main_dir, names) + write_summary(main_dir, rows) + atomic_json( + main_dir / "manifest.json", + { + "status": "complete", + "architecture": "two_block_predictor", + "prompt_ids": PROMPT_IDS, + "experiments": names, + "rollout_schedule": "FPPP", + "rollout_definition": SCHEDULE, + "generation_seed_reset_per_prompt": 0, + "gpu_shards": {"2": 22, "4": 21}, + "merged_shard_dir": str(shard_dir), + }, + ) + print(f"[merge] {len(rows)} pairs -> {main_dir / 'summary.csv'}", flush=True) + + +if __name__ == "__main__": + main() diff --git a/scripts/naive_vbench_policies.py b/scripts/naive_vbench_policies.py new file mode 100644 index 0000000000000000000000000000000000000000..c3ace23f3f88209e66c391b2eb87ef5fce221026 --- /dev/null +++ b/scripts/naive_vbench_policies.py @@ -0,0 +1,127 @@ +#!/usr/bin/env python3 +"""Pure policy definitions for the Extended-251 naive baselines.""" + +from __future__ import annotations + +from dataclasses import asdict, dataclass +from typing import Literal + + +PolicyFamily = Literal["full", "timestep_skipping", "velocity_reuse"] + + +def raw_timestep_schedule(num_steps: int) -> tuple[int, ...]: + """Return floor-spaced training noise levels in [1000, 0).""" + if num_steps < 1: + raise ValueError("num_steps must be positive") + return tuple(1000 * (num_steps - index) // num_steps for index in range(num_steps)) + + +@dataclass(frozen=True) +class NaivePolicy: + name: str + label: str + family: PolicyFamily + raw_timesteps: tuple[int, ...] + cache_pattern: str | None = None + + @property + def denoise_steps(self) -> int: + return len(self.raw_timesteps) + + @property + def full_calls_per_chunk(self) -> int: + if self.cache_pattern is None: + return self.denoise_steps + return self.cache_pattern.count("F") + + @property + def reuse_calls_per_chunk(self) -> int: + return 0 if self.cache_pattern is None else self.cache_pattern.count("R") + + def validate(self) -> None: + if not self.name or not self.label: + raise ValueError("Policy name and label must be non-empty") + if self.family == "velocity_reuse": + if self.cache_pattern is None: + raise ValueError(f"{self.name}: velocity reuse requires a cache pattern") + if len(self.cache_pattern) != 4 or self.cache_pattern[0] != "F": + raise ValueError(f"{self.name}: cache pattern must be four steps and start with F") + if set(self.cache_pattern) - {"F", "R"}: + raise ValueError(f"{self.name}: cache pattern may contain only F and R") + if self.raw_timesteps != (1000, 750, 500, 250): + raise ValueError(f"{self.name}: cache baselines must keep the original schedule") + elif self.cache_pattern is not None: + raise ValueError(f"{self.name}: only velocity-reuse policies have cache patterns") + + def to_dict(self) -> dict[str, object]: + value = asdict(self) + value.update( + denoise_steps=self.denoise_steps, + full_calls_per_chunk=self.full_calls_per_chunk, + reuse_calls_per_chunk=self.reuse_calls_per_chunk, + ) + return value + + +ORIGINAL_RAW_TIMESTEPS = (1000, 750, 500, 250) + +REFERENCE_POLICY = NaivePolicy( + name="ffff", + label="FFFF", + family="full", + raw_timesteps=ORIGINAL_RAW_TIMESTEPS, +) + +EVALUATION_POLICIES = ( + NaivePolicy( + name="naive_3step", + label="Naive 3-step / timestep skipping", + family="timestep_skipping", + raw_timesteps=raw_timestep_schedule(3), + ), + NaivePolicy( + name="naive_2step", + label="Naive 2-step / timestep skipping", + family="timestep_skipping", + raw_timesteps=raw_timestep_schedule(2), + ), + NaivePolicy( + name="naive_1step", + label="Naive 1-step / timestep skipping", + family="timestep_skipping", + raw_timesteps=raw_timestep_schedule(1), + ), + NaivePolicy( + name="naive_cache_3step_frff", + label="Naive Cache 3-step (FRFF)", + family="velocity_reuse", + raw_timesteps=ORIGINAL_RAW_TIMESTEPS, + cache_pattern="FRFF", + ), + NaivePolicy( + name="naive_cache_2step_frrf", + label="Naive Cache 2-step (FRRF)", + family="velocity_reuse", + raw_timesteps=ORIGINAL_RAW_TIMESTEPS, + cache_pattern="FRRF", + ), + NaivePolicy( + name="naive_cache_1step_frrr", + label="Naive Cache 1-step (FRRR)", + family="velocity_reuse", + raw_timesteps=ORIGINAL_RAW_TIMESTEPS, + cache_pattern="FRRR", + ), +) + +ALL_POLICIES = (REFERENCE_POLICY, *EVALUATION_POLICIES) +POLICIES_BY_NAME = {policy.name: policy for policy in ALL_POLICIES} +EVALUATION_STRATEGY_NAMES = tuple(policy.name for policy in EVALUATION_POLICIES) +ALL_STRATEGY_NAMES = tuple(policy.name for policy in ALL_POLICIES) + +for _policy in ALL_POLICIES: + _policy.validate() +if len(POLICIES_BY_NAME) != len(ALL_POLICIES): + raise ValueError("Naive policy names must be unique") + diff --git a/scripts/plot_frrf_chunk_metrics.py b/scripts/plot_frrf_chunk_metrics.py new file mode 100644 index 0000000000000000000000000000000000000000..98e3f321fd22fc341c2accc2551b96c1b8c7cd0e --- /dev/null +++ b/scripts/plot_frrf_chunk_metrics.py @@ -0,0 +1,164 @@ +#!/usr/bin/env python3 +"""Plot 10-prompt mean FRRF metrics versus the reused chunk. + +The three input ``per_prompt.csv`` files use slightly different identifiers +(``prompt_id`` for Self/Causal-Forcing and ``case`` for WorldPlay), but share +the metric columns. We deliberately aggregate from the per-prompt rows so +that PSNR is also an arithmetic mean over the ten prompts. +""" + +from __future__ import annotations + +import argparse +import csv +from collections import defaultdict +from pathlib import Path + +import matplotlib.pyplot as plt + + +DEFAULT_ROOTS = { + "Self-Forcing": Path( + "/data3/chenzhuo/workspace/Self-Forcing/outputs/" + "single_chunk_frrf_14chunks_first10" + ), + "Causal-Forcing": Path( + "/data3/chenzhuo/workspace/Causal-Forcing/outputs/" + "single_chunk_frrf_14chunks_first10" + ), + "HY-WorldPlay": Path( + "/data3/chenzhuo/workspace/HY-WorldPlay-DEV/outputs/" + "moviebench_single_chunk_frrf_14chunks_first10" + ), +} + +METRICS = ("psnr", "ssim", "lpips") +Y_LABELS = {"psnr": "PSNR (dB)", "ssim": "SSIM", "lpips": "LPIPS"} +COLORS = { + "Self-Forcing": "#1f77b4", + "Causal-Forcing": "#d62728", + "HY-WorldPlay": "#2ca02c", +} + + +def read_prompt_means(root: Path, num_chunks: int = 14) -> dict[str, list[float]]: + """Return arithmetic means over prompts for every metric and chunk.""" + csv_path = root / "per_prompt.csv" + if not csv_path.exists(): + raise FileNotFoundError(csv_path) + + # values[chunk][metric] -> list of prompt-level values + values: dict[int, dict[str, list[float]]] = defaultdict( + lambda: {metric: [] for metric in METRICS} + ) + with csv_path.open(newline="") as handle: + reader = csv.DictReader(handle) + required = {"reuse_chunk", *METRICS} + missing = required.difference(reader.fieldnames or ()) + if missing: + raise ValueError(f"{csv_path} is missing columns: {sorted(missing)}") + for row in reader: + chunk = int(row["reuse_chunk"]) + if not 0 <= chunk < num_chunks: + raise ValueError(f"unexpected reuse_chunk={chunk} in {csv_path}") + for metric in METRICS: + values[chunk][metric].append(float(row[metric])) + + result: dict[str, list[float]] = {} + for metric in METRICS: + means = [] + for chunk in range(num_chunks): + prompt_values = values[chunk][metric] + if len(prompt_values) != 10: + raise ValueError( + f"{csv_path}: chunk {chunk} has {len(prompt_values)} rows; " + "expected 10 prompts" + ) + means.append(sum(prompt_values) / len(prompt_values)) + result[metric] = means + return result + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument( + "--output", + type=Path, + default=Path( + "/data3/chenzhuo/workspace/Self-Forcing/outputs/plots/" + "frrf_14chunks_metrics_10prompt_mean.png" + ), + help="PNG output path (a PDF with the same stem is written too).", + ) + parser.add_argument("--num-chunks", type=int, default=14) + return parser.parse_args() + + +def main() -> None: + args = parse_args() + data = { + label: read_prompt_means(root, args.num_chunks) + for label, root in DEFAULT_ROOTS.items() + } + + plt.rcParams.update( + { + "font.size": 11, + "axes.labelsize": 12, + "axes.titlesize": 13, + "legend.fontsize": 10.5, + "xtick.labelsize": 10, + "ytick.labelsize": 10, + "savefig.bbox": "tight", + } + ) + fig, axes = plt.subplots(1, 3, figsize=(15.2, 4.7), sharex=True) + chunks = list(range(args.num_chunks)) + + for axis, metric in zip(axes, METRICS): + for label, values in data.items(): + axis.plot( + chunks, + values[metric], + color=COLORS[label], + marker="o", + markersize=4.5, + linewidth=2.0, + label=label, + ) + axis.set_title(metric.upper()) + axis.set_xlabel("Reuse chunk") + axis.set_ylabel(Y_LABELS[metric]) + axis.set_xticks(chunks) + axis.grid(True, linestyle="--", linewidth=0.7, alpha=0.35) + axis.set_axisbelow(True) + axis.spines["top"].set_visible(False) + axis.spines["right"].set_visible(False) + + # One shared legend for all three panels. + handles, labels = axes[0].get_legend_handles_labels() + fig.legend( + handles, + labels, + loc="upper center", + bbox_to_anchor=(0.5, 0.995), + ncol=3, + frameon=False, + ) + fig.suptitle( + "FRRF reuse-chunk error", + y=1.045, + fontsize=14, + fontweight="semibold", + ) + fig.tight_layout(rect=(0, 0, 1, 1.0), w_pad=2.0) + + args.output.parent.mkdir(parents=True, exist_ok=True) + fig.savefig(args.output, dpi=300) + fig.savefig(args.output.with_suffix(".pdf")) + print(f"saved {args.output}") + print(f"saved {args.output.with_suffix('.pdf')}") + + +if __name__ == "__main__": + main() diff --git a/scripts/plot_self_causal_evidence.py b/scripts/plot_self_causal_evidence.py new file mode 100644 index 0000000000000000000000000000000000000000..e06fbee886ff9338e1c8ed44a22617e667c50b00 --- /dev/null +++ b/scripts/plot_self_causal_evidence.py @@ -0,0 +1,776 @@ +#!/usr/bin/env python3 +"""Create compact publication-style figures for the Self/Causal evidence chain.""" + +from __future__ import annotations + +import argparse +import colorsys +import csv +from pathlib import Path + +import matplotlib + +matplotlib.use("Agg") +import matplotlib.pyplot as plt +import numpy as np +from matplotlib.lines import Line2D +from matplotlib.patches import Patch + + +SELF = "#0F766E" +CAUSAL = "#4F46E5" +INK = "#172033" +MUTED = "#64748B" +GRID = "#DCE3EC" +PALE = "#F1F5F9" +WARM = "#D97706" +MODEL_COLOR = {"self_forcing": SELF, "causal_forcing": CAUSAL} +MODEL_LABEL = {"self_forcing": "Self-Forcing", "causal_forcing": "Causal-Forcing"} +LAYERS = ("early", "middle", "late", "final") + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser() + parser.add_argument("--redundancy_csv", type=Path, required=True) + parser.add_argument("--motion_bins_csv", type=Path, required=True) + parser.add_argument("--motion_summary_csv", type=Path, required=True) + parser.add_argument("--native_gain_csv", type=Path, required=True) + parser.add_argument( + "--causal_late_sensitivity_csv", + type=Path, + help="Optional sensitivity summary used to replace only Causal late Ridge.", + ) + parser.add_argument("--sensitivity_scenario", default="exclude_two_folds") + parser.add_argument("--aligned_probe_csv", type=Path, required=True) + parser.add_argument( + "--unified_aligned_csv", + type=Path, + help="Optional four-layer 64-D/common-mask Ridge summary used for panels D/E.", + ) + parser.add_argument("--output_dir", type=Path, required=True) + return parser.parse_args() + + +def read_csv(path: Path) -> list[dict[str, str]]: + with path.open(newline="", encoding="utf-8") as handle: + return list(csv.DictReader(handle)) + + +def apply_causal_late_sensitivity( + native_rows: list[dict[str, str]], + sensitivity_path: Path | None, + scenario: str, +) -> bool: + """Replace only the Causal late Ridge aggregate with an explicit sensitivity result.""" + if sensitivity_path is None: + return False + sensitivity_rows = read_csv(sensitivity_path) + selected = [row for row in sensitivity_rows if row["scenario"] == scenario] + if len(selected) != 1: + raise ValueError(f"Expected one sensitivity row for {scenario!r}, found {len(selected)}") + source = selected[0] + targets = [ + row + for row in native_rows + if row["method"] == "linear" + and row["model_family"] == "causal_forcing" + and row["layer_role"] == "late" + and row["probe"] == "fusion_same" + ] + if len(targets) != 1: + raise ValueError(f"Expected one Causal late Ridge row, found {len(targets)}") + target = targets[0] + for key in ("gain_mean", "gain_ci95_low", "gain_ci95_high", "wins", "signflip_p"): + target[key] = source[key] + target["prompt_count"] = source["prompt_count"] + return True + + +def lighten(color: str, amount: float = 0.45) -> str: + r, g, b = matplotlib.colors.to_rgb(color) + h, l, s = colorsys.rgb_to_hls(r, g, b) + return matplotlib.colors.to_hex(colorsys.hls_to_rgb(h, 1 - amount * (1 - l), s * 0.85)) + + +def setup_style() -> None: + plt.rcParams.update({ + "font.family": "DejaVu Sans", + "font.size": 9, + "axes.titlesize": 10.5, + "axes.labelsize": 9, + "axes.titleweight": "semibold", + "axes.labelcolor": INK, + "axes.edgecolor": GRID, + "axes.linewidth": 0.8, + "xtick.color": MUTED, + "ytick.color": MUTED, + "xtick.labelsize": 8.5, + "ytick.labelsize": 8.5, + "grid.color": GRID, + "grid.linewidth": 0.7, + "grid.alpha": 0.65, + "legend.fontsize": 7.7, + "legend.frameon": False, + "figure.facecolor": "white", + "axes.facecolor": "white", + "savefig.facecolor": "white", + "savefig.bbox": "tight", + }) + + +def polish(ax: plt.Axes, grid_axis: str = "y") -> None: + ax.spines["top"].set_visible(False) + ax.spines["right"].set_visible(False) + ax.grid(axis=grid_axis, zorder=0) + ax.tick_params(length=0) + + +def panel_label(ax: plt.Axes, label: str) -> None: + ax.text( + -0.13, + 1.08, + label, + transform=ax.transAxes, + fontsize=12, + fontweight="bold", + color=INK, + va="top", + ) + + +def lookup_rows(rows: list[dict[str, str]], keys: tuple[str, ...]) -> dict[tuple[str, ...], dict[str, str]]: + return {tuple(row[key] for key in keys): row for row in rows} + + +def plot_redundancy( + ax: plt.Axes, + rows: list[dict[str, str]], + family: str, + variant_50: str, + panel: str, +) -> None: + table = lookup_rows(rows, ("model_family", "model_variant", "layer_role", "comparison")) + x = np.arange(len(LAYERS)) + color = MODEL_COLOR[family] + series = [ + (variant_50, "within_adjacent", "50-step · denoising", "#334155", "-", "o"), + (variant_50, "cross_boundary_to_all", "50-step · boundary", "#94A3B8", "--", "D"), + ("dmd4", "within_adjacent", "4-step · denoising", color, "-", "o"), + ("dmd4", "cross_boundary_to_all", "4-step · boundary", lighten(color), "--", "D"), + ] + for variant, comparison, label, line_color, linestyle, marker in series: + values, lows, highs = [], [], [] + for role in LAYERS: + row = table[(family, variant, role, comparison)] + value = float(row["token_cosine_mean"]) + values.append(value) + lows.append(value - float(row["ci95_low"])) + highs.append(float(row["ci95_high"]) - value) + ax.errorbar( + x, + values, + yerr=np.asarray([lows, highs]), + color=line_color, + linestyle=linestyle, + linewidth=1.65, + marker=marker, + markersize=4.3, + markeredgewidth=0.7, + capsize=2, + label=label, + zorder=3, + ) + ax.axvspan(-0.35, 2.35, color=PALE, alpha=0.58, zorder=-2) + ax.set_xticks(x, [value.title() for value in LAYERS]) + ax.set_ylim(0.62, 1.015) + ax.set_yticks([0.65, 0.75, 0.85, 0.95, 1.00]) + ax.set_ylabel("Token cosine") + ax.set_title(MODEL_LABEL[family], loc="left", pad=8) + polish(ax) + panel_label(ax, panel) + + +def plot_motion(ax: plt.Axes, rows: list[dict[str, str]], panel: str) -> None: + table = lookup_rows(rows, ("model", "action", "motion_bin")) + bins = ("low", "medium", "high") + x = np.arange(3) + for family in ("self_forcing", "causal_forcing"): + color = MODEL_COLOR[family] + raw = [float(table[(family, "none", motion)]["raw_cosine"]) for motion in bins] + flow = [float(table[(family, "none", motion)]["flow_aligned_cosine"]) for motion in bins] + ax.plot(x, raw, color=color, linewidth=1.8, marker="o", markersize=4.5, zorder=3) + ax.plot(x, flow, color=color, linewidth=1.45, linestyle="--", marker="^", markersize=4.5, zorder=3) + ax.fill_between(x, raw, flow, color=color, alpha=0.09, zorder=1) + ax.annotate( + f"+{100 * (flow[-1] - raw[-1]):.2f} pts", + (x[-1], flow[-1]), + xytext=(5, 3), + textcoords="offset points", + color=color, + fontsize=7.3, + fontweight="semibold", + ) + ax.set_xticks(x, ["Low", "Medium", "High"]) + ax.set_ylim(0.875, 0.942) + ax.set_ylabel("Boundary cosine") + ax.set_title("Motion exposes spatial mismatch", loc="left", pad=8) + handles = [ + Line2D([], [], color=SELF, marker="o", label="Self · raw"), + Line2D([], [], color=SELF, marker="^", linestyle="--", label="Self · flow"), + Line2D([], [], color=CAUSAL, marker="o", label="Causal · raw"), + Line2D([], [], color=CAUSAL, marker="^", linestyle="--", label="Causal · flow"), + ] + ax.legend(handles=handles, ncol=2, loc="lower left", columnspacing=0.8, handlelength=2.0) + polish(ax) + panel_label(ax, panel) + + +def plot_native_gain( + ax: plt.Axes, + rows: list[dict[str, str]], + panel: str, + sensitivity_applied: bool = False, +) -> None: + table = lookup_rows(rows, ("method", "model_family", "layer_role", "probe")) + x = np.arange(len(LAYERS)) + offsets = {"self_forcing": -0.11, "causal_forcing": 0.11} + for family in ("self_forcing", "causal_forcing"): + color = MODEL_COLOR[family] + for method, probe, marker, filled in ( + ("linear", "fusion_same", "o", True), + ("nonlinear", "both_correct", "D", False), + ): + values, lows, highs = [], [], [] + for role in LAYERS: + row = table[(method, family, role, probe)] + value = 100 * float(row["gain_mean"]) + values.append(value) + lows.append(value - 100 * float(row["gain_ci95_low"])) + highs.append(100 * float(row["gain_ci95_high"]) - value) + ax.errorbar( + x + offsets[family], + values, + yerr=np.asarray([lows, highs]), + linestyle="none", + marker=marker, + markersize=5.2 if method == "linear" else 4.6, + markerfacecolor=color if filled else "white", + markeredgecolor=color, + markeredgewidth=1.25, + ecolor=color, + elinewidth=1.1, + capsize=2.5, + zorder=4, + ) + ax.axhline(0, color=INK, linewidth=0.8, zorder=1) + ax.set_xticks(x, [value.title() for value in LAYERS]) + ax.set_ylabel("Held-out MSE reduction (%)") + ax.set_ylim(-2.8, 5.8) + title = "Native previous-boundary adds predictive value" + if sensitivity_applied: + title += "†" + ax.set_title(title, loc="left", pad=8) + handles = [ + Line2D([], [], color=SELF, marker="o", linestyle="none", label="Self · Ridge"), + Line2D([], [], color=SELF, marker="D", markerfacecolor="white", linestyle="none", label="Self · MLP"), + Line2D([], [], color=CAUSAL, marker="o", linestyle="none", label="Causal · Ridge"), + Line2D([], [], color=CAUSAL, marker="D", markerfacecolor="white", linestyle="none", label="Causal · MLP"), + ] + ax.legend(handles=handles, ncol=2, loc="lower left", columnspacing=1.0) + ax.text( + 0.995, + 0.03, + "95% prompt-bootstrap CI", + transform=ax.transAxes, + ha="right", + color=MUTED, + fontsize=7.3, + ) + polish(ax) + panel_label(ax, panel) + + +def plot_aligned_predictor(ax: plt.Axes, rows: list[dict[str, str]], panel: str) -> None: + table = lookup_rows(rows, ("model", "probe")) + x = np.arange(2) + width = 0.29 + for index, family in enumerate(("self_forcing", "causal_forcing")): + color = MODEL_COLOR[family] + for offset, probe, label, shade in ( + (-width / 2, "both_raw", "Raw boundary", lighten(color, 0.28)), + (width / 2, "both_flow", "Flow-aligned", color), + ): + row = table[(family, probe)] + value = 100 * float(row["mse_gain_vs_step_mean"]) + low = value - 100 * float(row["mse_gain_vs_step_ci95_low"]) + high = 100 * float(row["mse_gain_vs_step_ci95_high"]) - value + ax.bar( + x[index] + offset, + value, + width=width, + color=shade, + edgecolor=color if probe == "both_raw" else "white", + linewidth=0.9, + hatch="///" if probe == "both_raw" else None, + zorder=3, + ) + ax.errorbar( + x[index] + offset, + value, + yerr=np.asarray([[low], [high]]), + fmt="none", + ecolor=INK, + elinewidth=0.9, + capsize=2.2, + zorder=4, + ) + flow = table[(family, "both_flow")] + delta = 100 * float(flow["mse_gain_vs_raw_mean"]) + ymax = max( + 100 * float(table[(family, "both_raw")]["mse_gain_vs_step_ci95_high"]), + 100 * float(flow["mse_gain_vs_step_ci95_high"]), + ) + ax.text(x[index], ymax + 0.32, f"+{delta:.2f} pp", ha="center", color=color, fontsize=8, fontweight="semibold") + ax.axhline(0, color=INK, linewidth=0.8) + ax.set_xticks(x, ["Self", "Causal"]) + ax.set_ylim(0, 9.3) + ax.set_ylabel("MSE reduction vs step-only (%)") + ax.set_title("Alignment further improves prediction", loc="left", pad=8) + ax.legend( + handles=[ + Patch(facecolor="white", edgecolor=MUTED, hatch="///", label="Raw boundary"), + Patch(facecolor="#334155", edgecolor="white", label="Flow-aligned"), + ], + loc="lower right", + ) + polish(ax) + panel_label(ax, panel) + + +def plot_unified_raw_gain(ax: plt.Axes, rows: list[dict[str, str]], panel: str) -> None: + table = lookup_rows(rows, ("model", "layer_role", "probe")) + x = np.arange(len(LAYERS)) + offsets = {"self_forcing": -0.08, "causal_forcing": 0.08} + for family in ("self_forcing", "causal_forcing"): + values, lows, highs = [], [], [] + for role in LAYERS: + row = table[(family, role, "both_raw")] + value = 100 * float(row["mse_gain_vs_step_mean"]) + values.append(value) + lows.append(value - 100 * float(row["mse_gain_vs_step_ci95_low"])) + highs.append(100 * float(row["mse_gain_vs_step_ci95_high"]) - value) + ax.errorbar( + x + offsets[family], + values, + yerr=np.asarray([lows, highs]), + color=MODEL_COLOR[family], + linestyle="none", + marker="o", + markeredgecolor="white", + markeredgewidth=0.7, + markersize=5.3, + elinewidth=1.1, + capsize=2.5, + label=MODEL_LABEL[family], + zorder=4, + ) + ax.axhline(0, color=INK, linewidth=0.8) + ax.set_xticks(x, [value.title() for value in LAYERS]) + ax.set_ylim(0, 7.25) + ax.set_ylabel("MSE reduction vs step-only (%)") + ax.set_title("Raw previous-boundary adds predictive value", loc="left", pad=8) + ax.legend(ncol=2, loc="upper left", columnspacing=1.0) + ax.text( + 0.995, + 0.03, + "95% prompt-bootstrap CI", + transform=ax.transAxes, + ha="right", + color=MUTED, + fontsize=7.3, + ) + polish(ax) + panel_label(ax, panel) + + +def plot_unified_alignment_gain(ax: plt.Axes, rows: list[dict[str, str]], panel: str) -> None: + table = lookup_rows(rows, ("model", "layer_role", "probe")) + x = np.arange(len(LAYERS)) + for family in ("self_forcing", "causal_forcing"): + values, lows, highs = [], [], [] + for role in LAYERS: + row = table[(family, role, "both_flow")] + value = 100 * float(row["mse_gain_vs_raw_mean"]) + values.append(value) + lows.append(value - 100 * float(row["mse_gain_vs_raw_ci95_low"])) + highs.append(100 * float(row["mse_gain_vs_raw_ci95_high"]) - value) + ax.errorbar( + x, + values, + yerr=np.asarray([lows, highs]), + color=MODEL_COLOR[family], + linewidth=1.55, + marker="o" if family == "self_forcing" else "D", + markersize=4.7, + markeredgecolor="white", + markeredgewidth=0.6, + capsize=2.2, + label=MODEL_LABEL[family], + zorder=4, + ) + ax.axhline(0, color=INK, linewidth=0.8) + ax.set_xticks(x, ["Early", "Mid", "Late", "Final"]) + ax.set_ylim(0, 3.0) + ax.set_ylabel("Additional MSE reduction (%)") + ax.set_title("Flow alignment adds beyond raw", loc="left", pad=8) + ax.legend(loc="upper left") + polish(ax) + panel_label(ax, panel) + + +def plot_unified_combined_gain(ax: plt.Axes, rows: list[dict[str, str]], panel: str) -> None: + """Grouped raw/flow bars under one directly comparable y-axis.""" + table = lookup_rows(rows, ("model", "layer_role", "probe")) + positions = np.asarray([0, 1, 2, 3, 5, 6, 7, 8], dtype=float) + groups = [ + (family, role) + for family in ("self_forcing", "causal_forcing") + for role in LAYERS + ] + width = 0.34 + ax.axvspan(-0.55, 3.55, color=SELF, alpha=0.035, zorder=-3) + ax.axvspan(4.45, 8.55, color=CAUSAL, alpha=0.035, zorder=-3) + for index, (family, role) in enumerate(groups): + color = MODEL_COLOR[family] + raw = table[(family, role, "both_raw")] + flow = table[(family, role, "both_flow")] + raw_value = 100 * float(raw["mse_gain_vs_step_mean"]) + flow_value = 100 * float(flow["mse_gain_vs_step_mean"]) + raw_error = np.asarray([[ + raw_value - 100 * float(raw["mse_gain_vs_step_ci95_low"]) + ], [ + 100 * float(raw["mse_gain_vs_step_ci95_high"]) - raw_value + ]]) + flow_error = np.asarray([[ + flow_value - 100 * float(flow["mse_gain_vs_step_ci95_low"]) + ], [ + 100 * float(flow["mse_gain_vs_step_ci95_high"]) - flow_value + ]]) + x = positions[index] + ax.bar( + x - width / 2, + raw_value, + width=width, + facecolor=lighten(color, 0.27), + edgecolor=color, + linewidth=0.9, + hatch="///", + zorder=3, + ) + ax.bar( + x + width / 2, + flow_value, + width=width, + color=color, + edgecolor="white", + linewidth=0.7, + zorder=3, + ) + ax.errorbar( + x - width / 2, + raw_value, + yerr=raw_error, + fmt="none", + ecolor=INK, + elinewidth=0.85, + capsize=2.1, + zorder=4, + ) + ax.errorbar( + x + width / 2, + flow_value, + yerr=flow_error, + fmt="none", + ecolor=INK, + elinewidth=0.85, + capsize=2.1, + zorder=4, + ) + delta = 100 * float(flow["mse_gain_vs_raw_mean"]) + annotation_y = 100 * float(flow["mse_gain_vs_step_ci95_high"]) + 0.18 + ax.text( + x + width / 2, + annotation_y, + f"+{delta:.2f}%", + ha="center", + va="bottom", + color=color, + fontsize=7.0, + fontweight="semibold", + ) + ax.axhline(0, color=INK, linewidth=0.8) + ax.axvline(4.0, color=GRID, linewidth=0.9) + ax.set_xticks( + positions, + [ + "Early\nSelf", "Middle\nSelf", "Late\nSelf", "Final\nSelf", + "Early\nCausal", "Middle\nCausal", "Late\nCausal", "Final\nCausal", + ], + ) + ax.set_xlim(-0.65, 8.65) + ax.set_ylim(0, 10.5) + ax.set_ylabel("MSE reduction vs step-only (%)") + ax.set_title("Raw and flow-aligned boundaries add predictive value", loc="left", pad=8) + ax.legend( + handles=[ + Patch(facecolor=lighten(SELF, 0.27), edgecolor=SELF, hatch="///", label="Self · raw"), + Patch(facecolor=SELF, edgecolor="white", label="Self · flow-aligned"), + Patch(facecolor=lighten(CAUSAL, 0.27), edgecolor=CAUSAL, hatch="///", label="Causal · raw"), + Patch(facecolor=CAUSAL, edgecolor="white", label="Causal · flow-aligned"), + ], + ncol=4, + loc="upper left", + columnspacing=1.0, + handlelength=1.6, + ) + polish(ax) + panel_label(ax, panel) + + +def main_figure(args: argparse.Namespace, output: Path) -> None: + redundancy = read_csv(args.redundancy_csv) + motion = read_csv(args.motion_bins_csv) + native = read_csv(args.native_gain_csv) + sensitivity_applied = apply_causal_late_sensitivity( + native, + args.causal_late_sensitivity_csv, + args.sensitivity_scenario, + ) + aligned = read_csv(args.aligned_probe_csv) + unified = read_csv(args.unified_aligned_csv) if args.unified_aligned_csv else None + fig = plt.figure(figsize=(13.2, 7.35), constrained_layout=True) + grid = fig.add_gridspec(2, 3, height_ratios=(1.0, 1.07), width_ratios=(1, 1, 1.03)) + ax_a = fig.add_subplot(grid[0, 0]) + ax_b = fig.add_subplot(grid[0, 1], sharey=ax_a) + ax_c = fig.add_subplot(grid[0, 2]) + if unified is None: + ax_d = fig.add_subplot(grid[1, :2]) + ax_e = fig.add_subplot(grid[1, 2]) + else: + ax_d = fig.add_subplot(grid[1, :]) + plot_redundancy(ax_a, redundancy, "self_forcing", "wan14b50", "A") + plot_redundancy(ax_b, redundancy, "causal_forcing", "ar50", "B") + ax_b.set_ylabel("") + ax_b.tick_params(labelleft=False) + ax_a.legend( + ncol=1, + loc="center left", + bbox_to_anchor=(0.015, 0.50), + handlelength=2.0, + labelspacing=0.35, + frameon=True, + facecolor="white", + edgecolor="none", + framealpha=0.88, + ) + plot_motion(ax_c, motion, "C") + if unified is None: + plot_native_gain(ax_d, native, "D", sensitivity_applied=sensitivity_applied) + plot_aligned_predictor(ax_e, aligned, "E") + else: + plot_unified_combined_gain(ax_d, unified, "D") + fig.suptitle( + "Previous-chunk boundaries retain useful information—but spatial alignment matters", + x=0.012, + ha="left", + fontsize=14, + fontweight="bold", + color=INK, + ) + footer = ( + "10 prompts; prompt-first aggregation. Error bars show 95% prompt-bootstrap CI where available. " + "Boundary = previous chunk's final temporal slot broadcast to all current slots." + ) + if unified is not None: + footer += " D: 64-D full grid/common mask; bars vs step-only; labels = flow gain over raw; 9-train/1-test; no exclusions." + elif sensitivity_applied: + footer += " † Causal late Ridge excludes prompt 6/step 1 and prompt 9/step 2 only." + fig.text( + 0.012, + -0.012, + footer, + ha="left", + fontsize=7.8, + color=MUTED, + ) + fig.savefig(output / "self_causal_evidence_overview.png", dpi=240) + fig.savefig(output / "self_causal_evidence_overview.pdf") + plt.close(fig) + + +def plot_alignment_controls(ax: plt.Axes, rows: list[dict[str, str]], panel: str) -> None: + table = lookup_rows(rows, ("model", "action")) + controls = [ + ("global_gain", "Global shift"), + ("negated_gain", "Negated flow"), + ("shuffled_gain", "Shuffled flow"), + ("flow_gain", "Correct flow"), + ] + y = np.arange(len(controls)) + ax.axhspan(2.55, 3.45, color="#ECFDF5", zorder=-2) + values_by_family = {} + for family, offset in (("self_forcing", -0.09), ("causal_forcing", 0.09)): + row = table[(family, "none")] + values = [100 * float(row[key]) for key, _ in controls] + values_by_family[family] = values + ax.scatter(values, y + offset, s=30, color=MODEL_COLOR[family], edgecolor="white", linewidth=0.6, zorder=3, label=MODEL_LABEL[family]) + for index in range(len(controls)): + ax.plot( + [values_by_family["self_forcing"][index], values_by_family["causal_forcing"][index]], + [y[index] - 0.09, y[index] + 0.09], + color="#AAB7C8", + linewidth=1.0, + zorder=1, + ) + ax.set_yticks(y, [label for _, label in controls]) + ax.set_xlabel("Cosine recovery over raw (×100)") + ax.set_xlim(0, 1.82) + ax.set_title("Interpolation controls", loc="left", pad=8) + ax.legend(loc="upper left") + polish(ax, "x") + panel_label(ax, panel) + + +def plot_predictor_controls(ax: plt.Axes, rows: list[dict[str, str]], panel: str) -> None: + table = lookup_rows(rows, ("method", "model_family", "layer_role", "probe")) + controls = [ + ("fusion_same", "Correct boundary"), + ("fusion_wrong_step", "Wrong timestep"), + ("fusion_distant", "Distant boundary"), + ("fusion_token_shuffle", "Spatial shuffle"), + ("fusion_batch_shuffle", "Other video"), + ("fusion_zero", "Zero"), + ("fusion_noise", "Matched noise"), + ] + y = np.arange(len(controls)) + ax.axhspan(-0.43, 0.43, color="#ECFDF5", zorder=-2) + for family, offset in (("self_forcing", -0.10), ("causal_forcing", 0.10)): + values, lows, highs = [], [], [] + for probe, _ in controls: + row = table[("linear", family, "final", probe)] + value = 100 * float(row["gain_mean"]) + values.append(value) + lows.append(value - 100 * float(row["gain_ci95_low"])) + highs.append(100 * float(row["gain_ci95_high"]) - value) + ax.errorbar( + values, + y + offset, + xerr=np.asarray([lows, highs]), + fmt="o", + color=MODEL_COLOR[family], + markeredgecolor="white", + markeredgewidth=0.6, + markersize=5.0, + elinewidth=1.0, + capsize=2, + label=MODEL_LABEL[family], + zorder=3, + ) + ax.axvline(0, color=INK, linewidth=0.8) + ax.set_yticks(y, [label for _, label in controls]) + ax.set_xlabel("Final-layer held-out MSE reduction (%)") + ax.set_xlim(-0.45, 4.5) + ax.set_title("Predictor controls", loc="left", pad=8) + ax.legend(loc="lower right") + ax.invert_yaxis() + polish(ax, "x") + panel_label(ax, panel) + + +def plot_unified_predictor_controls(ax: plt.Axes, rows: list[dict[str, str]], panel: str) -> None: + table = lookup_rows(rows, ("model", "layer_role", "probe")) + controls = [ + ("both_raw", "Raw boundary"), + ("both_global", "Global shift"), + ("both_negated_flow", "Negated flow"), + ("both_shuffled_flow", "Shuffled flow"), + ("both_flow", "Correct flow"), + ] + y = np.arange(len(controls)) + ax.axhspan(3.57, 4.43, color="#ECFDF5", zorder=-2) + for family, offset in (("self_forcing", -0.10), ("causal_forcing", 0.10)): + values, lows, highs = [], [], [] + for probe, _ in controls: + row = table[(family, "final", probe)] + value = 100 * float(row["mse_gain_vs_step_mean"]) + values.append(value) + lows.append(value - 100 * float(row["mse_gain_vs_step_ci95_low"])) + highs.append(100 * float(row["mse_gain_vs_step_ci95_high"]) - value) + ax.errorbar( + values, + y + offset, + xerr=np.asarray([lows, highs]), + fmt="o", + color=MODEL_COLOR[family], + markeredgecolor="white", + markeredgewidth=0.6, + markersize=5.0, + elinewidth=1.0, + capsize=2, + label=MODEL_LABEL[family], + zorder=3, + ) + ax.axvline(0, color=INK, linewidth=0.8) + ax.set_yticks(y, [label for _, label in controls]) + ax.set_xlabel("Final-layer MSE reduction vs step-only (%)") + ax.set_xlim(0, 9.5) + ax.set_title("Unified predictor controls", loc="left", pad=8) + ax.legend(loc="upper left") + ax.invert_yaxis() + polish(ax, "x") + panel_label(ax, panel) + + +def controls_figure(args: argparse.Namespace, output: Path) -> None: + motion = read_csv(args.motion_summary_csv) + native = read_csv(args.native_gain_csv) + unified = read_csv(args.unified_aligned_csv) if args.unified_aligned_csv else None + fig, axes = plt.subplots(1, 2, figsize=(11.4, 4.25), constrained_layout=True) + plot_alignment_controls(axes[0], motion, "A") + if unified is None: + plot_predictor_controls(axes[1], native, "B") + else: + plot_unified_predictor_controls(axes[1], unified, "B") + fig.suptitle( + "Control experiments isolate correct spatial correspondence", + x=0.012, + ha="left", + fontsize=13, + fontweight="bold", + color=INK, + ) + fig.text( + 0.012, + -0.02, + "Green bands mark correct-flow conditions. Predictor error bars: 95% prompt-bootstrap CI.", + ha="left", + fontsize=7.8, + color=MUTED, + ) + fig.savefig(output / "self_causal_control_checks.png", dpi=240) + fig.savefig(output / "self_causal_control_checks.pdf") + plt.close(fig) + + +def main() -> None: + args = parse_args() + setup_style() + output = args.output_dir.resolve() + output.mkdir(parents=True, exist_ok=True) + main_figure(args, output) + controls_figure(args, output) + print(f"[complete] {output}", flush=True) + + +if __name__ == "__main__": + main() diff --git a/scripts/prepare_confidence_stage1_vbench.py b/scripts/prepare_confidence_stage1_vbench.py new file mode 100644 index 0000000000000000000000000000000000000000..069cb4547782b316726f04a0544962a5238230a4 --- /dev/null +++ b/scripts/prepare_confidence_stage1_vbench.py @@ -0,0 +1,83 @@ +#!/usr/bin/env python3 +"""Prepare the representative Stage-1 confidence-gating VBench conditions.""" + +from __future__ import annotations + +import argparse +import json +import os +from pathlib import Path + + +DIMENSIONS = [ + "subject_consistency", + "background_consistency", + "motion_smoothness", + "aesthetic_quality", + "imaging_quality", +] + + +def selected_name(root: Path, target: int) -> str: + value = json.loads((root / "validation/selected.json").read_text()) + matches = [ + row["config_name"] + for row in value["selected_dynamic"] + if int(row["target_accepts"]) == target + ] + if len(matches) != 1: + raise RuntimeError(f"Expected one target={target} selection under {root}: {matches}") + return str(matches[0]) + + +def main() -> None: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--experiment_root", type=Path, required=True) + parser.add_argument("--dataset_root", type=Path, required=True) + args = parser.parse_args() + root = args.experiment_root.resolve() + step12 = root / "dynamic_step12" + step123 = root / "dynamic_step123" + conditions = { + "ffff": step123 / "test/videos/ffff", + "fppf": step12 / "test/videos/fppf", + "step12_dynamic_k06": step12 / "test/videos" / selected_name(step12, 6), + "step12_dynamic_k08": step12 / "test/videos" / selected_name(step12, 8), + "step12_dynamic_k10": step12 / "test/videos" / selected_name(step12, 10), + "step123_dynamic_k06": step123 / "test/videos" / selected_name(step123, 6), + "step123_dynamic_k09": step123 / "test/videos" / selected_name(step123, 9), + "step123_dynamic_k12": step123 / "test/videos" / selected_name(step123, 12), + "step123_dynamic_k15": step123 / "test/videos" / selected_name(step123, 15), + } + output_root = root / "vbench/inputs" + for condition, source_dir in conditions.items(): + destination = output_root / condition + destination.mkdir(parents=True, exist_ok=True) + info = [] + for prompt_id in range(90, 100): + filename = f"prompt_{prompt_id:04d}.mp4" + source = (source_dir / filename).resolve() + if not source.is_file(): + raise FileNotFoundError(source) + target = destination / filename + if target.exists() or target.is_symlink(): + target.unlink() + target.symlink_to(os.path.relpath(source, destination)) + metadata = json.loads( + (args.dataset_root / f"prompt_{prompt_id:04d}/metadata.json").read_text() + ) + info.append( + { + "video_list": [filename], + "prompt_en": metadata["prompt"], + "dimension": DIMENSIONS, + } + ) + (destination / "full_info.json").write_text( + json.dumps(info, indent=2) + "\n", encoding="utf-8" + ) + print(f"[prepared] {condition}: {source_dir}") + + +if __name__ == "__main__": + main() diff --git a/scripts/prepare_layer17_dynamic_vbench.py b/scripts/prepare_layer17_dynamic_vbench.py new file mode 100644 index 0000000000000000000000000000000000000000..24637717c18e340fbe24115f05506e37f1676368 --- /dev/null +++ b/scripts/prepare_layer17_dynamic_vbench.py @@ -0,0 +1,107 @@ +#!/usr/bin/env python3 +"""Build VBench custom-input directories from dynamic-gate video outputs.""" + +from __future__ import annotations + +import argparse +import json +import os +from pathlib import Path + + +DIMENSIONS = [ + "subject_consistency", + "background_consistency", + "motion_smoothness", + "aesthetic_quality", + "imaging_quality", +] + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--dataset_root", type=Path, required=True) + parser.add_argument("--output_root", type=Path, required=True) + parser.add_argument( + "--condition", + action="append", + nargs=2, + metavar=("NAME", "VIDEO_DIR"), + required=True, + help="Condition name and directory containing prompt_XXXX.mp4 files.", + ) + parser.add_argument("--prompt_start", type=int, default=90) + parser.add_argument("--prompt_end", type=int, default=100) + parser.add_argument( + "--combined_name", + default=None, + help="Also build one flattened input directory for a single shared VBench run.", + ) + return parser.parse_args() + + +def main() -> None: + args = parse_args() + dataset_root = args.dataset_root.resolve() + output_root = args.output_root.resolve() + output_root.mkdir(parents=True, exist_ok=True) + combined_dir = output_root / args.combined_name if args.combined_name else None + combined_info = [] + if combined_dir is not None: + combined_dir.mkdir(parents=True, exist_ok=True) + + for name, source_text in args.condition: + source_dir = Path(source_text).resolve() + destination_dir = output_root / name + destination_dir.mkdir(parents=True, exist_ok=True) + full_info = [] + for prompt_id in range(args.prompt_start, args.prompt_end): + filename = f"prompt_{prompt_id:04d}.mp4" + source = source_dir / filename + if not source.is_file(): + raise FileNotFoundError(source) + destination = destination_dir / filename + if destination.is_symlink() or destination.exists(): + destination.unlink() + destination.symlink_to(os.path.relpath(source, destination_dir)) + + metadata_path = dataset_root / f"prompt_{prompt_id:04d}" / "metadata.json" + metadata = json.loads(metadata_path.read_text(encoding="utf-8")) + full_info.append( + { + "video_list": [filename], + "prompt_en": metadata["prompt"], + "dimension": DIMENSIONS, + } + ) + if combined_dir is not None: + combined_filename = f"{name}__{filename}" + combined_destination = combined_dir / combined_filename + if combined_destination.is_symlink() or combined_destination.exists(): + combined_destination.unlink() + combined_destination.symlink_to( + os.path.relpath(source, combined_dir) + ) + combined_info.append( + { + "video_list": [combined_filename], + "prompt_en": metadata["prompt"], + "dimension": DIMENSIONS, + } + ) + (destination_dir / "full_info.json").write_text( + json.dumps(full_info, indent=2) + "\n", encoding="utf-8" + ) + print(f"[prepared] {name}: {len(full_info)} videos -> {destination_dir}") + + if combined_dir is not None: + (combined_dir / "full_info.json").write_text( + json.dumps(combined_info, indent=2) + "\n", encoding="utf-8" + ) + print( + f"[prepared] combined: {len(combined_info)} videos -> {combined_dir}" + ) + + +if __name__ == "__main__": + main() diff --git a/scripts/probe_confidence_token_batch_size.py b/scripts/probe_confidence_token_batch_size.py new file mode 100644 index 0000000000000000000000000000000000000000..a843d41b76184a094f2fd42b1e59002153edb965 --- /dev/null +++ b/scripts/probe_confidence_token_batch_size.py @@ -0,0 +1,364 @@ +#!/usr/bin/env python3 +"""Probe one real frozen-Predictor + Confidence-token-head optimizer step.""" + +from __future__ import annotations + +import argparse +import json +import os +import sys +import time +import traceback +from pathlib import Path +from typing import Any + + +def _preparse_gpu() -> str: + parser = argparse.ArgumentParser(add_help=False) + parser.add_argument("--gpu", required=True) + args, _ = parser.parse_known_args() + os.environ["CUDA_VISIBLE_DEVICES"] = str(args.gpu) + return str(args.gpu) + + +PHYSICAL_GPU = _preparse_gpu() + +import torch +import torch.nn.functional as F +from safetensors import safe_open +from safetensors.torch import load_file +from torch.optim import AdamW + +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 ConfidenceTokenHead +from predictor_training.lazy_offline_data import ( + LazyLayer17Dataset, + collate_lazy_samples, +) +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 frozen_inputs, load_teacher +from scripts.train_layer17_stage1_lazy_ddp import move_training_batch +from utils.misc import set_seed + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--gpu", default=PHYSICAL_GPU) + parser.add_argument("--batch_size", type=int, required=True) + parser.add_argument("--prompt_start", type=int, default=0) + parser.add_argument("--chunk", type=int, default=6) + parser.add_argument("--target_step", type=int, default=3) + parser.add_argument( + "--dataset_root", + type=Path, + default=REPO_ROOT + / "offline_training_datasets" + / "predictor_offline_layer17_1000p_21f_seed0_no_chunk0", + ) + parser.add_argument( + "--predictor_weights", + type=Path, + default=REPO_ROOT + / "training_runs" + / "layer17_atc_chunk_stage1_1000p_4gpu_b16_2000steps" + / "checkpoint_step_2000" + / "predictor.safetensors", + ) + parser.add_argument( + "--predictor_input_variant", + choices=("auto", "self_forcing", "disca", "atc"), + default="auto", + ) + parser.add_argument( + "--checkpoint_path", + type=Path, + default=REPO_ROOT / "checkpoints/self_forcing_dmd.pt", + ) + parser.add_argument( + "--config_path", + type=Path, + default=REPO_ROOT / "configs/self_forcing_sid.yaml", + ) + parser.add_argument("--output", type=Path, required=True) + parser.add_argument("--learning_rate", type=float, default=3e-4) + parser.add_argument("--weight_decay", type=float, default=0.01) + parser.add_argument("--seed", type=int, default=0) + args = parser.parse_args() + if args.batch_size < 1: + parser.error("--batch_size must be positive") + if args.prompt_start < 0 or args.prompt_start + args.batch_size > 900: + parser.error("probe prompts must stay in the training split 0..899") + if not 1 <= args.chunk <= 6: + parser.error("--chunk must be in 1..6") + if not 1 <= args.target_step <= 3: + parser.error("--target_step must be in 1..3") + for name in ( + "dataset_root", + "predictor_weights", + "checkpoint_path", + "config_path", + "output", + ): + path = getattr(args, name).expanduser() + setattr( + args, + name, + path.resolve() if path.is_absolute() else (REPO_ROOT / path).resolve(), + ) + return args + + +def atomic_json(path: Path, value: dict[str, 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, sort_keys=True) + "\n", encoding="utf-8" + ) + os.replace(temporary, path) + + +def read_predictor_config( + path: Path, expected_input_variant: str = "auto" +) -> dict[str, Any]: + with safe_open(path, framework="pt", device="cpu") as handle: + metadata = handle.metadata() or {} + raw = metadata.get("predictor_config") + if raw is None: + if expected_input_variant != "self_forcing": + raise ValueError( + f"Missing predictor_config metadata: {path}; legacy concat " + "checkpoints require --predictor_input_variant self_forcing" + ) + return { + "source_layer": 17, + "input_variant": "self_forcing", + "gate_mode": "baseline", + "metadata_source": "explicit_legacy_concat_override", + } + config = json.loads(raw) + actual = str(config.get("input_variant", "self_forcing")) + if expected_input_variant != "auto" and actual != expected_input_variant: + raise ValueError( + f"Predictor input variant mismatch: expected={expected_input_variant} " + f"actual={actual}" + ) + return config + + +def load_predictor( + teacher: torch.nn.Module, + weights: Path, + device: torch.device, + expected_input_variant: str = "auto", +) -> tuple[SingleBlockPredictor, dict[str, Any]]: + config = read_predictor_config(weights, expected_input_variant) + source_layer = int(config.get("source_layer", 17)) + predictor = SingleBlockPredictor( + block=initialize_predictor_block( + teacher.blocks[source_layer], "teacher_full" + ), + dim=teacher.dim, + gradient_checkpointing=False, + input_variant=str(config.get("input_variant", "self_forcing")), + atc_previous_scope=config.get("atc_previous_scope", "chunk"), + atc_freq_dim=int(config.get("atc_freq_dim", 256)), + atc_mlp_hidden_dim=int(config.get("atc_mlp_hidden_dim", 3072)), + atc_gate_hidden_dim=int(config.get("atc_gate_hidden_dim", 512)), + atc_transport_residual_scale=float( + config.get("atc_transport_residual_scale", 0.1) + ), + atc_gate_initial_probability=float( + config.get("atc_gate_initial_probability", 0.3) + ), + atc_collect_diagnostics=False, + ) + predictor.load_state_dict(load_file(str(weights), device="cpu"), strict=True) + predictor.to(device=device, dtype=torch.bfloat16) + predictor.eval().requires_grad_(False) + return predictor, config + + +@torch.no_grad() +def extract_predictor_features( + predictor: SingleBlockPredictor, + batch: dict[str, Any], + teacher: torch.nn.Module, + device: torch.device, +) -> tuple[torch.Tensor, torch.Tensor]: + frozen = frozen_inputs(batch, teacher, device) + with torch.autocast(device_type="cuda", dtype=torch.bfloat16): + output = predictor( + current_tokens=frozen["current_tokens"], + anchor_hidden=batch["anchor_hidden"], + previous_hidden=batch["previous_hidden"], + timestep_modulation=frozen["timestep_modulation"], + grid_sizes=frozen["grid_sizes"], + freqs=frozen["freqs"], + history_k=batch["history_k"], + history_v=batch["history_v"], + cross_k=batch["cross_k"], + cross_v=batch["cross_v"], + current_start=batch["chunk"] * TOKENS_PER_CHUNK, + return_features=True, + condition_tokens=frozen["condition_tokens"], + anchor_distance=batch["anchor_distance"], + ) + if not isinstance(output, tuple): + raise RuntimeError("Predictor did not return (pred_hidden, transformed)") + return output + + +def hidden_nrmse(predicted: torch.Tensor, target: torch.Tensor) -> torch.Tensor: + error_energy = (predicted.float() - target.float()).square().sum(dim=(1, 2)) + target_energy = target.float().square().sum(dim=(1, 2)) + return torch.sqrt(error_energy / target_energy.clamp_min(1e-8)) + + +def memory_gib(value: int) -> float: + return value / 2**30 + + +def main() -> None: + args = parse_args() + started = time.perf_counter() + result: dict[str, Any] = { + "success": False, + "physical_gpu": str(args.gpu), + "batch_size": args.batch_size, + "prompt_ids": [args.prompt_start, args.prompt_start + args.batch_size - 1], + "split": {"train": "0..899", "validation": "900..999"}, + "chunk": args.chunk, + "target_step": args.target_step, + "predictor_weights": str(args.predictor_weights), + } + try: + torch.cuda.set_device(0) + device = torch.device("cuda", 0) + set_seed(args.seed) + torch.set_num_threads(4) + torch.set_num_interop_threads(1) + torch.backends.cuda.matmul.allow_tf32 = True + torch.set_float32_matmul_precision("high") + + print(f"[probe] gpu={args.gpu} batch={args.batch_size} loading teacher", flush=True) + teacher = load_teacher(args.checkpoint_path, args.config_path, device) + predictor, predictor_config = load_predictor( + teacher, + args.predictor_weights, + device, + expected_input_variant=args.predictor_input_variant, + ) + result["predictor_config"] = predictor_config + head = ConfidenceTokenHead(num_steps=3, dropout=0.1).to(device=device) + result["head_parameters"] = sum( + parameter.numel() for parameter in head.parameters() + ) + optimizer = AdamW( + head.parameters(), + lr=args.learning_rate, + betas=(0.9, 0.95), + weight_decay=args.weight_decay, + ) + + print(f"[probe] gpu={args.gpu} batch={args.batch_size} loading samples", flush=True) + dataset = LazyLayer17Dataset(args.dataset_root, layer_id=17) + cpu_batch = collate_lazy_samples( + [ + dataset[(prompt_id, args.chunk, args.target_step)] + for prompt_id in range( + args.prompt_start, args.prompt_start + args.batch_size + ) + ] + ) + torch.cuda.reset_peak_memory_stats() + batch = move_training_batch(cpu_batch, teacher, device) + del cpu_batch + torch.cuda.synchronize() + feature_started = time.perf_counter() + pred_hidden, transformed = extract_predictor_features( + predictor, batch, teacher, device + ) + target_log = torch.log( + hidden_nrmse(pred_hidden, batch["target_hidden"]) + 1e-6 + ).detach() + torch.cuda.synchronize() + result["predictor_forward_s"] = time.perf_counter() - feature_started + + chunk_position = torch.full( + (args.batch_size,), + (args.chunk - 1) / 5.0, + dtype=torch.float32, + device=device, + ) + step_id = torch.full( + (args.batch_size,), + args.target_step, + dtype=torch.long, + device=device, + ) + head_started = time.perf_counter() + optimizer.zero_grad(set_to_none=True) + with torch.autocast(device_type="cuda", dtype=torch.bfloat16): + predicted_log = head( + transformed_hidden=transformed, + pred_hidden=pred_hidden, + anchor_hidden=batch["anchor_hidden"], + chunk_position=chunk_position, + step_id=step_id, + ) + loss = F.smooth_l1_loss(predicted_log, target_log) + loss.backward() + grad_norm = torch.nn.utils.clip_grad_norm_(head.parameters(), 1.0) + optimizer.step() + torch.cuda.synchronize() + + free_bytes, total_bytes = torch.cuda.mem_get_info() + result.update( + success=True, + loss=float(loss.detach()), + grad_norm=float(grad_norm), + head_step_s=time.perf_counter() - head_started, + peak_allocated_gib=memory_gib(torch.cuda.max_memory_allocated()), + peak_reserved_gib=memory_gib(torch.cuda.max_memory_reserved()), + final_free_gib=memory_gib(free_bytes), + gpu_total_gib=memory_gib(total_bytes), + elapsed_s=time.perf_counter() - started, + ) + print( + f"[probe] PASS gpu={args.gpu} batch={args.batch_size} " + f"peak={result['peak_allocated_gib']:.2f}GiB " + f"reserved={result['peak_reserved_gib']:.2f}GiB", + flush=True, + ) + except torch.cuda.OutOfMemoryError as error: + result.update( + error_type="CUDAOutOfMemoryError", + error=str(error), + peak_allocated_gib=memory_gib(torch.cuda.max_memory_allocated()), + peak_reserved_gib=memory_gib(torch.cuda.max_memory_reserved()), + elapsed_s=time.perf_counter() - started, + ) + print(f"[probe] OOM gpu={args.gpu} batch={args.batch_size}: {error}", flush=True) + except Exception as error: + result.update( + error_type=type(error).__name__, + error=str(error), + traceback=traceback.format_exc(), + elapsed_s=time.perf_counter() - started, + ) + atomic_json(args.output, result) + raise + atomic_json(args.output, result) + if not result["success"]: + raise SystemExit(2) + + +if __name__ == "__main__": + main() diff --git a/scripts/run_aligned_conditional_probe_3models.py b/scripts/run_aligned_conditional_probe_3models.py new file mode 100644 index 0000000000000000000000000000000000000000..528e31b9564a05b4fee7ae91fb377894674b73ec --- /dev/null +++ b/scripts/run_aligned_conditional_probe_3models.py @@ -0,0 +1,501 @@ +#!/usr/bin/env python3 +"""Oracle flow-aligned conditional Ridge probes on three AR4 backbones. + +The target is the current chunk/current denoising-step full-grid feature. The +first input is the current chunk/previous-step feature. The second input is +the previous chunk/same-step boundary map, either raw or warped by global, +correct, negated, or spatially shuffled target-to-source flow. + +This is an oracle diagnostic because the flow is computed from the generated +current RGB frame. It tests whether alignment makes the previous-chunk route +more predictive; it is not an inference-time implementation. +""" + +from __future__ import annotations + +import argparse +import csv +import json +import os +from collections import defaultdict +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 numpy as np +import torch +import torch.nn.functional as F + +from analyze_fullgrid_bilinear_3models import ( + GridRun, + farneback, + load_causal_runs, + load_hy_runs, + load_self_runs, + resize_flow, + shuffled_flow, + warp, +) + + +PROBES = ( + "step_only", + "both_raw", + "both_global", + "both_flow", + "both_negated_flow", + "both_shuffled_flow", +) +LAYER_ROLES = {7: "early", 14: "middle", 22: "late", 29: "final"} + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser() + parser.add_argument("--gpu", default=PHYSICAL_GPU) + parser.add_argument("--self_root", type=Path, required=True) + parser.add_argument("--causal_root", type=Path, required=True) + parser.add_argument("--hy_root", type=Path) + parser.add_argument("--hy_cache_root", type=Path) + parser.add_argument("--output_root", type=Path, required=True) + parser.add_argument("--projection_dim", type=int, default=64) + parser.add_argument("--ridge", type=float, default=1e-4) + parser.add_argument("--seed", type=int, default=20260828) + parser.add_argument( + "--multilayer_self_causal", + action="store_true", + help="Use four-layer projected full grids for Self/Causal and skip HY.", + ) + return parser.parse_args() + + +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) + path.parent.mkdir(parents=True, exist_ok=True) + with path.open("w", newline="", encoding="utf-8") as handle: + writer = csv.DictWriter(handle, fieldnames=fields) + writer.writeheader() + writer.writerows(rows) + + +def load_multilayer_self_runs(root: Path) -> dict[str, list[GridRun]]: + result = {role: [] for role in LAYER_ROLES.values()} + for path in sorted((root / "runs").glob("prompt_*.pt")): + state = torch.load(path, map_location="cpu", weights_only=False) + projected = state.get("projected_by_layer", {}) + if not projected: + raise ValueError(f"No multilayer projected features in {path}") + anchors = np.load(path.with_suffix(".anchors.npz"), allow_pickle=False)["frames"] + by_layer: dict[int, dict[tuple[int, int], torch.Tensor]] = defaultdict(dict) + for key, tensor in projected.items(): + layer, chunk, step = (int(value) for value in key.split(":")) + by_layer[layer][(chunk, step)] = tensor.float() + for layer, role in LAYER_ROLES.items(): + if layer not in by_layer: + raise ValueError(f"Missing layer {layer} in {path}") + result[role].append( + GridRun( + "self_forcing", + "none", + int(state.get("run_index", len(result[role]))), + anchors, + int(state["num_frame_per_block"]), + by_layer[layer], + path, + ) + ) + return result + + +def load_multilayer_causal_runs(root: Path) -> dict[str, list[GridRun]]: + result = {role: [] for role in LAYER_ROLES.values()} + for run_dir in sorted((root / "runs").glob("prompt_*")): + path = run_dir / "feature_snapshots.pt" + anchor_path = run_dir / "rgb_anchor_frames.npz" + if not path.exists() or not anchor_path.exists(): + continue + state = torch.load(path, map_location="cpu", weights_only=False) + projected = state.get("projected", {}) + if not projected: + raise ValueError(f"No projected features in {path}") + anchors = np.load(anchor_path, allow_pickle=False)["frames"] + by_layer: dict[int, dict[tuple[int, int], torch.Tensor]] = defaultdict(dict) + for key, tensor in projected.items(): + layer, chunk, step = (int(value) for value in key.split(":")) + by_layer[layer][(chunk, step)] = tensor.float() + for layer, role in LAYER_ROLES.items(): + if layer not in by_layer: + raise ValueError(f"Missing layer {layer} in {path}") + result[role].append( + GridRun( + "causal_forcing", + "none", + int(state["prompt_id"]), + anchors, + 3, + by_layer[layer], + path, + ) + ) + return result + + +def columns(name: str, data: dict[str, torch.Tensor]) -> list[torch.Tensor]: + ones = torch.ones_like(data["step"]) + mapping = { + "step_only": [data["step"], ones], + "both_raw": [data["step"], data["raw"], ones], + "both_global": [data["step"], data["global"], ones], + "both_flow": [data["step"], data["flow"], ones], + "both_negated_flow": [data["step"], data["negated"], ones], + "both_shuffled_flow": [data["step"], data["shuffled"], ones], + } + return mapping[name] + + +def collect_prompt_step(run: GridRun, target_step: int) -> dict[str, torch.Tensor]: + collected: dict[str, list[torch.Tensor]] = defaultdict(list) + for chunk in range(1, run.chunks): + source_frame_index = chunk * run.chunk_size - 1 + source_frame = run.anchors[source_frame_index] + source_map = run.features[(chunk - 1, target_step)][-1].float() + target_maps = run.features[(chunk, target_step)].float() + step_maps = run.features[(chunk, target_step - 1)].float() + for slot in range(run.chunk_size): + target_frame = run.anchors[chunk * run.chunk_size + slot] + flow = resize_flow(farneback(target_frame, source_frame)) + global_flow = torch.zeros_like(flow) + global_flow[0].fill_(float(torch.median(flow[0]))) + global_flow[1].fill_(float(torch.median(flow[1]))) + control_flows = { + "global": global_flow, + "flow": flow, + "negated": -flow, + "shuffled": shuffled_flow( + flow, + seed=(run.prompt_id + 1) * 100000 + + chunk * 1000 + + slot * 10 + + target_step, + ), + } + aligned: dict[str, torch.Tensor] = {} + masks: list[torch.Tensor] = [] + for name, control_flow in control_flows.items(): + aligned[name], mask = warp(source_map, control_flow) + masks.append(mask) + common_mask = torch.stack(masks).all(dim=0) + if not bool(common_mask.any()): + continue + collected["target"].append(target_maps[slot][common_mask]) + collected["step"].append(step_maps[slot][common_mask]) + collected["raw"].append(source_map[common_mask]) + for name in control_flows: + collected[name].append(aligned[name][common_mask]) + result = {key: torch.cat(values, dim=0).contiguous() for key, values in collected.items()} + expected = {"target", "step", "raw", "global", "flow", "negated", "shuffled"} + if set(result) != expected: + raise ValueError(f"Incomplete aligned sample for {run.source}: {set(result)}") + return result + + +def ridge_sufficient_statistics( + data: dict[str, torch.Tensor], + probe: str, + device: torch.device, +) -> tuple[torch.Tensor, torch.Tensor]: + design = torch.stack(columns(probe, data), dim=-1).to(device=device, dtype=torch.float64) + target = data["target"].to(device=device, dtype=torch.float64) + gram = torch.einsum("ndp,ndq->dpq", design, design) + rhs = torch.einsum("ndp,nd->dp", design, target) + return gram, rhs + + +def solve_ridge( + gram: torch.Tensor, + rhs: torch.Tensor, + ridge: float, +) -> torch.Tensor: + parameter_count = gram.shape[-1] + device = gram.device + scale = gram.diagonal(dim1=-2, dim2=-1).mean(dim=-1).clamp_min(1e-8) + regularizer = torch.eye(parameter_count, dtype=torch.float64, device=device)[None] + regularizer = regularizer * (float(ridge) * scale[:, None, None]) + regularizer[:, -1, -1] = 0.0 + try: + weights = torch.linalg.solve(gram + regularizer, rhs.unsqueeze(-1)).squeeze(-1) + except torch.linalg.LinAlgError: + weights = (torch.linalg.pinv(gram + regularizer) @ rhs.unsqueeze(-1)).squeeze(-1) + return weights.float() + + +def evaluate( + data: dict[str, torch.Tensor], + probe: str, + weights: torch.Tensor, + device: torch.device, +) -> dict[str, float]: + design = torch.stack(columns(probe, data), dim=-1).to(device=device, dtype=torch.float32) + target = data["target"].to(device=device, dtype=torch.float32) + prediction = torch.einsum("ndp,dp->nd", design, weights) + error = prediction - target + mse = error.square().mean() + variance = (target - target.mean()).square().mean().clamp_min(1e-12) + nmse = mse / variance + cosine = F.cosine_similarity(prediction, target, dim=-1, eps=1e-8).mean() + return { + "mse": float(mse), + "nMSE": float(nmse), + "nRMSE": float(torch.sqrt(nmse)), + "r2": float(1.0 - nmse), + "cosine": float(cosine), + } + + +def bootstrap(values: list[float], seed: int, rounds: int = 10000) -> tuple[float, float, float]: + array = np.asarray(values, dtype=np.float64) + generator = np.random.default_rng(seed) + indices = generator.integers(0, len(array), size=(rounds, len(array))) + means = array[indices].mean(axis=1) + return float(array.mean()), float(np.quantile(means, 0.025)), float(np.quantile(means, 0.975)) + + +def summarize(rows: list[dict[str, Any]], seed: int) -> list[dict[str, Any]]: + groups: dict[tuple[str, str, str], list[dict[str, Any]]] = defaultdict(list) + for row in rows: + groups[(row["model"], row["layer_role"], row["probe"])].append(row) + output = [] + for (model, role, probe), selected in sorted(groups.items()): + by_prompt: dict[int, list[dict[str, Any]]] = defaultdict(list) + for row in selected: + by_prompt[int(row["held_out_prompt"])].append(row) + prompt_rows = [] + for prompt_id, values in sorted(by_prompt.items()): + item = {"prompt_id": prompt_id} + for metric in ("mse", "nMSE", "nRMSE", "r2", "cosine"): + item[metric] = float(np.mean([float(row[metric]) for row in values])) + prompt_rows.append(item) + item: dict[str, Any] = { + "model": model, + "layer_role": role, + "probe": probe, + "prompt_count": len(prompt_rows), + "fold_count": len(selected), + } + for metric in ("mse", "nMSE", "nRMSE", "r2", "cosine"): + stable = seed + sum(map(ord, model + role + probe + metric)) + avg, low, high = bootstrap([row[metric] for row in prompt_rows], stable) + item[f"{metric}_mean"] = avg + item[f"{metric}_ci95_low"] = low + item[f"{metric}_ci95_high"] = high + + output.append(item) + + prompt_metric: dict[tuple[str, str, str, int], dict[str, float]] = {} + grouped: dict[tuple[str, str, str, int], list[dict[str, Any]]] = defaultdict(list) + for row in rows: + grouped[ + (row["model"], row["layer_role"], row["probe"], int(row["held_out_prompt"])) + ].append(row) + for key, values in grouped.items(): + prompt_metric[key] = { + metric: float(np.mean([float(row[metric]) for row in values])) + for metric in ("mse", "nMSE", "nRMSE", "r2", "cosine") + } + for item in output: + model, role, probe = item["model"], item["layer_role"], item["probe"] + if probe == "step_only": + continue + gain_step, gain_raw = [], [] + prompt_ids = sorted( + prompt_id + for candidate_model, candidate_role, candidate_probe, prompt_id in prompt_metric + if candidate_model == model and candidate_role == role and candidate_probe == probe + ) + for prompt_id in prompt_ids: + current = prompt_metric[(model, role, probe, prompt_id)]["mse"] + step = prompt_metric[(model, role, "step_only", prompt_id)]["mse"] + raw = prompt_metric[(model, role, "both_raw", prompt_id)]["mse"] + gain_step.append((step - current) / max(step, 1e-12)) + gain_raw.append((raw - current) / max(raw, 1e-12)) + avg, low, high = bootstrap( + gain_step, seed + 300000 + sum(map(ord, model + role + probe)) + ) + item.update({ + "mse_gain_vs_step_mean": avg, + "mse_gain_vs_step_ci95_low": low, + "mse_gain_vs_step_ci95_high": high, + "mse_gain_vs_step_wins": int(sum(value > 0 for value in gain_step)), + }) + avg, low, high = bootstrap( + gain_raw, seed + 600000 + sum(map(ord, model + role + probe)) + ) + item.update({ + "mse_gain_vs_raw_mean": avg, + "mse_gain_vs_raw_ci95_low": low, + "mse_gain_vs_raw_ci95_high": high, + "mse_gain_vs_raw_wins": int(sum(value > 0 for value in gain_raw)), + }) + if probe == "both_flow": + for baseline_probe, label in ( + ("both_global", "global"), + ("both_negated_flow", "negated_flow"), + ("both_shuffled_flow", "shuffled_flow"), + ): + gains = [] + for prompt_id in prompt_ids: + current = prompt_metric[(model, role, probe, prompt_id)]["mse"] + baseline = prompt_metric[(model, role, baseline_probe, prompt_id)]["mse"] + gains.append((baseline - current) / max(baseline, 1e-12)) + avg, low, high = bootstrap( + gains, + seed + 900000 + sum(map(ord, model + role + baseline_probe)), + ) + item.update({ + f"mse_gain_vs_{label}_mean": avg, + f"mse_gain_vs_{label}_ci95_low": low, + f"mse_gain_vs_{label}_ci95_high": high, + f"mse_gain_vs_{label}_wins": int(sum(value > 0 for value in gains)), + }) + return output + + +def main() -> None: + args = parse_args() + output = args.output_root.resolve() + output.mkdir(parents=True, exist_ok=True) + device = torch.device("cuda:0") + if not torch.cuda.is_available(): + raise RuntimeError("CUDA is required for this experiment") + + if args.multilayer_self_causal: + model_runs = { + "self_forcing": load_multilayer_self_runs(args.self_root.resolve()), + "causal_forcing": load_multilayer_causal_runs(args.causal_root.resolve()), + } + else: + if args.hy_root is None or args.hy_cache_root is None: + raise ValueError("--hy_root and --hy_cache_root are required without multilayer mode") + model_runs = { + "self_forcing": {"final": load_self_runs(args.self_root.resolve())}, + "causal_forcing": {"final": load_causal_runs(args.causal_root.resolve())}, + "hy_static": { + "final": load_hy_runs( + args.hy_root.resolve(), + "static", + args.hy_cache_root.resolve(), + args.projection_dim, + device, + False, + ) + }, + } + fold_rows: list[dict[str, Any]] = [] + for model, role_runs in model_runs.items(): + for role, runs in role_runs.items(): + if len(runs) != 10: + raise ValueError(f"Expected 10 runs for {model}/{role}, found {len(runs)}") + for target_step in range(1, 4): + prepared = [collect_prompt_step(run, target_step) for run in runs] + token_counts = [int(data["target"].shape[0]) for data in prepared] + print( + f"[prepare] {model}/{role} step={target_step} tokens={token_counts}", + flush=True, + ) + statistics = { + probe: [ridge_sufficient_statistics(data, probe, device) for data in prepared] + for probe in PROBES + } + for held_out in range(10): + test = prepared[held_out] + for probe in PROBES: + grams, right_sides = zip(*statistics[probe]) + train_gram = torch.stack(grams).sum(dim=0) - grams[held_out] + train_rhs = torch.stack(right_sides).sum(dim=0) - right_sides[held_out] + weights = solve_ridge(train_gram, train_rhs, args.ridge) + values = evaluate(test, probe, weights, device) + fold_rows.append({ + "model": model, + "layer_role": role, + "target_step": target_step, + "held_out_prompt": held_out, + "train_prompts": 9, + "test_tokens": int(test["target"].shape[0]), + "probe": probe, + **values, + }) + print( + f"[fold] {model}/{role} step={target_step} heldout={held_out}", + flush=True, + ) + del prepared + torch.cuda.empty_cache() + + summary = summarize(fold_rows, args.seed) + write_csv(output / "aligned_probe_folds.csv", fold_rows) + write_csv(output / "aligned_probe_summary.csv", summary) + summary_lookup = { + (row["model"], row["layer_role"], row["probe"]): row for row in summary + } + report = [ + "# Oracle flow-aligned conditional Ridge probe", + "", + "All gains are prompt-wise relative MSE reductions averaged over 10 held-out prompts and three target denoising steps. Flow is computed from the generated current RGB frame and is therefore an oracle diagnostic.", + "", + "| model | layer | raw chunk vs step-only | flow-aligned vs step-only | flow-aligned vs raw chunk | flow-aligned vs shuffled flow | flow-vs-raw wins |", + "|---|---|---:|---:|---:|---:|---:|", + ] + for model, role_runs in model_runs.items(): + for role in role_runs: + raw = summary_lookup[(model, role, "both_raw")] + flow = summary_lookup[(model, role, "both_flow")] + report.append( + f"| {model} | {role} | {100 * raw['mse_gain_vs_step_mean']:.2f}% | " + f"{100 * flow['mse_gain_vs_step_mean']:.2f}% | " + f"{100 * flow['mse_gain_vs_raw_mean']:.2f}% | " + f"{100 * flow['mse_gain_vs_shuffled_flow_mean']:.2f}% | " + f"{flow['mse_gain_vs_raw_wins']}/10 |" + ) + report.extend([ + "", + "Self-Forcing and Causal-Forcing are evaluated at early, middle, late, and final layers under one identical feature space, mask, split, and Ridge capacity.", + "", + "The experiment uses the common intersection of in-bounds masks for every warp, so all predictor variants see identical target tokens. Full-grid features are fixed 64-D signed random projections; conclusions concern within-model paired gains rather than native-space or cross-model absolute errors.", + ]) + (output / "REPORT.md").write_text("\n".join(report) + "\n", encoding="utf-8") + config = { + "gpu": str(args.gpu), + "models": list(model_runs), + "layer_roles": {model: list(role_runs) for model, role_runs in model_runs.items()}, + "prompt_count": 10, + "target_steps": [1, 2, 3], + "probes": list(PROBES), + "ridge": args.ridge, + "projection_dim": args.projection_dim, + "grid": [30, 52], + "split": "leave-one-prompt-out (9 train, 1 test)", + "support": "intersection of in-bounds masks for global/correct/negated/shuffled warps", + "flow": "Farneback target RGB to previous-chunk boundary RGB; oracle diagnostic", + "row_count": len(fold_rows), + } + (output / "config.json").write_text(json.dumps(config, indent=2) + "\n", encoding="utf-8") + print(f"[complete] {output} rows={len(fold_rows)}", flush=True) + + +if __name__ == "__main__": + main() diff --git a/scripts/run_conditional_probe_offline.py b/scripts/run_conditional_probe_offline.py new file mode 100644 index 0000000000000000000000000000000000000000..aaa3fc813069f81fc0ec3bd78cf2ccd097692b67 --- /dev/null +++ b/scripts/run_conditional_probe_offline.py @@ -0,0 +1,436 @@ +#!/usr/bin/env python3 +"""Run grouped, channel-wise Linear/Ridge conditional probes offline. + +The input is the normalized per-prompt dataset produced by +``build_conditional_probe_dataset.py``. All controls are derived in memory from +the same canonical features, so no model forward is repeated and every probe +uses identical target tokens. +""" + +from __future__ import annotations + +import argparse +import csv +import json +import math +import os +import sys +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) + + +_preparse_gpu() + +import numpy as np +import torch +import torch.nn.functional as F + + +FAMILIES = ("self_forcing", "causal_forcing", "hy_worldplay") +ROLES = ("early", "middle", "late", "final") +LAYER_INDICES = { + "self_forcing": {"early": 7, "middle": 14, "late": 22, "final": 29}, + "causal_forcing": {"early": 7, "middle": 14, "late": 22, "final": 29}, + "hy_worldplay": {"early": 13, "middle": 26, "late": 40, "final": 53}, +} +PROBES = ( + "within_affine", + "within_quadratic", + "cross_affine", + "fusion_same", + "fusion_step_duplicate", + "fusion_distant", + "fusion_wrong_step", + "fusion_token_shuffle", + "fusion_batch_shuffle", + "fusion_zero", + "fusion_noise", +) + + +def write_csv(path: Path, rows: list[dict[str, Any]]) -> None: + if not rows: + return + path.parent.mkdir(parents=True, exist_ok=True) + 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, extrasaction="ignore") + writer.writeheader() + writer.writerows(rows) + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser() + parser.add_argument("--dataset_root", type=Path, required=True) + parser.add_argument("--output_dir", type=Path, required=True) + parser.add_argument("--num_prompts", type=int, default=10) + parser.add_argument("--chunks", type=int, default=4) + parser.add_argument("--steps", type=int, default=4) + parser.add_argument("--ridge", type=float, default=1e-4) + parser.add_argument("--seed", type=int, default=0) + parser.add_argument( + "--chunk_pairing", + choices=("matched_slot", "boundary_to_all"), + default="matched_slot", + help=( + "How the previous-chunk auxiliary feature is paired with the current chunk. " + "boundary_to_all broadcasts the previous chunk's final temporal slot to " + "all current temporal slots at matched spatial coordinates." + ), + ) + return parser.parse_args() + + +def load_family(root: Path, family: str, count: int) -> list[dict[str, Any]]: + runs = [] + for prompt_id in range(count): + pt_path = root / family / f"prompt_{prompt_id:04d}.pt" + npz_path = root / family / f"prompt_{prompt_id:04d}.npz" + path = npz_path if family == "hy_worldplay" and not pt_path.exists() else pt_path + if not path.exists(): + raise FileNotFoundError(path) + if path.suffix == ".npz": + data = np.load(path, allow_pickle=False) + raw: dict[int, dict[tuple[int, int], torch.Tensor]] = {} + for index, stage_value in enumerate(data["stages"]): + stage = str(stage_value) + if not stage.startswith("block_"): + continue + layer = int(stage.split("_")[-1]) + chunk = int(data["chunks"][index]) + step = int(data["steps"][index]) + raw.setdefault(layer, {})[(chunk, step)] = torch.from_numpy(data["features"][index]) + layer_roles = {13: "early", 26: "middle", 40: "late", 53: "final"} + features = {} + for layer, role in layer_roles.items(): + rows = [] + for chunk in range(4): + rows.append(torch.stack([raw[layer][(chunk, step)] for step in range(4)], dim=0)) + features[role] = torch.stack(rows, dim=0).contiguous() + run = { + "prompt_id": prompt_id, + "features": features, + "timesteps": np.asarray(data["timesteps"], dtype=np.float32), + "coords": np.asarray(data["coords"], dtype=np.int64), + "grid_shape": np.asarray(data["grid_shape"], dtype=np.int64), + } + else: + run = torch.load(path, map_location="cpu", weights_only=False) + if int(run.get("prompt_id", prompt_id)) != prompt_id: + raise ValueError(f"Prompt id mismatch in {path}") + runs.append(run) + return runs + + +def boundary_to_all(reference: torch.Tensor, run: dict[str, Any]) -> torch.Tensor: + """Broadcast the last temporal slot while preserving target token order.""" + coords = np.asarray(run.get("coords")) + if coords.ndim != 2 or coords.shape[1] != 3 or len(coords) != reference.shape[0]: + raise ValueError( + "boundary_to_all requires one (temporal,y,x) coordinate per feature token; " + f"coords={coords.shape}, features={tuple(reference.shape)}" + ) + slots = sorted(int(value) for value in np.unique(coords[:, 0])) + if not slots: + raise ValueError("No temporal slots in coordinates") + last_mask = coords[:, 0] == slots[-1] + source_coords = coords[last_mask, 1:] + source = reference[torch.from_numpy(last_mask)] + result = torch.empty_like(reference) + for slot in slots: + target_mask = coords[:, 0] == slot + target_coords = coords[target_mask, 1:] + if not np.array_equal(target_coords, source_coords): + raise ValueError( + f"Temporal slot {slot} does not share the boundary slot's spatial grid" + ) + result[torch.from_numpy(target_mask)] = source + return result + + +def samples( + run: dict[str, Any], + role: str, + chunks: int, + step: int, + chunk_pairing: str = "matched_slot", +) -> dict[str, torch.Tensor]: + values = run["features"][role].float() + if values.ndim != 4: + raise ValueError(f"Expected [chunk,step,token,channel], got {values.shape}") + if values.shape[0] < chunks or values.shape[1] < 4: + raise ValueError(f"Insufficient feature grid for {role}: {values.shape}") + targets, within, cross, distant, wrong = [], [], [], [], [] + # c>=2 is required for the distant-chunk control. Pool c=2 and c=3 for + # the common four-chunk protocol, while retaining chunk_id downstream. + for chunk in range(2, chunks): + targets.append(values[chunk, step]) + within.append(values[chunk, step - 1]) + if chunk_pairing == "boundary_to_all": + cross.append(boundary_to_all(values[chunk - 1, step], run)) + distant.append(boundary_to_all(values[chunk - 2, step], run)) + wrong.append(boundary_to_all(values[chunk - 1, step - 1], run)) + else: + cross.append(values[chunk - 1, step]) + distant.append(values[chunk - 2, step]) + wrong.append(values[chunk - 1, step - 1]) + return { + "target": torch.cat(targets, dim=0), + "within": torch.cat(within, dim=0), + "cross": torch.cat(cross, dim=0), + "distant": torch.cat(distant, dim=0), + "wrong": torch.cat(wrong, dim=0), + } + + +def token_shuffle(value: torch.Tensor, spatial_count: int) -> torch.Tensor: + # Roll spatial tokens independently inside every chunk/temporal slice, + # preserving the marginal feature distribution while breaking coordinate + # correspondence. This supports both 3-slot Self/Causal and 4-slot HY. + if spatial_count <= 0 or value.shape[0] % spatial_count: + return value.roll(shifts=max(1, value.shape[0] // 2), dims=0) + frames = value.reshape(-1, spatial_count, value.shape[1]) + return frames.roll(shifts=1, dims=1).reshape_as(value) + + +def noise_like(value: torch.Tensor, seed: int) -> torch.Tensor: + generator = torch.Generator(device="cpu").manual_seed(int(seed)) + noise = torch.randn(value.shape, generator=generator, dtype=value.dtype) + return noise * value.std(dim=0, keepdim=True).clamp_min(1e-6) + value.mean(dim=0, keepdim=True) + + +def columns( + name: str, + data: dict[str, torch.Tensor], + batch: torch.Tensor, + seed: int, + spatial_count: int, +): + within, cross = data["within"], data["cross"] + ones = torch.ones_like(within) + mapping = { + "within_affine": [within, ones], + "within_quadratic": [within, within.square(), ones], + "cross_affine": [cross, ones], + "fusion_same": [within, cross, ones], + "fusion_step_duplicate": [within, within, ones], + "fusion_distant": [within, data["distant"], ones], + "fusion_wrong_step": [within, data["wrong"], ones], + "fusion_token_shuffle": [within, token_shuffle(cross, spatial_count), ones], + "fusion_batch_shuffle": [within, batch, ones], + "fusion_zero": [within, torch.zeros_like(cross), ones], + "fusion_noise": [within, noise_like(cross, seed), ones], + } + return mapping[name] + + +def fit_ridge(features: list[torch.Tensor], target: torch.Tensor, ridge: float) -> torch.Tensor: + design = torch.stack(features, dim=-1).double() # [N,D,P] + y = target.double() + gram = torch.einsum("ndp,ndq->dpq", design, design) + rhs = torch.einsum("ndp,nd->dp", design, y) + p = gram.shape[-1] + scale = gram.diagonal(dim1=-2, dim2=-1).mean(dim=-1).clamp_min(1e-8) + reg = torch.eye(p, dtype=gram.dtype).unsqueeze(0) * (float(ridge) * scale[:, None, None]) + # The last column is the explicit bias and is not regularized. + reg[:, -1, -1] = 0.0 + try: + return torch.linalg.solve(gram + reg, rhs.unsqueeze(-1)).squeeze(-1).float() + except torch.linalg.LinAlgError: + return (torch.linalg.pinv(gram + reg) @ rhs.unsqueeze(-1)).squeeze(-1).float() + + +def predict(features: list[torch.Tensor], weights: torch.Tensor) -> torch.Tensor: + return torch.einsum("ndp,dp->nd", torch.stack(features, dim=-1).float(), weights) + + +def metrics(pred: torch.Tensor, target: torch.Tensor) -> dict[str, float]: + pred, target = pred.float(), target.float() + error = pred - target + mse = error.square().mean() + centered = target - target.mean() + variance = centered.square().mean().clamp_min(1e-12) + nmse = mse / variance + cosine = F.cosine_similarity(pred.reshape(1, -1), target.reshape(1, -1), dim=1, eps=1e-8)[0] + return { + "mse": float(mse), + "nMSE": float(nmse), + "nRMSE": float(torch.sqrt(nmse)), + "r2": float(1.0 - nmse), + "cosine": float(cosine), + } + + +def bootstrap(values: list[float], seed: int, rounds: int = 4000): + values = np.asarray(values, dtype=np.float64) + rng = np.random.default_rng(seed) + if values.size == 0: + return float("nan"), float("nan"), float("nan") + draws = rng.integers(0, values.size, size=(rounds, values.size)) + means = values[draws].mean(axis=1) + return float(values.mean()), float(np.quantile(means, 0.025)), float(np.quantile(means, 0.975)) + + +def main() -> None: + args = parse_args() + args.dataset_root = args.dataset_root.resolve() + args.output_dir = args.output_dir.resolve() + args.output_dir.mkdir(parents=True, exist_ok=True) + all_rows: list[dict[str, Any]] = [] + config = { + "dataset_root": str(args.dataset_root), + "num_prompts": args.num_prompts, + "chunks": args.chunks, + "steps": args.steps, + "target_chunks": list(range(2, args.chunks)), + "ridge": args.ridge, + "chunk_pairing": args.chunk_pairing, + "outer_split": "test prompt p; donor (p+1)%N also excluded from training", + "probes": list(PROBES), + } + for family in FAMILIES: + print(f"[load] {family}", flush=True) + runs = load_family(args.dataset_root, family, args.num_prompts) + for layer_index, role in enumerate(ROLES): + coords = np.asarray(runs[0].get("coords")) + slots = sorted(int(value) for value in np.unique(coords[:, 0])) + if not slots or len(coords) != runs[0]["features"][role].shape[2]: + raise ValueError( + f"Invalid temporal coordinates for {family}/{role}: " + f"coords={coords.shape}, tokens={runs[0]['features'][role].shape[2]}" + ) + spatial_count = int((coords[:, 0] == slots[0]).sum()) + for step in range(1, args.steps): + prepared = [ + samples(run, role, args.chunks, step, args.chunk_pairing) + for run in runs + ] + for held_out in range(args.num_prompts): + donor = (held_out + 1) % args.num_prompts + train_ids = [i for i in range(args.num_prompts) if i not in {held_out, donor}] + train_data = {key: torch.cat([prepared[i][key] for i in train_ids], dim=0) for key in prepared[0]} + test_data = prepared[held_out] + donor_data = prepared[donor] + train_batch = torch.cat( + [prepared[train_ids[(position + 1) % len(train_ids)]]["cross"] + for position in range(len(train_ids))], + dim=0, + ) + for probe in PROBES: + train_cols = columns( + probe, + train_data, + train_batch, + seed=args.seed + held_out * 100 + step, + spatial_count=spatial_count, + ) + test_cols = columns( + probe, + test_data, + donor_data["cross"], + seed=args.seed + 10000 + held_out * 100 + step, + spatial_count=spatial_count, + ) + weights = fit_ridge(train_cols, train_data["target"], args.ridge) + pred = predict(test_cols, weights) + row = { + "model_family": family, + "layer_role": role, + "layer_index": LAYER_INDICES[family][role], + "target_step": step, + "held_out_prompt": held_out, + "other_video_prompt": donor, + "train_prompts": len(train_ids), + "test_tokens": int(test_data["target"].shape[0]), + "probe": probe, + **metrics(pred, test_data["target"]), + } + all_rows.append(row) + if held_out % 2 == 0: + print( + f"[progress] {family} {role} step={step} heldout={held_out}", + flush=True, + ) + + write_csv(args.output_dir / "linear_probe_folds.csv", all_rows) + summary_rows = [] + for family in FAMILIES: + for role in ROLES: + for step in range(1, args.steps): + for probe in PROBES: + selected = [ + row for row in all_rows + if row["model_family"] == family + and row["layer_role"] == role + and row["target_step"] == step + and row["probe"] == probe + ] + if not selected: + continue + item = { + "model_family": family, + "layer_role": role, + "target_step": step, + "probe": probe, + "prompt_count": len(selected), + } + for metric in ("mse", "nMSE", "nRMSE", "r2", "cosine"): + stable_seed = ( + args.seed + + 100000 * FAMILIES.index(family) + + 10000 * ROLES.index(role) + + 100 * int(step) + + sum(ord(ch) for ch in probe) + + sum(ord(ch) for ch in metric) + ) + mean, low, high = bootstrap( + [float(row[metric]) for row in selected], + stable_seed, + ) + item[f"{metric}_mean"] = mean + item[f"{metric}_ci95_low"] = low + item[f"{metric}_ci95_high"] = high + baseline = [ + row for row in all_rows + if row["model_family"] == family + and row["layer_role"] == role + and row["target_step"] == step + and row["probe"] == "within_affine" + ] + if baseline: + gains = [ + (float(base["mse"]) - float(cur["mse"])) / max(float(base["mse"]), 1e-12) + for base, cur in zip( + sorted(baseline, key=lambda row: row["held_out_prompt"]), + sorted(selected, key=lambda row: row["held_out_prompt"]), + ) + ] + mean, low, high = bootstrap(gains, args.seed + 700000 + step) + item.update({ + "gain_vs_within_affine_mean": mean, + "gain_vs_within_affine_ci95_low": low, + "gain_vs_within_affine_ci95_high": high, + "gain_vs_within_affine_wins": sum(value > 0 for value in gains), + }) + summary_rows.append(item) + write_csv(args.output_dir / "linear_probe_summary.csv", summary_rows) + config["fold_rows"] = len(all_rows) + config["summary_rows"] = len(summary_rows) + (args.output_dir / "config.json").write_text(json.dumps(config, indent=2) + "\n", encoding="utf-8") + print(f"[complete] {args.output_dir} rows={len(all_rows)}", flush=True) + + +if __name__ == "__main__": + main() diff --git a/scripts/run_single_block_init_sweep.py b/scripts/run_single_block_init_sweep.py new file mode 100644 index 0000000000000000000000000000000000000000..9a79ac02e8f2fb228829758eb8256c96b299f821 --- /dev/null +++ b/scripts/run_single_block_init_sweep.py @@ -0,0 +1,1082 @@ +#!/usr/bin/env python3 +"""Train and compare single-block Self-Forcing Predictor initializations.""" + +from __future__ import annotations + +import argparse +import csv +import hashlib +import json +import math +import os +import random +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() + # Respect torchrun's inherited multi-GPU visibility. Standalone invocations + # keep the historical GPU-2 default, while an explicit --gpu still wins. + gpu_was_explicit = any( + value == "--gpu" or value.startswith("--gpu=") for value in sys.argv[1:] + ) + if gpu_was_explicit or "CUDA_VISIBLE_DEVICES" not in os.environ: + os.environ["CUDA_VISIBLE_DEVICES"] = str(args.gpu) + return str(args.gpu) + + +PHYSICAL_GPU = _preparse_gpu() + +import torch +import torch.nn.functional as F +from omegaconf import OmegaConf +from safetensors.torch import save_file +from torch.optim import AdamW + +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 ( + OfflinePredictorStore, + TOKENS_PER_CHUNK, +) +from predictor_training.single_block import ( + SingleBlockPredictor, + TripleFeatureFusion, + initialize_predictor_block, +) +from utils.misc import set_seed +from utils.wan_wrapper import WanDiffusionWrapper +from wan.modules.model import sinusoidal_embedding_1d + + +OTHER_METHODS = ( + "random_full", + "teacher_identity", + "random_identity", + "full_zero", +) + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--gpu", default=PHYSICAL_GPU) + parser.add_argument( + "--dataset_root", + type=Path, + default=Path("outputs/predictor_offline_100_all_blocks"), + ) + parser.add_argument( + "--checkpoint_path", + type=Path, + default=Path("checkpoints/self_forcing_dmd.pt"), + ) + parser.add_argument( + "--config_path", + type=Path, + default=Path("configs/self_forcing_sid.yaml"), + ) + parser.add_argument( + "--output_dir", + type=Path, + default=Path("outputs/single_block_init_sweep"), + ) + parser.add_argument( + "--teacher_layers", + type=int, + nargs="*", + default=None, + help="Teacher source layers. Omit to sweep 0..29.", + ) + parser.add_argument("--max_steps", type=int, default=1000) + parser.add_argument("--batch_size", type=int, default=32) + parser.add_argument("--eval_batch_size", type=int, default=10) + parser.add_argument("--eval_every", type=int, default=100) + parser.add_argument("--log_every", type=int, default=20) + parser.add_argument("--save_every", type=int, default=100) + parser.add_argument("--train_prompts", type=int, default=80) + parser.add_argument("--val_prompts", type=int, default=20) + parser.add_argument("--seed", type=int, default=0) + parser.add_argument("--fusion_lr", type=float, default=1e-4) + parser.add_argument("--block_lr", type=float, default=1e-5) + parser.add_argument("--weight_decay", type=float, default=0.01) + parser.add_argument("--hidden_weight", type=float, default=0.1) + parser.add_argument("--flow_weight", type=float, default=1.0) + parser.add_argument( + "--gate_mode", choices=("baseline", "learned", "constant"), + default="baseline", + ) + parser.add_argument("--gate_hidden_dim", type=int, default=128) + parser.add_argument("--gate_initial_bias", type=float, default=4.6) + parser.add_argument("--gate_floor", type=float, default=0.0) + parser.add_argument("--gate_lr", type=float, default=None) + parser.add_argument("--gate_freeze_steps", type=int, default=0) + parser.add_argument("--constant_gate", type=float, default=1.0) + parser.add_argument("--grad_clip", type=float, default=1.0) + parser.add_argument("--fusion_warmup_steps", type=int, default=100) + parser.add_argument("--block_freeze_steps", type=int, default=100) + parser.add_argument("--block_warmup_steps", type=int, default=100) + parser.add_argument( + "--gradient_checkpointing", + action=argparse.BooleanOptionalAction, + default=True, + ) + parser.add_argument( + "--run_other_initializations", + action=argparse.BooleanOptionalAction, + default=True, + ) + parser.add_argument( + "--save_final_weights", + action=argparse.BooleanOptionalAction, + default=True, + ) + args = parser.parse_args() + if args.max_steps < 1: + parser.error("--max_steps must be positive") + if args.train_prompts < 1 or args.val_prompts < 1: + parser.error("Train and validation prompt counts must be positive") + if args.train_prompts + args.val_prompts > 100: + parser.error("The offline dataset contains 100 prompts") + if args.batch_size > args.train_prompts: + parser.error("--batch_size cannot exceed --train_prompts") + if args.eval_batch_size > args.val_prompts: + args.eval_batch_size = args.val_prompts + if not 0.0 <= args.constant_gate <= 1.0: + parser.error("--constant_gate must be in [0, 1]") + if not 0.0 <= args.gate_floor < 1.0: + parser.error("--gate_floor must be in [0, 1)") + if args.gate_freeze_steps < 0: + parser.error("--gate_freeze_steps must be non-negative") + if args.gate_lr is None: + args.gate_lr = args.fusion_lr + 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) + "\n", + encoding="utf-8", + ) + os.replace(temporary, path) + + +def append_jsonl(path: Path, value: dict[str, Any]) -> None: + path.parent.mkdir(parents=True, exist_ok=True) + with path.open("a", encoding="utf-8") as handle: + handle.write(json.dumps(value, sort_keys=True) + "\n") + + +def parameter_norm(parameters: list[torch.nn.Parameter]) -> float: + total = 0.0 + for parameter in parameters: + total += float(parameter.detach().float().square().sum()) + return math.sqrt(total) + + +def gradient_norm(parameters: list[torch.nn.Parameter]) -> float: + total = 0.0 + for parameter in parameters: + if parameter.grad is not None: + total += float(parameter.grad.detach().float().square().sum()) + return math.sqrt(total) + + +def build_shared_nonblock_state(seed: int) -> dict[str, dict[str, torch.Tensor]]: + set_seed(seed) + fusion = TripleFeatureFusion(1536) + residual = torch.nn.Linear(1536, 1536) + torch.nn.init.zeros_(residual.weight) + torch.nn.init.zeros_(residual.bias) + return { + "fusion": { + key: value.detach().clone() + for key, value in fusion.state_dict().items() + }, + "residual_out": { + key: value.detach().clone() + for key, value in residual.state_dict().items() + }, + } + + +def load_teacher( + checkpoint_path: Path, + config_path: Path, + device: torch.device, +) -> torch.nn.Module: + config = OmegaConf.merge( + OmegaConf.load(REPO_ROOT / "configs/default_config.yaml"), + OmegaConf.load(config_path), + ) + wrapper = WanDiffusionWrapper( + **getattr(config, "model_kwargs", {}), + is_causal=True, + ) + checkpoint = torch.load( + checkpoint_path, + map_location="cpu", + weights_only=False, + mmap=True, + ) + wrapper.load_state_dict(checkpoint["generator_ema"], strict=True) + del checkpoint + wrapper.to(device=device, dtype=torch.bfloat16) + wrapper.eval().requires_grad_(False) + model = wrapper.model + if len(model.blocks) != 30 or model.dim != 1536: + raise ValueError( + f"Expected Wan 1.3B with 30×1536 blocks, got " + f"{len(model.blocks)}×{model.dim}" + ) + if model.freqs.device != device: + model.freqs = model.freqs.to(device) + return model + + +class BatchSchedule: + """Precompute identical prompt/group batches for every initialization.""" + + def __init__( + self, + prompt_ids: list[int], + batch_size: int, + steps: int, + seed: int, + ) -> None: + self.entries: list[tuple[int, int, list[int]]] = [] + rng = random.Random(seed) + groups = [ + (chunk, target_step) + for chunk in range(1, 7) + for target_step in range(1, 4) + ] + while len(self.entries) < steps: + epoch_groups = groups.copy() + rng.shuffle(epoch_groups) + for chunk, target_step in epoch_groups: + prompts = rng.sample(prompt_ids, batch_size) + self.entries.append((chunk, target_step, prompts)) + if len(self.entries) == steps: + break + + def fingerprint(self) -> str: + payload = json.dumps(self.entries, separators=(",", ":")).encode() + return hashlib.sha256(payload).hexdigest() + + +def lr_values( + step: int, + max_steps: int, + fusion_lr: float, + block_lr: float, + fusion_warmup_steps: int, + block_freeze_steps: int, + block_warmup_steps: int, +) -> tuple[float, float]: + if step < fusion_warmup_steps: + fusion_factor = float(step + 1) / max(1, fusion_warmup_steps) + else: + progress = (step - fusion_warmup_steps) / max( + 1, max_steps - fusion_warmup_steps + ) + fusion_factor = 0.5 * (1.0 + math.cos(math.pi * min(1.0, progress))) + + if step < block_freeze_steps: + block_factor = 0.0 + elif step < block_freeze_steps + block_warmup_steps: + block_factor = float(step - block_freeze_steps + 1) / max( + 1, block_warmup_steps + ) + else: + progress = ( + step - block_freeze_steps - block_warmup_steps + ) / max(1, max_steps - block_freeze_steps - block_warmup_steps) + block_factor = 0.5 * ( + 1.0 + math.cos(math.pi * min(1.0, progress)) + ) + return fusion_lr * fusion_factor, block_lr * block_factor + + +def move_batch( + batch: dict[str, Any], + device: torch.device, +) -> dict[str, Any]: + output = { + key: value.to( + device=device, + dtype=( + torch.bfloat16 + if value.is_floating_point() and key != "timestep" + else value.dtype + ), + non_blocking=False, + ) + for key, value in batch.items() + if isinstance(value, torch.Tensor) + } + output.update( + { + "prompt_ids": batch["prompt_ids"], + "chunk": batch["chunk"], + "anchor_step": batch["anchor_step"], + "target_step": batch["target_step"], + } + ) + return output + + +def frozen_inputs( + batch: dict[str, Any], + teacher: torch.nn.Module, + device: torch.device, +) -> dict[str, torch.Tensor]: + noisy = batch["noisy_latent"] + batch_size = noisy.shape[0] + with torch.no_grad(), torch.autocast( + device_type="cuda", dtype=torch.bfloat16 + ): + current_tokens = teacher.patch_embedding( + noisy.permute(0, 2, 1, 3, 4) + ).flatten(2).transpose(1, 2) + timestep = batch["timestep"] + 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 + ) + tokens_per_frame = 30 * 52 + condition_tokens = ( + condition_per_frame[:, :, None, :] + .expand(batch_size, timestep.shape[1], tokens_per_frame, teacher.dim) + .reshape(batch_size, -1, teacher.dim) + ) + grid_sizes = torch.tensor( + [[3, 30, 52]] * batch_size, + dtype=torch.long, + device="cpu", + ) + return { + "current_tokens": current_tokens, + "timestep_modulation": timestep_modulation, + "head_embedding": head_embedding, + "condition_tokens": condition_tokens, + "grid_sizes": grid_sizes, + "freqs": teacher.freqs, + } + + +def hidden_to_flow( + pred_hidden: torch.Tensor, + head_embedding: torch.Tensor, + grid_sizes: torch.Tensor, + teacher: torch.nn.Module, +) -> torch.Tensor: + with torch.autocast(device_type="cuda", dtype=torch.bfloat16): + head_tokens = teacher.head(pred_hidden, head_embedding) + flow_channels_first = torch.stack( + teacher.unpatchify(head_tokens, grid_sizes) + ) + return flow_channels_first.permute(0, 2, 1, 3, 4) + + +def forward_predictor( + model: SingleBlockPredictor, + batch: dict[str, Any], + teacher: torch.nn.Module, + device: torch.device, +) -> tuple[torch.Tensor, torch.Tensor]: + frozen = frozen_inputs(batch, teacher, device) + with torch.autocast(device_type="cuda", dtype=torch.bfloat16): + pred_hidden = model( + current_tokens=frozen["current_tokens"], + anchor_hidden=batch["anchor_hidden"], + previous_hidden=batch["previous_hidden"], + timestep_modulation=frozen["timestep_modulation"], + grid_sizes=frozen["grid_sizes"], + freqs=frozen["freqs"], + history_k=batch["history_k"], + history_v=batch["history_v"], + cross_k=batch["cross_k"], + cross_v=batch["cross_v"], + current_start=batch["chunk"] * TOKENS_PER_CHUNK, + condition_tokens=frozen["condition_tokens"], + anchor_distance=batch.get("anchor_distance"), + ) + pred_flow = hidden_to_flow( + pred_hidden, + frozen["head_embedding"], + frozen["grid_sizes"], + teacher, + ) + return pred_hidden, pred_flow + + +@torch.inference_mode() +def evaluate( + model: SingleBlockPredictor, + store: OfflinePredictorStore, + val_prompt_ids: list[int], + batch_size: int, + teacher: torch.nn.Module, + device: torch.device, + hidden_weight: float, + flow_weight: float, +) -> dict[str, float]: + model.eval() + hidden_squared = 0.0 + hidden_elements = 0 + flow_squared = 0.0 + flow_elements = 0 + gate_sum = 0.0 + gate_square_sum = 0.0 + gate_count = 0 + gate_histogram = torch.zeros(10, dtype=torch.int64) + started = time.perf_counter() + for chunk in range(1, 7): + for target_step in range(1, 4): + for start in range(0, len(val_prompt_ids), batch_size): + prompt_ids = val_prompt_ids[start : start + batch_size] + batch = move_batch( + store.batch(prompt_ids, chunk, target_step), device + ) + pred_hidden, pred_flow = forward_predictor( + model, batch, teacher, device + ) + gate = model.fusion.last_gate + if gate is None: + raise RuntimeError("Fusion did not expose gate values") + gate_float = gate.float() + gate_sum += float(gate_float.sum()) + gate_square_sum += float(gate_float.square().sum()) + gate_count += gate_float.numel() + gate_histogram += torch.histc( + gate_float, bins=10, min=0.0, max=1.0 + ).to(device="cpu", dtype=torch.int64) + hidden_error = pred_hidden.float() - batch["target_hidden"].float() + flow_error = pred_flow.float() - batch["target_flow"].float() + hidden_squared += float(hidden_error.square().sum()) + hidden_elements += hidden_error.numel() + flow_squared += float(flow_error.square().sum()) + flow_elements += flow_error.numel() + del batch, pred_hidden, pred_flow, hidden_error, flow_error + hidden_mse = hidden_squared / hidden_elements + flow_mse = flow_squared / flow_elements + gate_mean = gate_sum / gate_count + gate_variance = max(0.0, gate_square_sum / gate_count - gate_mean**2) + model.train() + return { + "hidden_mse": hidden_mse, + "flow_mse": flow_mse, + "total_loss": hidden_weight * hidden_mse + flow_weight * flow_mse, + "gate_mean": gate_mean, + "gate_std": math.sqrt(gate_variance), + "gate_histogram_10_bins": gate_histogram.tolist(), + "gate_count": gate_count, + "eval_time_s": time.perf_counter() - started, + } + + +def normalized_auc( + evaluations: list[dict[str, Any]], key: str, max_steps: int +) -> float: + if len(evaluations) < 2: + return float(evaluations[0][key]) + area = 0.0 + for left, right in zip(evaluations, evaluations[1:]): + width = int(right["step"]) - int(left["step"]) + area += width * (float(left[key]) + float(right[key])) * 0.5 + return area / max_steps + + +def save_predictor_weights( + model: SingleBlockPredictor, path: Path +) -> None: + tensors = { + key: value.detach().to(device="cpu").contiguous() + for key, value in model.state_dict().items() + } + temporary = path.with_suffix(path.suffix + ".tmp") + save_file(tensors, temporary) + os.replace(temporary, path) + + +def make_model( + teacher: torch.nn.Module, + source_layer: int, + method: str, + shared_nonblock_state: dict[str, dict[str, torch.Tensor]], + seed: int, + gradient_checkpointing: bool, + device: torch.device, + gate_mode: str = "baseline", + gate_hidden_dim: int = 128, + gate_initial_bias: float = 4.6, + gate_floor: float = 0.0, + constant_gate: float = 1.0, + input_variant: str = "self_forcing", + atc_previous_scope: str = "chunk", + atc_freq_dim: int = 256, + atc_mlp_hidden_dim: int = 3072, + atc_gate_hidden_dim: int = 512, + atc_transport_residual_scale: float = 0.1, + atc_gate_initial_probability: float = 0.3, +) -> SingleBlockPredictor: + # Random block baselines use the same seed independently of the shared + # fusion initialization. + set_seed(seed) + block = initialize_predictor_block( + teacher.blocks[source_layer], + method, + ) + model = SingleBlockPredictor( + block, + dim=1536, + gradient_checkpointing=gradient_checkpointing, + input_variant=input_variant, + gate_mode=gate_mode, + gate_hidden_dim=gate_hidden_dim, + gate_initial_bias=gate_initial_bias, + gate_floor=gate_floor, + constant_gate=constant_gate, + atc_previous_scope=atc_previous_scope, + atc_freq_dim=atc_freq_dim, + atc_mlp_hidden_dim=atc_mlp_hidden_dim, + atc_gate_hidden_dim=atc_gate_hidden_dim, + atc_transport_residual_scale=atc_transport_residual_scale, + atc_gate_initial_probability=atc_gate_initial_probability, + ) + fusion_state = shared_nonblock_state["fusion"] + if input_variant == "disca": + # Keep all overlapping initialization exactly matched to the + # Self-Forcing baseline, while physically removing the third D-wide + # previous-chunk input channel from proj_in. + disca_state = { + key: value + for key, value in fusion_state.items() + if key.startswith(("current_norm.", "anchor_norm.", "proj_out.")) + } + disca_state["proj_in.weight"] = fusion_state["proj_in.weight"][ + :, : 2 * teacher.dim + ].clone() + disca_state["proj_in.bias"] = fusion_state["proj_in.bias"].clone() + model.fusion.load_state_dict(disca_state, strict=True) + elif input_variant == "self_forcing": + incompatible = model.fusion.load_state_dict(fusion_state, strict=False) + unexpected = list(incompatible.unexpected_keys) + missing = [ + key + for key in incompatible.missing_keys + if not key.startswith("gate.") + ] + if unexpected or missing: + raise RuntimeError( + f"Shared fusion state mismatch: missing={missing}, " + f"unexpected={unexpected}" + ) + elif input_variant != "atc": + raise ValueError(f"Unknown input variant: {input_variant}") + model.residual_out.load_state_dict( + shared_nonblock_state["residual_out"], strict=True + ) + return model.to(device=device) + + +def run_experiment( + *, + name: str, + method: str, + source_layer: int, + args: argparse.Namespace, + store: OfflinePredictorStore, + teacher: torch.nn.Module, + train_prompt_ids: list[int], + val_prompt_ids: list[int], + schedule: BatchSchedule, + shared_nonblock_state: dict[str, dict[str, torch.Tensor]], + device: torch.device, +) -> dict[str, Any]: + run_dir = args.output_dir / name + 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": + print(f"[run] skip completed {name}", flush=True) + return existing + run_dir.mkdir(parents=True, exist_ok=True) + log_path = run_dir / "train_log.jsonl" + + model = make_model( + teacher, + source_layer, + method, + shared_nonblock_state, + args.seed, + args.gradient_checkpointing, + device, + args.gate_mode, + args.gate_hidden_dim, + args.gate_initial_bias, + args.gate_floor, + args.constant_gate, + ) + fusion_parameters = model.fusion_parameters_without_gate() + gate_parameters = model.gate_parameters() + block_parameters = model.block_parameters() + optimizer_groups = [ + {"params": fusion_parameters, "lr": args.fusion_lr}, + ] + if gate_parameters: + optimizer_groups.append({"params": gate_parameters, "lr": args.gate_lr}) + block_group_index = len(optimizer_groups) + optimizer_groups.append({"params": block_parameters, "lr": args.block_lr}) + optimizer = AdamW( + optimizer_groups, + betas=(0.9, 0.95), + weight_decay=args.weight_decay, + ) + initial_block_norm = parameter_norm(block_parameters) + evaluations: list[dict[str, Any]] = [] + start_step = 0 + latest_path = run_dir / "training_latest.pt" + if latest_path.exists(): + state = torch.load( + latest_path, map_location="cpu", weights_only=False + ) + model.load_state_dict(state["model"], strict=True) + optimizer.load_state_dict(state["optimizer"]) + evaluations = state["evaluations"] + start_step = int(state["step"]) + print(f"[run] resume {name} at step {start_step}", flush=True) + + model.set_block_trainable(start_step >= args.block_freeze_steps) + config = { + "name": name, + "initialization_method": method, + "source_layer": source_layer, + "source_definition": ( + "block weights, clean-history K/V projection/input, and text K/V " + "all use this generator_ema Teacher layer" + ), + "seed": args.seed, + "train_prompt_ids": train_prompt_ids, + "val_prompt_ids": val_prompt_ids, + "max_steps": args.max_steps, + "batch_size": args.batch_size, + "eval_batch_size": args.eval_batch_size, + "fusion_lr": args.fusion_lr, + "block_lr": args.block_lr, + "fusion_warmup_steps": args.fusion_warmup_steps, + "block_freeze_steps": args.block_freeze_steps, + "block_warmup_steps": args.block_warmup_steps, + "weight_decay": args.weight_decay, + "hidden_weight": args.hidden_weight, + "flow_weight": args.flow_weight, + "gradient_checkpointing": args.gradient_checkpointing, + "gate_mode": args.gate_mode, + "gate_hidden_dim": args.gate_hidden_dim, + "gate_initial_bias": args.gate_initial_bias, + "gate_floor": args.gate_floor, + "gate_lr": args.gate_lr, + "gate_freeze_steps": args.gate_freeze_steps, + "constant_gate": args.constant_gate, + "batch_schedule_sha256": schedule.fingerprint(), + "initial_block_parameter_norm": initial_block_norm, + "trainable_parameters": sum( + parameter.numel() for parameter in model.parameters() + ), + } + atomic_json(run_dir / "config.json", config) + print( + f"[run] {name}: method={method}, source={source_layer}, " + f"start={start_step}", + flush=True, + ) + + if not evaluations: + initial_eval = evaluate( + model, + store, + val_prompt_ids, + args.eval_batch_size, + teacher, + device, + args.hidden_weight, + args.flow_weight, + ) + evaluations.append({"step": 0, **initial_eval}) + print( + f"[eval] {name} step=0 flow={initial_eval['flow_mse']:.8f} " + f"hidden={initial_eval['hidden_mse']:.8f}", + flush=True, + ) + + optimizer.zero_grad(set_to_none=True) + model.train() + run_started = time.perf_counter() + for step in range(start_step, args.max_steps): + block_enabled = step >= args.block_freeze_steps + if any(parameter.requires_grad != block_enabled for parameter in block_parameters): + model.set_block_trainable(block_enabled) + + fusion_lr, block_lr = lr_values( + step, + args.max_steps, + args.fusion_lr, + args.block_lr, + args.fusion_warmup_steps, + args.block_freeze_steps, + args.block_warmup_steps, + ) + optimizer.param_groups[0]["lr"] = fusion_lr + if gate_parameters: + gate_lr, _ = lr_values( + max(0, step - args.gate_freeze_steps), + max(1, args.max_steps - args.gate_freeze_steps), + args.gate_lr, args.block_lr, + args.fusion_warmup_steps, 0, args.block_warmup_steps, + ) + if step < args.gate_freeze_steps: + gate_lr = 0.0 + optimizer.param_groups[1]["lr"] = gate_lr + else: + gate_lr = 0.0 + optimizer.param_groups[block_group_index]["lr"] = block_lr + + chunk, target_step, prompt_ids = schedule.entries[step] + batch = move_batch( + store.batch(prompt_ids, chunk, target_step), device + ) + step_started = time.perf_counter() + pred_hidden, pred_flow = forward_predictor( + model, batch, teacher, device + ) + hidden_loss = F.mse_loss( + pred_hidden.float(), batch["target_hidden"].float() + ) + flow_loss = F.mse_loss( + pred_flow.float(), batch["target_flow"].float() + ) + loss = ( + args.hidden_weight * hidden_loss + + args.flow_weight * flow_loss + ) + loss.backward() + if step < args.gate_freeze_steps: + for parameter in gate_parameters: + parameter.grad = None + fusion_grad_norm = gradient_norm(fusion_parameters) + block_grad_norm = gradient_norm(block_parameters) + total_grad_norm = torch.nn.utils.clip_grad_norm_( + model.parameters(), args.grad_clip + ) + optimizer.step() + optimizer.zero_grad(set_to_none=True) + completed_step = step + 1 + + if completed_step == 1 or completed_step % args.log_every == 0: + torch.cuda.synchronize() + record = { + "step": completed_step, + "train_total_loss": float(loss.detach()), + "train_hidden_mse": float(hidden_loss.detach()), + "train_flow_mse": float(flow_loss.detach()), + "fusion_grad_norm": fusion_grad_norm, + "block_grad_norm": block_grad_norm, + "total_grad_norm_before_clip": float(total_grad_norm), + "fusion_lr": fusion_lr, + "gate_lr": gate_lr, + "block_lr": block_lr, + "chunk": chunk, + "target_step": target_step, + "step_time_s": time.perf_counter() - step_started, + "peak_gpu_gib": torch.cuda.max_memory_allocated() / (1024**3), + } + append_jsonl(log_path, record) + print( + f"[train] {name} {completed_step}/{args.max_steps} " + f"flow={record['train_flow_mse']:.8f} " + f"time={record['step_time_s']:.2f}s " + f"mem={record['peak_gpu_gib']:.1f}G", + flush=True, + ) + + should_eval = ( + completed_step % args.eval_every == 0 + or completed_step == args.max_steps + ) + if should_eval: + validation = evaluate( + model, + store, + val_prompt_ids, + args.eval_batch_size, + teacher, + device, + args.hidden_weight, + args.flow_weight, + ) + evaluations.append({"step": completed_step, **validation}) + print( + f"[eval] {name} step={completed_step} " + f"flow={validation['flow_mse']:.8f} " + f"hidden={validation['hidden_mse']:.8f}", + flush=True, + ) + + should_save = ( + completed_step % args.save_every == 0 + or completed_step == args.max_steps + ) + if should_save: + temporary = latest_path.with_suffix(".pt.tmp") + torch.save( + { + "model": { + key: value.detach().cpu() + for key, value in model.state_dict().items() + }, + "optimizer": optimizer.state_dict(), + "evaluations": evaluations, + "step": completed_step, + }, + temporary, + ) + os.replace(temporary, latest_path) + + del ( + batch, + pred_hidden, + pred_flow, + hidden_loss, + flow_loss, + loss, + total_grad_norm, + ) + + final = evaluations[-1] + result = { + "status": "complete", + **config, + "final_val_hidden_mse": final["hidden_mse"], + "final_val_flow_mse": final["flow_mse"], + "final_val_total_loss": final["total_loss"], + "val_hidden_mse_auc": normalized_auc( + evaluations, "hidden_mse", args.max_steps + ), + "val_flow_mse_auc": normalized_auc( + evaluations, "flow_mse", args.max_steps + ), + "val_total_loss_auc": normalized_auc( + evaluations, "total_loss", args.max_steps + ), + "evaluations": evaluations, + "training_time_s": time.perf_counter() - run_started, + "final_block_parameter_norm": parameter_norm(block_parameters), + } + if args.save_final_weights: + save_predictor_weights(model, run_dir / "predictor_final.safetensors") + atomic_json(metrics_path, result) + if latest_path.exists(): + latest_path.unlink() + del model, optimizer + torch.cuda.empty_cache() + return result + + +def write_summary(output_dir: Path, results: list[dict[str, Any]]) -> None: + rows = [ + { + "name": result["name"], + "initialization_method": result["initialization_method"], + "source_layer": result["source_layer"], + "final_val_flow_mse": result["final_val_flow_mse"], + "final_val_hidden_mse": result["final_val_hidden_mse"], + "final_val_total_loss": result["final_val_total_loss"], + "val_flow_mse_auc": result["val_flow_mse_auc"], + "val_hidden_mse_auc": result["val_hidden_mse_auc"], + "val_total_loss_auc": result["val_total_loss_auc"], + "training_time_s": result["training_time_s"], + } + for result in results + ] + rows.sort(key=lambda row: float(row["final_val_flow_mse"])) + atomic_json(output_dir / "summary.json", rows) + temporary = output_dir / "summary.csv.tmp" + with temporary.open("w", newline="", encoding="utf-8") as handle: + writer = csv.DictWriter(handle, fieldnames=list(rows[0])) + writer.writeheader() + writer.writerows(rows) + os.replace(temporary, output_dir / "summary.csv") + + +def main() -> None: + args = parse_args() + args.dataset_root = resolve(args.dataset_root) + args.checkpoint_path = resolve(args.checkpoint_path) + args.config_path = resolve(args.config_path) + args.output_dir = resolve(args.output_dir) + args.output_dir.mkdir(parents=True, exist_ok=True) + layers = ( + list(range(30)) + if args.teacher_layers is None or len(args.teacher_layers) == 0 + else sorted(set(args.teacher_layers)) + ) + if any(layer < 0 or layer >= 30 for layer in layers): + raise ValueError(f"Invalid Teacher layers: {layers}") + + set_seed(args.seed) + torch.backends.cuda.matmul.allow_tf32 = True + torch.backends.cudnn.allow_tf32 = True + torch.set_float32_matmul_precision("high") + device = torch.device("cuda") + + train_prompt_ids = list(range(args.train_prompts)) + val_prompt_ids = list( + range(80, 80 + args.val_prompts) + if args.train_prompts == 80 + else range(args.train_prompts, args.train_prompts + args.val_prompts) + ) + prompt_ids = sorted(set(train_prompt_ids + val_prompt_ids)) + schedule = BatchSchedule( + train_prompt_ids, + args.batch_size, + args.max_steps, + args.seed, + ) + shared_nonblock_state = build_shared_nonblock_state(args.seed) + + atomic_json( + args.output_dir / "sweep_config.json", + { + **{ + key: str(value) if isinstance(value, Path) else value + for key, value in vars(args).items() + }, + "teacher_layers": layers, + "train_prompt_ids": train_prompt_ids, + "val_prompt_ids": val_prompt_ids, + "batch_schedule_sha256": schedule.fingerprint(), + "other_initializations": list(OTHER_METHODS), + }, + ) + + print("[setup] loading frozen generator_ema", flush=True) + teacher = load_teacher( + args.checkpoint_path, + args.config_path, + device, + ) + print("[setup] loading common offline trajectories into RAM", flush=True) + store = OfflinePredictorStore(args.dataset_root, prompt_ids) + + results: list[dict[str, Any]] = [] + for layer in layers: + name = f"{args.gate_mode}_layer_{layer:02d}" + metrics_path = args.output_dir / name / "metrics.json" + if metrics_path.exists(): + existing = json.loads(metrics_path.read_text(encoding="utf-8")) + if existing.get("status") == "complete": + results.append(existing) + print(f"[sweep] already complete: {name}", flush=True) + continue + store.load_layer_cache(layer, teacher, device) + results.append( + run_experiment( + name=name, + method="teacher_full", + source_layer=layer, + args=args, + store=store, + teacher=teacher, + train_prompt_ids=train_prompt_ids, + val_prompt_ids=val_prompt_ids, + schedule=schedule, + shared_nonblock_state=shared_nonblock_state, + device=device, + ) + ) + write_summary(args.output_dir, results) + + teacher_results = [ + result + for result in results + if result["initialization_method"] == "teacher_full" + ] + if not teacher_results: + raise RuntimeError("No Teacher layer experiment completed") + best_teacher = min( + teacher_results, + key=lambda result: float(result["final_val_flow_mse"]), + ) + best_layer = int(best_teacher["source_layer"]) + atomic_json( + args.output_dir / "best_teacher_layer.json", + { + "source_layer": best_layer, + "selection_metric": "final_val_flow_mse", + "value": best_teacher["final_val_flow_mse"], + "run": best_teacher["name"], + }, + ) + print( + f"[sweep] best Teacher layer={best_layer}, " + f"flow={best_teacher['final_val_flow_mse']:.8f}", + flush=True, + ) + + if args.run_other_initializations: + store.load_layer_cache(best_layer, teacher, device) + for method in OTHER_METHODS: + name = f"{method}_source_{best_layer:02d}" + results.append( + run_experiment( + name=name, + method=method, + source_layer=best_layer, + args=args, + store=store, + teacher=teacher, + train_prompt_ids=train_prompt_ids, + val_prompt_ids=val_prompt_ids, + schedule=schedule, + shared_nonblock_state=shared_nonblock_state, + device=device, + ) + ) + write_summary(args.output_dir, results) + + write_summary(args.output_dir, results) + print( + f"[sweep] complete: {len(results)} experiments -> " + f"{args.output_dir / 'summary.csv'}", + flush=True, + ) + + +if __name__ == "__main__": + main() diff --git a/scripts/run_three_block_consecutive_sweep.py b/scripts/run_three_block_consecutive_sweep.py new file mode 100644 index 0000000000000000000000000000000000000000..f844518caebfb58ab37b37ae7b1de23e208d135b --- /dev/null +++ b/scripts/run_three_block_consecutive_sweep.py @@ -0,0 +1,651 @@ +#!/usr/bin/env python3 +"""Train three-block Predictors initialized from consecutive Teacher layers.""" + +# ruff: noqa: E402 -- CUDA_VISIBLE_DEVICES must be set before importing torch. + +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 torch +import torch.nn.functional as F +from safetensors.torch import save_file +from torch.optim import AdamW + +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 OfflinePredictorStore, TOKENS_PER_CHUNK +from predictor_training.three_block import ThreeBlockPredictor +from scripts.run_single_block_init_sweep import ( + BatchSchedule, + append_jsonl, + atomic_json, + build_shared_nonblock_state, + frozen_inputs, + gradient_norm, + hidden_to_flow, + load_teacher, + lr_values, + move_batch, + normalized_auc, + parameter_norm, +) +from utils.misc import set_seed + + +Triple = tuple[int, int, int] + + +def default_triples() -> list[Triple]: + return [(index, index + 1, index + 2) for index in range(28)] + + +def parse_triple(value: str) -> Triple: + try: + values = tuple(int(item) for item in value.split(",")) + except ValueError as error: + raise argparse.ArgumentTypeError( + f"Triple must look like 0,1,2; got {value!r}" + ) from error + if len(values) != 3: + raise argparse.ArgumentTypeError( + f"Triple must contain exactly three layers; got {value!r}" + ) + triple = (values[0], values[1], values[2]) + if not ( + 0 <= triple[0] + and triple[1] == triple[0] + 1 + and triple[2] == triple[1] + 1 + and triple[2] < 30 + ): + raise argparse.ArgumentTypeError( + "Triple must be three consecutive layers within 0..29; " + f"got {triple}" + ) + return triple + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--gpu", default=PHYSICAL_GPU) + parser.add_argument( + "--dataset_root", + type=Path, + default=Path("outputs/predictor_offline_100_all_blocks"), + ) + parser.add_argument( + "--checkpoint_path", + type=Path, + default=Path("checkpoints/self_forcing_dmd.pt"), + ) + parser.add_argument( + "--config_path", type=Path, default=Path("configs/self_forcing_sid.yaml") + ) + parser.add_argument( + "--output_dir", + type=Path, + default=Path("outputs/three_block_consecutive_sweep"), + ) + parser.add_argument( + "--triples", + type=parse_triple, + nargs="*", + default=None, + help=( + "Optional subset such as 0,1,2 1,2,3. " + "Omit to train all 28 consecutive triples." + ), + ) + parser.add_argument("--max_triples", type=int, default=None) + parser.add_argument("--max_steps", type=int, default=1000) + parser.add_argument("--batch_size", type=int, default=32) + parser.add_argument("--eval_batch_size", type=int, default=10) + parser.add_argument("--eval_every", type=int, default=100) + parser.add_argument("--log_every", type=int, default=20) + parser.add_argument("--save_every", type=int, default=100) + parser.add_argument("--train_prompts", type=int, default=80) + parser.add_argument("--val_prompts", type=int, default=20) + parser.add_argument("--seed", type=int, default=0) + parser.add_argument("--fusion_lr", type=float, default=1e-4) + parser.add_argument("--block_lr", type=float, default=1e-5) + parser.add_argument("--weight_decay", type=float, default=0.01) + parser.add_argument("--hidden_weight", type=float, default=0.1) + parser.add_argument("--flow_weight", type=float, default=1.0) + parser.add_argument("--grad_clip", type=float, default=1.0) + parser.add_argument("--fusion_warmup_steps", type=int, default=100) + parser.add_argument("--block_freeze_steps", type=int, default=100) + parser.add_argument("--block_warmup_steps", type=int, default=100) + parser.add_argument( + "--gradient_checkpointing", + action=argparse.BooleanOptionalAction, + default=True, + ) + parser.add_argument( + "--save_final_weights", + action=argparse.BooleanOptionalAction, + default=True, + ) + args = parser.parse_args() + if args.max_steps < 1: + parser.error("--max_steps must be positive") + if args.train_prompts < 1 or args.val_prompts < 1: + parser.error("Prompt counts must be positive") + if args.train_prompts + args.val_prompts > 100: + parser.error("The offline dataset contains 100 prompts") + if args.batch_size > args.train_prompts: + parser.error("--batch_size cannot exceed --train_prompts") + if args.eval_batch_size > args.val_prompts: + args.eval_batch_size = args.val_prompts + if args.max_triples is not None and args.max_triples < 1: + parser.error("--max_triples 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 experiment_name(triple: Triple) -> str: + return "triple_" + "_".join(f"{layer:02d}" for layer in triple) + + +def make_model( + teacher: torch.nn.Module, + triple: Triple, + shared_nonblock_state: dict[str, dict[str, torch.Tensor]], + seed: int, + gradient_checkpointing: bool, + device: torch.device, +) -> ThreeBlockPredictor: + set_seed(seed) + model = ThreeBlockPredictor( + [teacher.blocks[layer] for layer in triple], + dim=teacher.dim, + gradient_checkpointing=gradient_checkpointing, + ) + model.fusion.load_state_dict(shared_nonblock_state["fusion"], strict=True) + model.residual_out.load_state_dict( + shared_nonblock_state["residual_out"], strict=True + ) + return model.to(device=device) + + +def forward_predictor( + model: ThreeBlockPredictor, + batch: dict[str, Any], + teacher: torch.nn.Module, + device: torch.device, +) -> tuple[torch.Tensor, torch.Tensor]: + frozen = frozen_inputs(batch, teacher, device) + with torch.autocast(device_type="cuda", dtype=torch.bfloat16): + pred_hidden = model( + current_tokens=frozen["current_tokens"], + anchor_hidden=batch["anchor_hidden"], + previous_hidden=batch["previous_hidden"], + timestep_modulation=frozen["timestep_modulation"], + grid_sizes=frozen["grid_sizes"], + freqs=frozen["freqs"], + history_ks=[batch[f"history_k_{index}"] for index in range(3)], + history_vs=[batch[f"history_v_{index}"] for index in range(3)], + cross_ks=[batch[f"cross_k_{index}"] for index in range(3)], + cross_vs=[batch[f"cross_v_{index}"] for index in range(3)], + current_start=batch["chunk"] * TOKENS_PER_CHUNK, + ) + pred_flow = hidden_to_flow( + pred_hidden, + frozen["head_embedding"], + frozen["grid_sizes"], + teacher, + ) + return pred_hidden, pred_flow + + +@torch.inference_mode() +def evaluate( + model: ThreeBlockPredictor, + store: OfflinePredictorStore, + triple: Triple, + val_prompt_ids: list[int], + batch_size: int, + teacher: torch.nn.Module, + device: torch.device, + hidden_weight: float, + flow_weight: float, +) -> dict[str, float]: + model.eval() + hidden_squared = 0.0 + hidden_elements = 0 + flow_squared = 0.0 + flow_elements = 0 + started = time.perf_counter() + for chunk in range(1, 7): + for target_step in range(1, 4): + for start in range(0, len(val_prompt_ids), batch_size): + prompt_ids = val_prompt_ids[start : start + batch_size] + batch = move_batch( + store.batch_layers(prompt_ids, chunk, target_step, triple), + device, + ) + pred_hidden, pred_flow = forward_predictor( + model, batch, teacher, device + ) + hidden_error = pred_hidden.float() - batch["target_hidden"].float() + flow_error = pred_flow.float() - batch["target_flow"].float() + hidden_squared += float(hidden_error.square().sum()) + hidden_elements += hidden_error.numel() + flow_squared += float(flow_error.square().sum()) + flow_elements += flow_error.numel() + del batch, pred_hidden, pred_flow, hidden_error, flow_error + hidden_mse = hidden_squared / hidden_elements + flow_mse = flow_squared / flow_elements + model.train() + return { + "hidden_mse": hidden_mse, + "flow_mse": flow_mse, + "total_loss": hidden_weight * hidden_mse + flow_weight * flow_mse, + "eval_time_s": time.perf_counter() - started, + } + + +def save_predictor_weights(model: ThreeBlockPredictor, path: Path) -> None: + tensors = { + key: value.detach().to(device="cpu").contiguous() + for key, value in model.state_dict().items() + } + temporary = path.with_suffix(path.suffix + ".tmp") + save_file(tensors, temporary) + os.replace(temporary, path) + + +def run_experiment( + *, + triple: Triple, + args: argparse.Namespace, + store: OfflinePredictorStore, + teacher: torch.nn.Module, + train_prompt_ids: list[int], + val_prompt_ids: list[int], + schedule: BatchSchedule, + shared_nonblock_state: dict[str, dict[str, torch.Tensor]], + device: torch.device, +) -> dict[str, Any]: + name = experiment_name(triple) + run_dir = args.output_dir / name + 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": + print(f"[run] skip completed {name}", flush=True) + return existing + run_dir.mkdir(parents=True, exist_ok=True) + log_path = run_dir / "train_log.jsonl" + + model = make_model( + teacher, + triple, + shared_nonblock_state, + args.seed, + args.gradient_checkpointing, + device, + ) + fusion_parameters = model.fusion_parameters() + block_parameters = model.block_parameters() + optimizer = AdamW( + [ + {"params": fusion_parameters, "lr": args.fusion_lr}, + {"params": block_parameters, "lr": args.block_lr}, + ], + betas=(0.9, 0.95), + weight_decay=args.weight_decay, + ) + initial_block_norm = parameter_norm(block_parameters) + evaluations: list[dict[str, Any]] = [] + start_step = 0 + latest_path = run_dir / "training_latest.pt" + if latest_path.exists(): + state = torch.load(latest_path, map_location="cpu", weights_only=False) + model.load_state_dict(state["model"], strict=True) + optimizer.load_state_dict(state["optimizer"]) + evaluations = state["evaluations"] + start_step = int(state["step"]) + print(f"[run] resume {name} at step {start_step}", flush=True) + + model.set_blocks_trainable(start_step >= args.block_freeze_steps) + config = { + "name": name, + "architecture": "three_block_predictor", + "initialization_method": "teacher_full", + "source_layers": list(triple), + "triple_kind": "consecutive", + "source_definition": ( + "Predictor block i, its clean-history K/V, and its text K/V use " + "generator_ema source_layers[i]" + ), + "seed": args.seed, + "train_prompt_ids": train_prompt_ids, + "val_prompt_ids": val_prompt_ids, + "max_steps": args.max_steps, + "batch_size": args.batch_size, + "eval_batch_size": args.eval_batch_size, + "fusion_lr": args.fusion_lr, + "block_lr": args.block_lr, + "fusion_warmup_steps": args.fusion_warmup_steps, + "block_freeze_steps": args.block_freeze_steps, + "block_warmup_steps": args.block_warmup_steps, + "weight_decay": args.weight_decay, + "hidden_weight": args.hidden_weight, + "flow_weight": args.flow_weight, + "gradient_checkpointing": args.gradient_checkpointing, + "batch_schedule_sha256": schedule.fingerprint(), + "initial_block_parameter_norm": initial_block_norm, + "trainable_parameters": sum( + parameter.numel() for parameter in model.parameters() + ), + } + atomic_json(run_dir / "config.json", config) + print(f"[run] {name}: start={start_step}", flush=True) + + if not evaluations: + initial_eval = evaluate( + model, + store, + triple, + val_prompt_ids, + args.eval_batch_size, + teacher, + device, + args.hidden_weight, + args.flow_weight, + ) + evaluations.append({"step": 0, **initial_eval}) + print( + f"[eval] {name} step=0 flow={initial_eval['flow_mse']:.8f} " + f"hidden={initial_eval['hidden_mse']:.8f}", + flush=True, + ) + + optimizer.zero_grad(set_to_none=True) + model.train() + run_started = time.perf_counter() + for step in range(start_step, args.max_steps): + block_enabled = step >= args.block_freeze_steps + if any( + parameter.requires_grad != block_enabled + for parameter in block_parameters + ): + model.set_blocks_trainable(block_enabled) + + fusion_lr, block_lr = lr_values( + step, + args.max_steps, + args.fusion_lr, + args.block_lr, + args.fusion_warmup_steps, + args.block_freeze_steps, + args.block_warmup_steps, + ) + optimizer.param_groups[0]["lr"] = fusion_lr + optimizer.param_groups[1]["lr"] = block_lr + + chunk, target_step, prompt_ids = schedule.entries[step] + batch = move_batch( + store.batch_layers(prompt_ids, chunk, target_step, triple), device + ) + step_started = time.perf_counter() + pred_hidden, pred_flow = forward_predictor(model, batch, teacher, device) + hidden_loss = F.mse_loss( + pred_hidden.float(), batch["target_hidden"].float() + ) + flow_loss = F.mse_loss(pred_flow.float(), batch["target_flow"].float()) + loss = args.hidden_weight * hidden_loss + args.flow_weight * flow_loss + loss.backward() + fusion_grad_norm = gradient_norm(fusion_parameters) + block_grad_norm = gradient_norm(block_parameters) + per_block_grad_norms = [ + gradient_norm(list(block.parameters())) for block in model.blocks + ] + total_grad_norm = torch.nn.utils.clip_grad_norm_( + model.parameters(), args.grad_clip + ) + optimizer.step() + optimizer.zero_grad(set_to_none=True) + completed_step = step + 1 + + if completed_step == 1 or completed_step % args.log_every == 0: + torch.cuda.synchronize() + record = { + "step": completed_step, + "train_total_loss": float(loss.detach()), + "train_hidden_mse": float(hidden_loss.detach()), + "train_flow_mse": float(flow_loss.detach()), + "fusion_grad_norm": fusion_grad_norm, + "block_grad_norm": block_grad_norm, + "total_grad_norm_before_clip": float(total_grad_norm), + "fusion_lr": fusion_lr, + "block_lr": block_lr, + "chunk": chunk, + "target_step": target_step, + "step_time_s": time.perf_counter() - step_started, + "peak_gpu_gib": torch.cuda.max_memory_allocated() / (1024**3), + } + record.update( + { + f"block_{index}_grad_norm": value + for index, value in enumerate(per_block_grad_norms) + } + ) + append_jsonl(log_path, record) + print( + f"[train] {name} {completed_step}/{args.max_steps} " + f"flow={record['train_flow_mse']:.8f} " + f"time={record['step_time_s']:.2f}s " + f"mem={record['peak_gpu_gib']:.1f}G", + flush=True, + ) + + if completed_step % args.eval_every == 0 or completed_step == args.max_steps: + validation = evaluate( + model, + store, + triple, + val_prompt_ids, + args.eval_batch_size, + teacher, + device, + args.hidden_weight, + args.flow_weight, + ) + evaluations.append({"step": completed_step, **validation}) + print( + f"[eval] {name} step={completed_step} " + f"flow={validation['flow_mse']:.8f} " + f"hidden={validation['hidden_mse']:.8f}", + flush=True, + ) + + if completed_step % args.save_every == 0 or completed_step == args.max_steps: + temporary = latest_path.with_suffix(".pt.tmp") + torch.save( + { + "model": { + key: value.detach().cpu() + for key, value in model.state_dict().items() + }, + "optimizer": optimizer.state_dict(), + "evaluations": evaluations, + "step": completed_step, + }, + temporary, + ) + os.replace(temporary, latest_path) + + del batch, pred_hidden, pred_flow, hidden_loss, flow_loss, loss + del total_grad_norm + + final = evaluations[-1] + result = { + "status": "complete", + **config, + "final_val_hidden_mse": final["hidden_mse"], + "final_val_flow_mse": final["flow_mse"], + "final_val_total_loss": final["total_loss"], + "val_hidden_mse_auc": normalized_auc( + evaluations, "hidden_mse", args.max_steps + ), + "val_flow_mse_auc": normalized_auc( + evaluations, "flow_mse", args.max_steps + ), + "val_total_loss_auc": normalized_auc( + evaluations, "total_loss", args.max_steps + ), + "evaluations": evaluations, + "training_time_s": time.perf_counter() - run_started, + "final_block_parameter_norm": parameter_norm(block_parameters), + } + if args.save_final_weights: + save_predictor_weights(model, run_dir / "predictor_final.safetensors") + atomic_json(metrics_path, result) + if latest_path.exists(): + latest_path.unlink() + del model, optimizer + torch.cuda.empty_cache() + return result + + +def write_summary(output_dir: Path, results: list[dict[str, Any]]) -> None: + rows = [ + { + "name": result["name"], + "source_layer_1": result["source_layers"][0], + "source_layer_2": result["source_layers"][1], + "source_layer_3": result["source_layers"][2], + "final_val_flow_mse": result["final_val_flow_mse"], + "final_val_hidden_mse": result["final_val_hidden_mse"], + "final_val_total_loss": result["final_val_total_loss"], + "val_flow_mse_auc": result["val_flow_mse_auc"], + "val_hidden_mse_auc": result["val_hidden_mse_auc"], + "val_total_loss_auc": result["val_total_loss_auc"], + "training_time_s": result["training_time_s"], + } + for result in results + ] + rows.sort(key=lambda row: float(row["final_val_flow_mse"])) + 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.dataset_root = resolve(args.dataset_root) + args.checkpoint_path = resolve(args.checkpoint_path) + args.config_path = resolve(args.config_path) + args.output_dir = resolve(args.output_dir) + args.output_dir.mkdir(parents=True, exist_ok=True) + + triples = ( + default_triples() + if args.triples is None + else list(dict.fromkeys(args.triples)) + ) + if args.max_triples is not None: + triples = triples[: args.max_triples] + train_prompt_ids = list(range(args.train_prompts)) + val_prompt_ids = list( + range(args.train_prompts, args.train_prompts + args.val_prompts) + ) + all_prompt_ids = train_prompt_ids + val_prompt_ids + device = torch.device("cuda") + set_seed(args.seed) + torch.set_grad_enabled(True) + + print("[setup] loading frozen generator_ema", flush=True) + teacher = load_teacher(args.checkpoint_path, args.config_path, device) + print("[setup] loading common offline trajectories into RAM", flush=True) + store = OfflinePredictorStore(args.dataset_root, all_prompt_ids) + schedule = BatchSchedule( + train_prompt_ids, args.batch_size, args.max_steps, args.seed + ) + shared_nonblock_state = build_shared_nonblock_state(args.seed) + manifest = { + "status": "running", + "architecture": "three_block_predictor", + "initialization_method": "teacher_full", + "triples": [list(triple) for triple in triples], + "num_triples": len(triples), + "default_consecutive_triples": 28, + "max_steps": args.max_steps, + "seed": args.seed, + "train_prompt_ids": train_prompt_ids, + "val_prompt_ids": val_prompt_ids, + "batch_schedule_sha256": schedule.fingerprint(), + } + atomic_json(args.output_dir / "sweep_manifest.json", manifest) + + results: list[dict[str, Any]] = [] + for index, triple in enumerate(triples, start=1): + name = experiment_name(triple) + metrics_path = args.output_dir / name / "metrics.json" + if metrics_path.exists(): + existing = json.loads(metrics_path.read_text(encoding="utf-8")) + if existing.get("status") == "complete": + print(f"[sweep] {index}/{len(triples)} skip {name}", flush=True) + results.append(existing) + continue + print( + f"[data] {index}/{len(triples)} loading layers {triple}", flush=True + ) + store.load_layer_caches(triple, teacher, device) + result = run_experiment( + triple=triple, + args=args, + store=store, + teacher=teacher, + train_prompt_ids=train_prompt_ids, + val_prompt_ids=val_prompt_ids, + schedule=schedule, + shared_nonblock_state=shared_nonblock_state, + device=device, + ) + results.append(result) + write_summary(args.output_dir, results) + + manifest["status"] = "complete" + atomic_json(args.output_dir / "sweep_manifest.json", manifest) + write_summary(args.output_dir, results) + print( + f"[complete] {len(results)} triples -> {args.output_dir / 'summary.csv'}", + flush=True, + ) + + +if __name__ == "__main__": + main() diff --git a/scripts/run_two_block_pair_sweep.py b/scripts/run_two_block_pair_sweep.py new file mode 100644 index 0000000000000000000000000000000000000000..daca1e3d6afc486e5541f5a51a2eb7cb1381f19a --- /dev/null +++ b/scripts/run_two_block_pair_sweep.py @@ -0,0 +1,632 @@ +#!/usr/bin/env python3 +"""Train two-block Predictors over mirror and consecutive Teacher pairs.""" + +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 torch +import torch.nn.functional as F +from safetensors.torch import save_file +from torch.optim import AdamW + +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 OfflinePredictorStore, TOKENS_PER_CHUNK +from predictor_training.two_block import TwoBlockPredictor +from scripts.run_single_block_init_sweep import ( + BatchSchedule, + append_jsonl, + atomic_json, + build_shared_nonblock_state, + frozen_inputs, + gradient_norm, + hidden_to_flow, + load_teacher, + lr_values, + move_batch, + normalized_auc, + parameter_norm, +) +from utils.misc import set_seed + + +def default_pairs() -> list[tuple[int, int]]: + mirror = [(index, 29 - index) for index in range(15)] + consecutive = [(index, index + 1) for index in range(29)] + return list(dict.fromkeys(mirror + consecutive)) + + +def pair_kind(pair: tuple[int, int]) -> str: + mirror = pair[0] + pair[1] == 29 + consecutive = pair[1] == pair[0] + 1 + if mirror and consecutive: + return "mirror+consecutive" + if mirror: + return "mirror" + if consecutive: + return "consecutive" + return "custom" + + +def parse_pair(value: str) -> tuple[int, int]: + try: + left_text, right_text = value.split(",", maxsplit=1) + pair = (int(left_text), int(right_text)) + except Exception as error: + raise argparse.ArgumentTypeError( + f"Pair must look like 0,29; got {value!r}" + ) from error + if not (0 <= pair[0] < pair[1] < 30): + raise argparse.ArgumentTypeError( + f"Pair must satisfy 0 <= first < second < 30; got {pair}" + ) + return pair + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--gpu", default=PHYSICAL_GPU) + parser.add_argument( + "--dataset_root", + type=Path, + default=Path("outputs/predictor_offline_100_all_blocks"), + ) + parser.add_argument( + "--checkpoint_path", + type=Path, + default=Path("checkpoints/self_forcing_dmd.pt"), + ) + parser.add_argument( + "--config_path", type=Path, default=Path("configs/self_forcing_sid.yaml") + ) + parser.add_argument( + "--output_dir", type=Path, default=Path("outputs/two_block_pair_sweep") + ) + parser.add_argument( + "--pairs", + type=parse_pair, + nargs="*", + default=None, + help="Pairs such as 0,29 1,28. Omit for all mirror+consecutive pairs.", + ) + parser.add_argument("--max_pairs", type=int, default=None) + parser.add_argument("--max_steps", type=int, default=1000) + parser.add_argument("--batch_size", type=int, default=32) + parser.add_argument("--eval_batch_size", type=int, default=10) + parser.add_argument("--eval_every", type=int, default=100) + parser.add_argument("--log_every", type=int, default=20) + parser.add_argument("--save_every", type=int, default=100) + parser.add_argument("--train_prompts", type=int, default=80) + parser.add_argument("--val_prompts", type=int, default=20) + parser.add_argument("--seed", type=int, default=0) + parser.add_argument("--fusion_lr", type=float, default=1e-4) + parser.add_argument("--block_lr", type=float, default=1e-5) + parser.add_argument("--weight_decay", type=float, default=0.01) + parser.add_argument("--hidden_weight", type=float, default=0.1) + parser.add_argument("--flow_weight", type=float, default=1.0) + parser.add_argument("--grad_clip", type=float, default=1.0) + parser.add_argument("--fusion_warmup_steps", type=int, default=100) + parser.add_argument("--block_freeze_steps", type=int, default=100) + parser.add_argument("--block_warmup_steps", type=int, default=100) + parser.add_argument( + "--gradient_checkpointing", + action=argparse.BooleanOptionalAction, + default=True, + ) + parser.add_argument( + "--save_final_weights", + action=argparse.BooleanOptionalAction, + default=True, + ) + args = parser.parse_args() + if args.max_steps < 1: + parser.error("--max_steps must be positive") + if args.train_prompts < 1 or args.val_prompts < 1: + parser.error("Prompt counts must be positive") + if args.train_prompts + args.val_prompts > 100: + parser.error("The offline dataset contains 100 prompts") + if args.batch_size > args.train_prompts: + parser.error("--batch_size cannot exceed --train_prompts") + if args.eval_batch_size > args.val_prompts: + args.eval_batch_size = args.val_prompts + return args + + +def resolve(path: Path) -> Path: + path = path.expanduser() + return path.resolve() if path.is_absolute() else (REPO_ROOT / path).resolve() + + +def make_model( + teacher: torch.nn.Module, + pair: tuple[int, int], + shared_nonblock_state: dict[str, dict[str, torch.Tensor]], + seed: int, + gradient_checkpointing: bool, + device: torch.device, +) -> TwoBlockPredictor: + set_seed(seed) + model = TwoBlockPredictor( + [teacher.blocks[pair[0]], teacher.blocks[pair[1]]], + dim=teacher.dim, + gradient_checkpointing=gradient_checkpointing, + ) + model.fusion.load_state_dict(shared_nonblock_state["fusion"], strict=True) + model.residual_out.load_state_dict( + shared_nonblock_state["residual_out"], strict=True + ) + return model.to(device=device) + + +def forward_predictor( + model: TwoBlockPredictor, + batch: dict[str, Any], + teacher: torch.nn.Module, + device: torch.device, +) -> tuple[torch.Tensor, torch.Tensor]: + frozen = frozen_inputs(batch, teacher, device) + with torch.autocast(device_type="cuda", dtype=torch.bfloat16): + pred_hidden = model( + current_tokens=frozen["current_tokens"], + anchor_hidden=batch["anchor_hidden"], + previous_hidden=batch["previous_hidden"], + timestep_modulation=frozen["timestep_modulation"], + grid_sizes=frozen["grid_sizes"], + freqs=frozen["freqs"], + history_ks=[batch["history_k_0"], batch["history_k_1"]], + history_vs=[batch["history_v_0"], batch["history_v_1"]], + cross_ks=[batch["cross_k_0"], batch["cross_k_1"]], + cross_vs=[batch["cross_v_0"], batch["cross_v_1"]], + current_start=batch["chunk"] * TOKENS_PER_CHUNK, + ) + pred_flow = hidden_to_flow( + pred_hidden, + frozen["head_embedding"], + frozen["grid_sizes"], + teacher, + ) + return pred_hidden, pred_flow + + +@torch.inference_mode() +def evaluate( + model: TwoBlockPredictor, + store: OfflinePredictorStore, + pair: tuple[int, int], + val_prompt_ids: list[int], + batch_size: int, + teacher: torch.nn.Module, + device: torch.device, + hidden_weight: float, + flow_weight: float, +) -> dict[str, float]: + model.eval() + hidden_squared = 0.0 + hidden_elements = 0 + flow_squared = 0.0 + flow_elements = 0 + started = time.perf_counter() + for chunk in range(1, 7): + for target_step in range(1, 4): + for start in range(0, len(val_prompt_ids), batch_size): + prompt_ids = val_prompt_ids[start : start + batch_size] + batch = move_batch( + store.batch_layers(prompt_ids, chunk, target_step, pair), device + ) + pred_hidden, pred_flow = forward_predictor( + model, batch, teacher, device + ) + hidden_error = pred_hidden.float() - batch["target_hidden"].float() + flow_error = pred_flow.float() - batch["target_flow"].float() + hidden_squared += float(hidden_error.square().sum()) + hidden_elements += hidden_error.numel() + flow_squared += float(flow_error.square().sum()) + flow_elements += flow_error.numel() + del batch, pred_hidden, pred_flow, hidden_error, flow_error + hidden_mse = hidden_squared / hidden_elements + flow_mse = flow_squared / flow_elements + model.train() + return { + "hidden_mse": hidden_mse, + "flow_mse": flow_mse, + "total_loss": hidden_weight * hidden_mse + flow_weight * flow_mse, + "eval_time_s": time.perf_counter() - started, + } + + +def save_predictor_weights(model: TwoBlockPredictor, path: Path) -> None: + tensors = { + key: value.detach().to(device="cpu").contiguous() + for key, value in model.state_dict().items() + } + temporary = path.with_suffix(path.suffix + ".tmp") + save_file(tensors, temporary) + os.replace(temporary, path) + + +def run_experiment( + *, + pair: tuple[int, int], + args: argparse.Namespace, + store: OfflinePredictorStore, + teacher: torch.nn.Module, + train_prompt_ids: list[int], + val_prompt_ids: list[int], + schedule: BatchSchedule, + shared_nonblock_state: dict[str, dict[str, torch.Tensor]], + device: torch.device, +) -> dict[str, Any]: + name = f"pair_{pair[0]:02d}_{pair[1]:02d}" + run_dir = args.output_dir / name + 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": + print(f"[run] skip completed {name}", flush=True) + return existing + run_dir.mkdir(parents=True, exist_ok=True) + log_path = run_dir / "train_log.jsonl" + + model = make_model( + teacher, + pair, + shared_nonblock_state, + args.seed, + args.gradient_checkpointing, + device, + ) + fusion_parameters = model.fusion_parameters() + block_parameters = model.block_parameters() + optimizer = AdamW( + [ + {"params": fusion_parameters, "lr": args.fusion_lr}, + {"params": block_parameters, "lr": args.block_lr}, + ], + betas=(0.9, 0.95), + weight_decay=args.weight_decay, + ) + initial_block_norm = parameter_norm(block_parameters) + evaluations: list[dict[str, Any]] = [] + start_step = 0 + latest_path = run_dir / "training_latest.pt" + if latest_path.exists(): + state = torch.load(latest_path, map_location="cpu", weights_only=False) + model.load_state_dict(state["model"], strict=True) + optimizer.load_state_dict(state["optimizer"]) + evaluations = state["evaluations"] + start_step = int(state["step"]) + print(f"[run] resume {name} at step {start_step}", flush=True) + + model.set_blocks_trainable(start_step >= args.block_freeze_steps) + config = { + "name": name, + "architecture": "two_block_predictor", + "initialization_method": "teacher_full", + "source_layers": list(pair), + "pair_kind": pair_kind(pair), + "source_definition": ( + "Predictor block i, its clean-history K/V, and its text K/V use " + "generator_ema source_layers[i]" + ), + "seed": args.seed, + "train_prompt_ids": train_prompt_ids, + "val_prompt_ids": val_prompt_ids, + "max_steps": args.max_steps, + "batch_size": args.batch_size, + "eval_batch_size": args.eval_batch_size, + "fusion_lr": args.fusion_lr, + "block_lr": args.block_lr, + "fusion_warmup_steps": args.fusion_warmup_steps, + "block_freeze_steps": args.block_freeze_steps, + "block_warmup_steps": args.block_warmup_steps, + "weight_decay": args.weight_decay, + "hidden_weight": args.hidden_weight, + "flow_weight": args.flow_weight, + "gradient_checkpointing": args.gradient_checkpointing, + "batch_schedule_sha256": schedule.fingerprint(), + "initial_block_parameter_norm": initial_block_norm, + "trainable_parameters": sum( + parameter.numel() for parameter in model.parameters() + ), + } + atomic_json(run_dir / "config.json", config) + print( + f"[run] {name}: kind={config['pair_kind']} start={start_step}", flush=True + ) + + if not evaluations: + initial_eval = evaluate( + model, + store, + pair, + val_prompt_ids, + args.eval_batch_size, + teacher, + device, + args.hidden_weight, + args.flow_weight, + ) + evaluations.append({"step": 0, **initial_eval}) + print( + f"[eval] {name} step=0 flow={initial_eval['flow_mse']:.8f} " + f"hidden={initial_eval['hidden_mse']:.8f}", + flush=True, + ) + + optimizer.zero_grad(set_to_none=True) + model.train() + run_started = time.perf_counter() + for step in range(start_step, args.max_steps): + block_enabled = step >= args.block_freeze_steps + if any( + parameter.requires_grad != block_enabled + for parameter in block_parameters + ): + model.set_blocks_trainable(block_enabled) + + fusion_lr, block_lr = lr_values( + step, + args.max_steps, + args.fusion_lr, + args.block_lr, + args.fusion_warmup_steps, + args.block_freeze_steps, + args.block_warmup_steps, + ) + optimizer.param_groups[0]["lr"] = fusion_lr + optimizer.param_groups[1]["lr"] = block_lr + + chunk, target_step, prompt_ids = schedule.entries[step] + batch = move_batch( + store.batch_layers(prompt_ids, chunk, target_step, pair), device + ) + step_started = time.perf_counter() + pred_hidden, pred_flow = forward_predictor(model, batch, teacher, device) + hidden_loss = F.mse_loss( + pred_hidden.float(), batch["target_hidden"].float() + ) + flow_loss = F.mse_loss(pred_flow.float(), batch["target_flow"].float()) + loss = args.hidden_weight * hidden_loss + args.flow_weight * flow_loss + loss.backward() + fusion_grad_norm = gradient_norm(fusion_parameters) + block_grad_norm = gradient_norm(block_parameters) + block_0_grad_norm = gradient_norm(list(model.blocks[0].parameters())) + block_1_grad_norm = gradient_norm(list(model.blocks[1].parameters())) + total_grad_norm = torch.nn.utils.clip_grad_norm_( + model.parameters(), args.grad_clip + ) + optimizer.step() + optimizer.zero_grad(set_to_none=True) + completed_step = step + 1 + + if completed_step == 1 or completed_step % args.log_every == 0: + torch.cuda.synchronize() + record = { + "step": completed_step, + "train_total_loss": float(loss.detach()), + "train_hidden_mse": float(hidden_loss.detach()), + "train_flow_mse": float(flow_loss.detach()), + "fusion_grad_norm": fusion_grad_norm, + "block_grad_norm": block_grad_norm, + "block_0_grad_norm": block_0_grad_norm, + "block_1_grad_norm": block_1_grad_norm, + "total_grad_norm_before_clip": float(total_grad_norm), + "fusion_lr": fusion_lr, + "block_lr": block_lr, + "chunk": chunk, + "target_step": target_step, + "step_time_s": time.perf_counter() - step_started, + "peak_gpu_gib": torch.cuda.max_memory_allocated() / (1024**3), + } + append_jsonl(log_path, record) + print( + f"[train] {name} {completed_step}/{args.max_steps} " + f"flow={record['train_flow_mse']:.8f} " + f"time={record['step_time_s']:.2f}s " + f"mem={record['peak_gpu_gib']:.1f}G", + flush=True, + ) + + if completed_step % args.eval_every == 0 or completed_step == args.max_steps: + validation = evaluate( + model, + store, + pair, + val_prompt_ids, + args.eval_batch_size, + teacher, + device, + args.hidden_weight, + args.flow_weight, + ) + evaluations.append({"step": completed_step, **validation}) + print( + f"[eval] {name} step={completed_step} " + f"flow={validation['flow_mse']:.8f} " + f"hidden={validation['hidden_mse']:.8f}", + flush=True, + ) + + if completed_step % args.save_every == 0 or completed_step == args.max_steps: + temporary = latest_path.with_suffix(".pt.tmp") + torch.save( + { + "model": { + key: value.detach().cpu() + for key, value in model.state_dict().items() + }, + "optimizer": optimizer.state_dict(), + "evaluations": evaluations, + "step": completed_step, + }, + temporary, + ) + os.replace(temporary, latest_path) + + del batch, pred_hidden, pred_flow, hidden_loss, flow_loss, loss + del total_grad_norm + + final = evaluations[-1] + result = { + "status": "complete", + **config, + "final_val_hidden_mse": final["hidden_mse"], + "final_val_flow_mse": final["flow_mse"], + "final_val_total_loss": final["total_loss"], + "val_hidden_mse_auc": normalized_auc( + evaluations, "hidden_mse", args.max_steps + ), + "val_flow_mse_auc": normalized_auc( + evaluations, "flow_mse", args.max_steps + ), + "val_total_loss_auc": normalized_auc( + evaluations, "total_loss", args.max_steps + ), + "evaluations": evaluations, + "training_time_s": time.perf_counter() - run_started, + "final_block_parameter_norm": parameter_norm(block_parameters), + } + if args.save_final_weights: + save_predictor_weights(model, run_dir / "predictor_final.safetensors") + atomic_json(metrics_path, result) + if latest_path.exists(): + latest_path.unlink() + del model, optimizer + torch.cuda.empty_cache() + return result + + +def write_summary(output_dir: Path, results: list[dict[str, Any]]) -> None: + rows = [ + { + "name": result["name"], + "source_layer_1": result["source_layers"][0], + "source_layer_2": result["source_layers"][1], + "pair_kind": result["pair_kind"], + "final_val_flow_mse": result["final_val_flow_mse"], + "final_val_hidden_mse": result["final_val_hidden_mse"], + "final_val_total_loss": result["final_val_total_loss"], + "val_flow_mse_auc": result["val_flow_mse_auc"], + "val_hidden_mse_auc": result["val_hidden_mse_auc"], + "val_total_loss_auc": result["val_total_loss_auc"], + "training_time_s": result["training_time_s"], + } + for result in results + ] + rows.sort(key=lambda row: float(row["final_val_flow_mse"])) + 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.dataset_root = resolve(args.dataset_root) + args.checkpoint_path = resolve(args.checkpoint_path) + args.config_path = resolve(args.config_path) + args.output_dir = resolve(args.output_dir) + args.output_dir.mkdir(parents=True, exist_ok=True) + + pairs = default_pairs() if args.pairs is None else list(dict.fromkeys(args.pairs)) + if args.max_pairs is not None: + pairs = pairs[: args.max_pairs] + train_prompt_ids = list(range(args.train_prompts)) + val_prompt_ids = list( + range(args.train_prompts, args.train_prompts + args.val_prompts) + ) + all_prompt_ids = train_prompt_ids + val_prompt_ids + device = torch.device("cuda") + set_seed(args.seed) + torch.set_grad_enabled(True) + + print("[setup] loading frozen generator_ema", flush=True) + teacher = load_teacher(args.checkpoint_path, args.config_path, device) + print("[setup] loading common offline trajectories into RAM", flush=True) + store = OfflinePredictorStore(args.dataset_root, all_prompt_ids) + schedule = BatchSchedule( + train_prompt_ids, args.batch_size, args.max_steps, args.seed + ) + shared_nonblock_state = build_shared_nonblock_state(args.seed) + manifest = { + "status": "running", + "pairs": [list(pair) for pair in pairs], + "num_pairs": len(pairs), + "mirror_pairs": 15, + "consecutive_pairs": 29, + "deduplicated_overlap": [14, 15], + "max_steps": args.max_steps, + "seed": args.seed, + "train_prompt_ids": train_prompt_ids, + "val_prompt_ids": val_prompt_ids, + "batch_schedule_sha256": schedule.fingerprint(), + } + atomic_json(args.output_dir / "sweep_manifest.json", manifest) + + results: list[dict[str, Any]] = [] + for index, pair in enumerate(pairs, start=1): + name = f"pair_{pair[0]:02d}_{pair[1]:02d}" + metrics_path = args.output_dir / name / "metrics.json" + if metrics_path.exists(): + existing = json.loads(metrics_path.read_text(encoding="utf-8")) + if existing.get("status") == "complete": + print(f"[sweep] {index}/{len(pairs)} skip {name}", flush=True) + results.append(existing) + continue + print( + f"[data] {index}/{len(pairs)} loading pair {pair[0]},{pair[1]}", + flush=True, + ) + store.load_layer_caches(pair, teacher, device) + result = run_experiment( + pair=pair, + args=args, + store=store, + teacher=teacher, + train_prompt_ids=train_prompt_ids, + val_prompt_ids=val_prompt_ids, + schedule=schedule, + shared_nonblock_state=shared_nonblock_state, + device=device, + ) + results.append(result) + write_summary(args.output_dir, results) + + manifest["status"] = "complete" + atomic_json(args.output_dir / "sweep_manifest.json", manifest) + write_summary(args.output_dir, results) + print( + f"[complete] {len(results)} pairs -> {args.output_dir / 'summary.csv'}", + flush=True, + ) + + +if __name__ == "__main__": + main() diff --git a/scripts/run_wan14b_timestep_chunk_cosine.py b/scripts/run_wan14b_timestep_chunk_cosine.py new file mode 100644 index 0000000000000000000000000000000000000000..8f8abed324fa318b6b8705a52642fd96a6b3486a --- /dev/null +++ b/scripts/run_wan14b_timestep_chunk_cosine.py @@ -0,0 +1,422 @@ +#!/usr/bin/env python3 +"""Run native Wan2.1-T2V-14B/50-step feature cosine analysis. + +Wan14B evaluates the whole 21-latent video in one forward. The token sequence +is split into seven three-latent temporal slices. Consequently ``cross_same`` +means adjacent full-video temporal slices, not a previously computed AR chunk. +Only the conditional CFG stream is recorded. +""" + +from __future__ import annotations + +import argparse +import csv +import gc +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 torch +import torch.nn.functional as F + +ROOT = Path(__file__).resolve().parents[1] +if str(ROOT) not in sys.path: + sys.path.insert(0, str(ROOT)) + +from utils.misc import set_seed +from utils.wan_wrapper import WanTextEncoder +from wan.modules.model import WanModel +from wan.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser() + parser.add_argument("--gpu", default=PHYSICAL_GPU) + parser.add_argument( + "--checkpoint_dir", type=Path, default=Path("wan_models/Wan2.1-T2V-14B") + ) + parser.add_argument( + "--prompt_path", type=Path, default=Path("prompts/MovieGenVideoBench_extended.txt") + ) + parser.add_argument("--output_dir", type=Path, required=True) + parser.add_argument("--num_prompts", type=int, default=10) + parser.add_argument( + "--prompt_ids", + default="", + help="Optional comma-separated global prompt ids for multi-GPU sharding.", + ) + parser.add_argument("--sampling_steps", type=int, default=50) + parser.add_argument("--layers", type=int, nargs="+", default=[9, 19, 29, 39]) + parser.add_argument("--max_tokens", type=int, default=240) + parser.add_argument( + "--chunk_pairing", + choices=["matched_slot", "boundary_to_all"], + default="matched_slot", + help=( + "Cross-chunk token pairing. boundary_to_all compares all current " + "temporal slots with the previous slice's last slot." + ), + ) + parser.add_argument("--seed", type=int, default=0) + parser.add_argument("--guide_scale", type=float, default=5.0) + parser.add_argument("--shift", type=float, default=5.0) + parser.add_argument("--overwrite", action="store_true") + return parser.parse_args() + + +def resolve(path: Path) -> Path: + return path if path.is_absolute() else (ROOT / path).resolve() + + +def write_csv(path: Path, rows: list[dict[str, Any]]) -> None: + if not rows: + return + path.parent.mkdir(parents=True, exist_ok=True) + 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, extrasaction="ignore") + writer.writeheader() + writer.writerows(rows) + + +def cosine_metrics(reference: torch.Tensor, target: torch.Tensor) -> dict[str, float]: + x = reference.float() + y = target.float() + values = F.cosine_similarity(x, y, dim=-1, eps=1e-8) + quantiles = torch.quantile(values, torch.tensor([0.1, 0.5, 0.9])) + return { + "token_cosine_mean": float(values.mean()), + "token_cosine_p10": float(quantiles[0]), + "token_cosine_p50": float(quantiles[1]), + "token_cosine_p90": float(quantiles[2]), + "global_cosine": float(F.cosine_similarity(x.flatten()[None], y.flatten()[None])), + } + + +class NativeWanRecorder: + def __init__( + self, + model: WanModel, + layers: list[int], + max_tokens: int, + chunk_pairing: str = "matched_slot", + ): + self.model = model + self.layers = sorted(set(layers)) + self.max_tokens = int(max_tokens) + self.chunk_pairing = str(chunk_pairing) + self.handles = [] + self.active = False + self.prompt_id = -1 + self.step = -1 + self.timestep = 0.0 + self.rows: list[dict[str, Any]] = [] + self.previous: dict[tuple[int, int], torch.Tensor] = {} + self.previous_timestep: float | None = None + self.matched_pairs = {12: 0, 25: 12, 37: 25} + self.matched_history: dict[ + tuple[int, int, int], tuple[torch.Tensor, float] + ] = {} + for layer in self.layers: + if not 0 <= layer < len(model.blocks): + raise ValueError(f"Layer {layer} outside 0..{len(model.blocks)-1}") + self.handles.append(model.blocks[layer].register_forward_hook(self._hook(layer))) + + def reset(self, prompt_id: int) -> None: + self.active = False + self.prompt_id = int(prompt_id) + self.step = -1 + self.timestep = 0.0 + self.rows = [] + self.previous = {} + self.previous_timestep = None + self.matched_history = {} + + def begin(self, step: int, timestep: float) -> None: + self.step = int(step) + self.timestep = float(timestep) + self.active = True + + def end(self) -> None: + self.active = False + self.previous_timestep = self.timestep + + def close(self) -> None: + for handle in self.handles: + handle.remove() + self.handles.clear() + + def _sample(self, tokens: torch.Tensor) -> torch.Tensor: + if self.chunk_pairing == "boundary_to_all": + if tokens.shape[0] != 3 * 30 * 52: + raise ValueError(f"Expected 3x30x52 tokens, got {tokens.shape[0]}") + spatial_count = max(1, self.max_tokens // 3) + if 30 * 52 <= spatial_count: + indices = torch.arange(30 * 52, device=tokens.device) + else: + indices = torch.linspace( + 0, 30 * 52 - 1, spatial_count, device=tokens.device + ).round().long().unique() + return ( + tokens.reshape(3, 30 * 52, -1) + .index_select(1, indices) + .detach() + .to("cpu", torch.float16) + ) + if tokens.shape[0] <= self.max_tokens: + indices = torch.arange(tokens.shape[0], device=tokens.device) + else: + indices = torch.linspace( + 0, tokens.shape[0] - 1, self.max_tokens, device=tokens.device + ).round().long().unique() + return tokens.index_select(0, indices).detach().to("cpu", torch.float16) + + def _hook(self, layer: int): + def hook(_module, _inputs, output) -> None: + if not self.active: + return + feature = output[0] if isinstance(output, (tuple, list)) else output + if not isinstance(feature, torch.Tensor) or feature.ndim != 3: + raise ValueError(f"Unexpected layer {layer} output: {type(feature)}") + # 21 latent frames x (30*52) tokens, grouped into 7 slices of 3 frames. + per_frame = 30 * 52 + expected = 21 * per_frame + if feature.shape[1] < expected: + raise ValueError(f"Expected at least {expected} tokens, got {feature.shape[1]}") + current: list[torch.Tensor] = [] + for chunk in range(7): + start = chunk * 3 * per_frame + stop = start + 3 * per_frame + current.append(self._sample(feature[0, start:stop])) + for chunk, sampled in enumerate(current): + base = { + "model_family": "self_forcing", + "model_variant": "wan14b_50", + "chunk_semantics": "full_video_temporal_slice", + "prompt_id": self.prompt_id, + "layer_index": layer, + "chunk": chunk, + "target_step": self.step, + "target_timestep": self.timestep, + "sample_tokens": int(sampled.numel() // sampled.shape[-1]), + "feature_dim": int(sampled.shape[1]), + } + within = self.previous.get((layer, chunk)) + if within is not None: + self.rows.append({ + **base, + "comparison": "within_adjacent", + "reference_chunk": chunk, + "reference_step": self.step - 1, + "reference_timestep": self.previous_timestep, + **cosine_metrics(within, sampled), + }) + matched_reference_step = self.matched_pairs.get(self.step) + matched = ( + None + if matched_reference_step is None + else self.matched_history.get((layer, chunk, matched_reference_step)) + ) + if matched is not None: + matched_feature, matched_timestep = matched + self.rows.append({ + **base, + "comparison": "within_matched_gap", + "reference_chunk": chunk, + "reference_step": matched_reference_step, + "reference_timestep": matched_timestep, + **cosine_metrics(matched_feature, sampled), + }) + if chunk > 0: + cross_reference = ( + current[chunk - 1][-1:].expand_as(sampled) + if self.chunk_pairing == "boundary_to_all" + else current[chunk - 1] + ) + self.rows.append({ + **base, + "comparison": ( + "cross_boundary_to_all" + if self.chunk_pairing == "boundary_to_all" + else "cross_same" + ), + "reference_chunk": chunk - 1, + "reference_step": self.step, + "reference_timestep": self.timestep, + **cosine_metrics(cross_reference, sampled), + }) + self.previous[(layer, chunk)] = sampled + if self.step in {0, 12, 25, 37}: + self.matched_history[(layer, chunk, self.step)] = ( + sampled, + self.timestep, + ) + + return hook + + +@torch.inference_mode() +def encode_prompts(prompts: list[str]) -> tuple[torch.Tensor, torch.Tensor]: + print("[load] text encoder", flush=True) + encoder = WanTextEncoder().to(device="cuda", dtype=torch.bfloat16).eval() + conditional = encoder(text_prompts=prompts)["prompt_embeds"].detach().to("cpu", torch.bfloat16) + negative = encoder(text_prompts=[ + "色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止," + "整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指," + "画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合" + ])["prompt_embeds"].detach().to("cpu", torch.bfloat16) + del encoder + gc.collect() + torch.cuda.empty_cache() + return conditional, negative + + +@torch.inference_mode() +def main() -> None: + args = parse_args() + checkpoint_dir = resolve(args.checkpoint_dir) + prompt_path = resolve(args.prompt_path) + args.output_dir = resolve(args.output_dir) + args.output_dir.mkdir(parents=True, exist_ok=True) + all_prompts = [line.strip() for line in prompt_path.read_text(encoding="utf-8").splitlines() if line.strip()] + prompt_ids = ( + [int(value) for value in args.prompt_ids.split(",") if value.strip()] + if args.prompt_ids + else list(range(args.num_prompts)) + ) + if not prompt_ids or min(prompt_ids) < 0 or max(prompt_ids) >= len(all_prompts): + raise ValueError(f"Invalid prompt ids {prompt_ids} for {len(all_prompts)} prompts") + prompts = [all_prompts[index] for index in prompt_ids] + conditional, negative = encode_prompts(prompts) + print("[load] Wan14B in bfloat16", flush=True) + model = WanModel.from_pretrained( + checkpoint_dir, + torch_dtype=torch.bfloat16, + low_cpu_mem_usage=True, + ).eval().requires_grad_(False).to("cuda") + recorder = NativeWanRecorder( + model, + args.layers, + args.max_tokens, + chunk_pairing=args.chunk_pairing, + ) + scheduler = FlowUniPCMultistepScheduler( + num_train_timesteps=1000, shift=1, use_dynamic_shifting=False + ) + scheduler.set_timesteps(args.sampling_steps, device="cuda", shift=args.shift) + timesteps = scheduler.timesteps + manifest = { + "model_family": "self_forcing", + "model_variant": "wan14b_50", + "chunk_semantics": "full_video_temporal_slice", + "checkpoint_dir": str(checkpoint_dir), + "prompt_path": str(prompt_path), + "num_prompts": len(prompts), + "prompt_ids": prompt_ids, + "num_frames": 21, + "num_chunks": 7, + "sampling_steps": args.sampling_steps, + "timesteps": [float(value) for value in timesteps.detach().cpu()], + "layers": args.layers, + "layer_roles": ["early", "middle", "late", "final"], + "max_tokens": args.max_tokens, + "chunk_pairing": args.chunk_pairing, + "seed": args.seed, + "guide_scale": args.guide_scale, + "shift": args.shift, + "physical_gpu": args.gpu, + "conditional_stream_only": True, + } + (args.output_dir / "run_manifest.json").write_text( + json.dumps(manifest, indent=2, ensure_ascii=False) + "\n", encoding="utf-8" + ) + all_rows: list[dict[str, Any]] = [] + seq_len = 21 * 30 * 52 + try: + for local_index, (prompt_id, prompt) in enumerate(zip(prompt_ids, prompts)): + run_dir = args.output_dir / "runs" / f"prompt_{prompt_id:04d}" + pair_path = run_dir / "feature_pair_metrics.csv" + if pair_path.exists() and not args.overwrite: + with pair_path.open(newline="", encoding="utf-8") as handle: + all_rows.extend(csv.DictReader(handle)) + print(f"[skip] prompt {prompt_id}", flush=True) + continue + run_dir.mkdir(parents=True, exist_ok=True) + set_seed(args.seed) + generator = torch.Generator(device="cuda").manual_seed(args.seed) + latents = [torch.randn( + 16, 21, 60, 104, + dtype=torch.float32, device="cuda", generator=generator, + )] + cond = conditional[local_index].to("cuda") + uncond = negative[0].to("cuda") + recorder.reset(prompt_id) + scheduler.set_timesteps(args.sampling_steps, device="cuda", shift=args.shift) + torch.cuda.reset_peak_memory_stats() + torch.cuda.synchronize() + started = time.perf_counter() + with torch.amp.autocast("cuda", dtype=torch.bfloat16): + for step, timestep in enumerate(scheduler.timesteps): + t = torch.stack([timestep]) + recorder.begin(step, float(timestep)) + pred_cond = model(latents, t=t, context=[cond], seq_len=seq_len)[0] + recorder.end() + pred_uncond = model(latents, t=t, context=[uncond], seq_len=seq_len)[0] + prediction = pred_uncond + args.guide_scale * (pred_cond - pred_uncond) + next_latent = scheduler.step( + prediction.unsqueeze(0), + timestep, + latents[0].unsqueeze(0), + return_dict=False, + generator=generator, + )[0] + latents = [next_latent.squeeze(0)] + torch.cuda.synchronize() + elapsed = time.perf_counter() - started + rows = list(recorder.rows) + write_csv(pair_path, rows) + metadata = { + "prompt_id": prompt_id, + "prompt": prompt, + "elapsed_s": elapsed, + "peak_gpu_gib": torch.cuda.max_memory_allocated() / 1024**3, + "pair_rows": len(rows), + } + (run_dir / "run_metadata.json").write_text( + json.dumps(metadata, indent=2, ensure_ascii=False) + "\n", encoding="utf-8" + ) + all_rows.extend(rows) + print( + f"[done] prompt={prompt_id} rows={len(rows)} elapsed={elapsed:.1f}s " + f"peak={metadata['peak_gpu_gib']:.1f}GiB", + flush=True, + ) + del latents, cond, uncond + torch.cuda.empty_cache() + finally: + recorder.close() + write_csv(args.output_dir / "feature_pair_metrics.csv", all_rows) + print(f"[complete] {args.output_dir}", flush=True) + + +if __name__ == "__main__": + main() diff --git a/scripts/summarize_boundary_to_all_cosine.py b/scripts/summarize_boundary_to_all_cosine.py new file mode 100644 index 0000000000000000000000000000000000000000..4cf8417bac8060fca85ad9f733e3bef3c7222c1a --- /dev/null +++ b/scripts/summarize_boundary_to_all_cosine.py @@ -0,0 +1,209 @@ +#!/usr/bin/env python3 +"""Summarize matched-step and previous-boundary-to-all cosine consistently.""" + +from __future__ import annotations + +import argparse +import csv +import json +from collections import defaultdict +from pathlib import Path +from typing import Any + +import numpy as np +import torch +import torch.nn.functional as F + + +ROLE_MAP = { + ("self_forcing", "wan14b50"): {9: "early", 19: "middle", 29: "late", 39: "final"}, + ("self_forcing", "dmd4"): {7: "early", 14: "middle", 22: "late", 29: "final"}, + ("causal_forcing", "ar50"): {7: "early", 14: "middle", 22: "late", 29: "final"}, + ("causal_forcing", "dmd4"): {7: "early", 14: "middle", 22: "late", 29: "final"}, +} +ORDER = [ + ("self_forcing", "wan14b50"), + ("self_forcing", "dmd4"), + ("causal_forcing", "ar50"), + ("causal_forcing", "dmd4"), +] + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser() + parser.add_argument("--self_dmd_root", type=Path, required=True) + parser.add_argument("--self_wan_root", type=Path, required=True) + parser.add_argument("--causal_ar_root", type=Path, required=True) + parser.add_argument("--causal_dmd_root", type=Path, required=True) + parser.add_argument("--output_root", type=Path, required=True) + return parser.parse_args() + + +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) + path.parent.mkdir(parents=True, exist_ok=True) + with path.open("w", newline="", encoding="utf-8") as handle: + writer = csv.DictWriter(handle, fieldnames=fields) + writer.writeheader() + writer.writerows(rows) + + +def self_dmd_rows(root: Path) -> list[dict[str, Any]]: + rows = [] + role_map = ROLE_MAP[("self_forcing", "dmd4")] + for path in sorted((root / "runs").glob("prompt_*.pt")): + state = torch.load(path, map_location="cpu", weights_only=False) + prompt_id = int(state["run_index"]) + coords = state["sample_coords"]["hidden"] + spatial_counts = [int((coords[:, 0] == slot).sum()) for slot in range(3)] + if len(set(spatial_counts)) != 1: + raise ValueError(f"Unbalanced temporal sampling in {path}: {spatial_counts}") + spatial_count = spatial_counts[0] + for layer, role in role_map.items(): + records = state["records"][f"block_{layer}_hidden"] + for chunk in range(1, 7): + for step in range(1, 4): + target = records[f"{chunk}:{step}"].float().reshape(3, spatial_count, -1) + previous_step = records[f"{chunk}:{step - 1}"].float().reshape( + 3, spatial_count, -1 + ) + previous_chunk = records[f"{chunk - 1}:{step}"].float().reshape( + 3, spatial_count, -1 + ) + for comparison, reference in ( + ("within_adjacent", previous_step), + ("cross_boundary_to_all", previous_chunk[-1:].expand_as(target)), + ): + rows.append({ + "model_family": "self_forcing", + "model_variant": "dmd4", + "layer_role": role, + "layer_index": layer, + "prompt_id": prompt_id, + "comparison": comparison, + "target_chunk": chunk, + "target_step": step, + "token_cosine_mean": float( + F.cosine_similarity(target, reference, dim=-1).mean() + ), + }) + return rows + + +def direct_rows(root: Path, family: str, variant: str) -> list[dict[str, Any]]: + rows = [] + role_map = ROLE_MAP[(family, variant)] + paths = sorted((root / "runs").glob("prompt_*/feature_pair_metrics.csv")) + if not paths: + paths = sorted(root.glob("shard_*/runs/prompt_*/feature_pair_metrics.csv")) + for path in paths: + with path.open(newline="", encoding="utf-8") as handle: + for raw in csv.DictReader(handle): + comparison = raw["comparison"] + if comparison not in {"within_adjacent", "cross_boundary_to_all"}: + continue + step = int(raw["target_step"]) + if step < 1: + continue + layer = int(raw["layer_index"]) + rows.append({ + "model_family": family, + "model_variant": variant, + "layer_role": role_map[layer], + "layer_index": layer, + "prompt_id": int(raw["prompt_id"]), + "comparison": comparison, + "target_chunk": int(raw["chunk"]), + "target_step": step, + "token_cosine_mean": float(raw["token_cosine_mean"]), + }) + return rows + + +def bootstrap(values: list[float], seed: int, rounds: int = 10000) -> tuple[float, float, float]: + array = np.asarray(values, dtype=np.float64) + generator = np.random.default_rng(seed) + indices = generator.integers(0, len(array), size=(rounds, len(array))) + means = array[indices].mean(axis=1) + return float(array.mean()), float(np.quantile(means, 0.025)), float(np.quantile(means, 0.975)) + + +def main() -> None: + args = parse_args() + output = args.output_root.resolve() + output.mkdir(parents=True, exist_ok=True) + rows = self_dmd_rows(args.self_dmd_root.resolve()) + rows += direct_rows(args.self_wan_root.resolve(), "self_forcing", "wan14b50") + rows += direct_rows(args.causal_ar_root.resolve(), "causal_forcing", "ar50") + rows += direct_rows(args.causal_dmd_root.resolve(), "causal_forcing", "dmd4") + + prompt_groups: dict[tuple[Any, ...], list[float]] = defaultdict(list) + for row in rows: + key = ( + row["model_family"], row["model_variant"], row["layer_role"], + row["layer_index"], row["prompt_id"], row["comparison"], + ) + prompt_groups[key].append(float(row["token_cosine_mean"])) + prompt_rows = [] + for key, values in prompt_groups.items(): + prompt_rows.append({ + **dict(zip( + ["model_family", "model_variant", "layer_role", "layer_index", "prompt_id", "comparison"], + key, + )), + "pair_count": len(values), + "token_cosine_mean": float(np.mean(values)), + }) + + groups: dict[tuple[Any, ...], list[float]] = defaultdict(list) + for row in prompt_rows: + key = (row["model_family"], row["model_variant"], row["layer_role"], row["comparison"]) + groups[key].append(float(row["token_cosine_mean"])) + summary = [] + for key, values in groups.items(): + avg, low, high = bootstrap(values, 20260828 + sum(map(ord, "".join(map(str, key))))) + summary.append({ + **dict(zip(["model_family", "model_variant", "layer_role", "comparison"], key)), + "prompt_count": len(values), + "token_cosine_mean": avg, + "ci95_low": low, + "ci95_high": high, + }) + + lookup = { + (row["model_family"], row["model_variant"], row["layer_role"], row["comparison"]): row + for row in summary + } + table_rows = [] + for family, variant in ORDER: + item: dict[str, Any] = {"model_family": family, "model_variant": variant} + for role in ("early", "middle", "late", "final"): + step = lookup[(family, variant, role, "within_adjacent")] + chunk = lookup[(family, variant, role, "cross_boundary_to_all")] + item[f"{role}_step"] = step["token_cosine_mean"] + item[f"{role}_chunk"] = chunk["token_cosine_mean"] + table_rows.append(item) + + write_csv(output / "boundary_pair_metrics.csv", rows) + write_csv(output / "boundary_prompt_summary.csv", prompt_rows) + write_csv(output / "boundary_cosine_summary.csv", summary) + write_csv(output / "boundary_cosine_table.csv", table_rows) + config = { + "pairing": "current slot s vs previous chunk last slot; matched spatial coordinate", + "support": "target_chunk>=1, target_step>=1", + "prompt_first_aggregation": True, + "prompt_count_expected": 10, + "row_count": len(rows), + } + (output / "config.json").write_text(json.dumps(config, indent=2) + "\n", encoding="utf-8") + print(f"[complete] {output}: rows={len(rows)}", flush=True) + + +if __name__ == "__main__": + main() diff --git a/scripts/summarize_conditional_probe.py b/scripts/summarize_conditional_probe.py new file mode 100644 index 0000000000000000000000000000000000000000..77a9caf1e5464784967b8b146cf80bfe6f5a5842 --- /dev/null +++ b/scripts/summarize_conditional_probe.py @@ -0,0 +1,418 @@ +#!/usr/bin/env python3 +"""Summarize Linear/Ridge and nonlinear conditional probe fold results.""" + +from __future__ import annotations + +import argparse +import csv +import itertools +import json +import math +from collections import defaultdict +from pathlib import Path +from typing import Any + +import matplotlib + +matplotlib.use("Agg") +import matplotlib.pyplot as plt +import numpy as np + + +FAMILIES = ("self_forcing", "causal_forcing", "hy_worldplay") +ROLES = ("early", "middle", "late", "final") + + +def read_csv(path: Path) -> list[dict[str, Any]]: + with path.open(newline="", encoding="utf-8") as handle: + rows = list(csv.DictReader(handle)) + for row in rows: + for key in ("target_step", "held_out_prompt", "seed", "layer_index", "test_tokens"): + if key in row: + row[key] = int(row[key]) + for key in ("mse", "nMSE", "nRMSE", "r2", "cosine"): + row[key] = float(row[key]) + return rows + + +def write_csv(path: Path, rows: list[dict[str, Any]]) -> None: + if not rows: + return + path.parent.mkdir(parents=True, exist_ok=True) + 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, extrasaction="ignore") + writer.writeheader() + writer.writerows(rows) + + +def bootstrap(values: list[float], seed: int, rounds: int = 10000): + values = np.asarray(values, dtype=np.float64) + rng = np.random.default_rng(seed) + if values.size == 0: + return float("nan"), float("nan"), float("nan") + indices = rng.integers(0, values.size, size=(rounds, values.size)) + means = values[indices].mean(axis=1) + return float(values.mean()), float(np.quantile(means, 0.025)), float(np.quantile(means, 0.975)) + + +def exact_signflip(values: list[float]) -> float: + values = np.asarray(values, dtype=np.float64) + values = values[np.isfinite(values)] + if not values.size: + return float("nan") + observed = abs(float(values.mean())) + exceed = 0 + total = 1 << int(values.size) + for mask in range(total): + signed = np.asarray( + [value if (mask >> index) & 1 else -value for index, value in enumerate(values)] + ) + if abs(float(signed.mean())) >= observed - 1e-15: + exceed += 1 + return float((exceed + 1) / (total + 1)) + + +def prompt_metric_rows(rows: list[dict[str, Any]]) -> list[dict[str, Any]]: + """Average seeds within each prompt, keeping prompt as statistical unit.""" + grouped: dict[tuple, list[dict[str, Any]]] = defaultdict(list) + for row in rows: + grouped[ + ( + row["method"], row["model_family"], row["layer_role"], + row["target_step"], row["probe"], row["held_out_prompt"], + ) + ].append(row) + result = [] + for key, values in sorted(grouped.items(), key=lambda item: tuple(map(str, item[0]))): + method, family, role, step, probe, prompt = key + item = { + "method": method, + "model_family": family, + "layer_role": role, + "target_step": step, + "probe": probe, + "held_out_prompt": prompt, + "seed_count": len(values), + } + for metric in ("mse", "nMSE", "nRMSE", "r2", "cosine"): + item[metric] = float(np.mean([float(value[metric]) for value in values])) + result.append(item) + return result + + +def gain_rows(prompt_rows: list[dict[str, Any]]) -> list[dict[str, Any]]: + baseline_name = {"linear": "within_affine", "nonlinear": "step_only"} + grouped = defaultdict(dict) + for row in prompt_rows: + grouped[ + (row["method"], row["model_family"], row["layer_role"], row["target_step"], row["held_out_prompt"]) + ][row["probe"]] = row + result = [] + for key, probes in sorted(grouped.items(), key=lambda item: tuple(map(str, item[0]))): + method, family, role, step, prompt = key + baseline = probes.get(baseline_name[method]) + if baseline is None: + continue + for probe, current in probes.items(): + reference_mse = float(baseline["mse"]) + current_mse = float(current["mse"]) + result.append({ + "method": method, + "model_family": family, + "layer_role": role, + "target_step": step, + "held_out_prompt": prompt, + "baseline_probe": baseline_name[method], + "probe": probe, + "baseline_mse": reference_mse, + "probe_mse": current_mse, + "gain": (reference_mse - current_mse) / max(reference_mse, 1e-12), + "delta_r2": float(current["r2"]) - float(baseline["r2"]), + }) + return result + + +def summarize_gains(gains: list[dict[str, Any]]) -> list[dict[str, Any]]: + grouped = defaultdict(list) + for row in gains: + grouped[(row["method"], row["model_family"], row["layer_role"], row["target_step"], row["probe"])].append(row) + result = [] + for key, values in sorted(grouped.items(), key=lambda item: tuple(map(str, item[0]))): + method, family, role, step, probe = key + values = sorted(values, key=lambda row: row["held_out_prompt"]) + numbers = [float(row["gain"]) for row in values] + mean, low, high = bootstrap(numbers, 1000 + sum(ord(ch) for ch in str(key))) + result.append({ + "method": method, + "model_family": family, + "layer_role": role, + "target_step": step, + "probe": probe, + "baseline_probe": values[0]["baseline_probe"], + "prompt_count": len(numbers), + "gain_mean": mean, + "gain_ci95_low": low, + "gain_ci95_high": high, + "wins": int(sum(number > 0 for number in numbers)), + "signflip_p": exact_signflip(numbers), + }) + return result + + +def summarize_layer_gains(gains: list[dict[str, Any]]) -> list[dict[str, Any]]: + """Average target steps inside each prompt, then summarize across prompts.""" + prompt_groups = defaultdict(list) + for row in gains: + prompt_groups[ + ( + row["method"], row["model_family"], row["layer_role"], + row["probe"], row["held_out_prompt"], row["baseline_probe"], + ) + ].append(float(row["gain"])) + groups = defaultdict(list) + for key, values in prompt_groups.items(): + method, family, role, probe, _prompt, baseline = key + groups[(method, family, role, probe, baseline)].append(float(np.mean(values))) + result = [] + for key, values in sorted(groups.items(), key=lambda item: tuple(map(str, item[0]))): + method, family, role, probe, baseline = key + mean, low, high = bootstrap(values, 17000 + sum(ord(ch) for ch in str(key))) + result.append({ + "method": method, + "model_family": family, + "layer_role": role, + "probe": probe, + "baseline_probe": baseline, + "target_step_count": 3, + "prompt_count": len(values), + "gain_mean": mean, + "gain_ci95_low": low, + "gain_ci95_high": high, + "wins": int(sum(value > 0 for value in values)), + "signflip_p": exact_signflip(values), + }) + return result + + +def summarize_metrics(prompt_rows: list[dict[str, Any]]) -> list[dict[str, Any]]: + grouped = defaultdict(list) + for row in prompt_rows: + grouped[(row["method"], row["model_family"], row["layer_role"], row["target_step"], row["probe"])].append(row) + result = [] + for key, values in sorted(grouped.items(), key=lambda item: tuple(map(str, item[0]))): + method, family, role, step, probe = key + item = { + "method": method, + "model_family": family, + "layer_role": role, + "target_step": step, + "probe": probe, + "prompt_count": len(values), + } + for metric_index, metric in enumerate(("mse", "nMSE", "nRMSE", "r2", "cosine")): + numbers = [float(row[metric]) for row in values] + mean, low, high = bootstrap(numbers, 9000 + metric_index + sum(ord(ch) for ch in str(key))) + item[f"{metric}_mean"] = mean + item[f"{metric}_ci95_low"] = low + item[f"{metric}_ci95_high"] = high + result.append(item) + return result + + +def plot_primary(gains: list[dict[str, Any]], output: Path) -> None: + fig, axes = plt.subplots(1, 2, figsize=(14, 5), sharey=True) + for ax, method, title in zip(axes, ("linear", "nonlinear"), ("Linear/Ridge", "Nonlinear MLP")): + selected = [ + row for row in gains + if row["method"] == method and row["probe"] in ({"fusion_same"} if method == "linear" else {"both_correct"}) + ] + x = np.arange(len(ROLES)) + width = 0.24 + for family_index, family in enumerate(FAMILIES): + values = [] + lows = [] + highs = [] + for role in ROLES: + group = [row for row in selected if row["model_family"] == family and row["layer_role"] == role] + nums = [float(row["gain"]) for row in group] + mean, low, high = bootstrap(nums, 1234 + family_index * 100 + ROLES.index(role)) + values.append(mean) + lows.append(mean - low) + highs.append(high - mean) + pos = x + (family_index - (len(FAMILIES) - 1) / 2) * width + ax.bar(pos, values, width, yerr=[lows, highs], capsize=3, label=family) + ax.axhline(0, color="black", linewidth=0.8) + ax.set_xticks(x, ROLES) + ax.set_ylabel("Gain vs step-only baseline" if method == "nonlinear" else "Gain vs within-affine baseline") + ax.set_title(title) + ax.grid(axis="y", alpha=0.25) + axes[1].legend(fontsize=9) + fig.tight_layout() + fig.savefig(output, dpi=180) + plt.close(fig) + + +def plot_controls(gains: list[dict[str, Any]], output: Path) -> None: + probes = ["fusion_same", "fusion_step_duplicate", "fusion_wrong_step", "fusion_distant", "fusion_batch_shuffle", "fusion_zero", "fusion_noise"] + labels = { + "fusion_same": "correct", + "fusion_step_duplicate": "step duplicate", + "fusion_wrong_step": "wrong step", + "fusion_distant": "distant", + "fusion_batch_shuffle": "other video", + "fusion_zero": "zero", + "fusion_noise": "noise", + } + fig, axes = plt.subplots( + 1, + len(FAMILIES), + figsize=(5.7 * len(FAMILIES), 5), + sharey=True, + ) + axes = np.atleast_1d(axes) + for ax, family in zip(axes, FAMILIES): + values = [] + errors = [] + for probe in probes: + group = [row for row in gains if row["method"] == "linear" and row["model_family"] == family and row["layer_role"] == "final" and row["probe"] == probe] + nums = [float(row["gain"]) for row in group] + mean, low, high = bootstrap(nums, 4000 + probes.index(probe)) + values.append(mean) + errors.append((mean - low, high - mean)) + y = np.arange(len(probes)) + ax.errorbar(values, y, xerr=np.asarray(errors).T, fmt="o", capsize=3) + ax.axvline(0, color="black", linewidth=0.8) + ax.set_yticks(y, [labels[p] for p in probes]) + ax.set_title(family) + ax.grid(axis="x", alpha=0.25) + axes[0].set_xlabel("Linear MSE gain") + fig.tight_layout() + fig.savefig(output, dpi=180) + plt.close(fig) + + +def build_report(metrics: list[dict[str, Any]], gains: list[dict[str, Any]], output: Path, config: dict[str, Any]) -> None: + lines = [ + "# Conditional prediction and incremental chunk information", + "", + "This report uses 10 prompt-grouped held-out folds. The test prompt and its", + "other-video donor are excluded from training; nonlinear seeds are averaged", + "within prompt before confidence intervals are computed.", + "", + "The primary endpoint is `fusion_same` vs `within_affine` for Linear/Ridge", + "and `both_correct` vs `step_only` for the nonlinear MLP.", + "", + "| method | family | role | step | probe | gain | 95% CI | wins | sign-flip p |", + "|---|---|---|---:|---|---:|---|---:|---:|", + ] + primary = [ + row for row in gains + if (row["method"] == "linear" and row["probe"] == "fusion_same") + or (row["method"] == "nonlinear" and row["probe"] == "both_correct") + ] + for row in primary: + lines.append( + f"| {row['method']} | {row['model_family']} | {row['layer_role']} | {row['target_step']} | " + f"{row['probe']} | {row['gain_mean']:.4f} | [{row['gain_ci95_low']:.4f}, {row['gain_ci95_high']:.4f}] | " + f"{row['wins']}/{row['prompt_count']} | {row['signflip_p']:.4f} |" + ) + lines += [ + "", + "Interpretation: positive gain means that adding the auxiliary feature lowers held-out MSE.", + "The exact sign-flip test treats prompt, not token, as the independent unit.", + "Absolute MSE is not compared across model families because feature dimensions and", + "conditioning paths differ.", + "", + "## Files", + "", + "- `probe_folds_unified.csv`", + "- `probe_prompt_averaged.csv`", + "- `probe_metrics_summary.csv`", + "- `probe_gain_summary.csv`", + "- `conditional_gain_by_layer.png`", + "- `conditional_control_comparison.png`", + "", + "```json", + json.dumps(config, indent=2, ensure_ascii=False), + "```", + ] + output.write_text("\n".join(lines) + "\n", encoding="utf-8") + + +def main() -> None: + global FAMILIES + parser = argparse.ArgumentParser() + parser.add_argument("--linear_csv", type=Path, required=True) + parser.add_argument("--nonlinear_dir", type=Path, required=True) + parser.add_argument("--output_dir", type=Path, required=True) + parser.add_argument("--families", default=",".join(FAMILIES)) + parser.add_argument( + "--chunk_pairing", + choices=("matched_slot", "boundary_to_all"), + default="matched_slot", + ) + args = parser.parse_args() + requested_families = tuple( + value.strip() for value in args.families.split(",") if value.strip() + ) + unknown = set(requested_families) - set(FAMILIES) + if not requested_families or unknown: + raise ValueError(f"Invalid families: {requested_families}; unknown={sorted(unknown)}") + FAMILIES = requested_families + args.output_dir.mkdir(parents=True, exist_ok=True) + + linear = read_csv(args.linear_csv) + linear = [row for row in linear if row["model_family"] in FAMILIES] + linear = [{**row, "method": "linear"} for row in linear] + nonlinear = [] + for family in FAMILIES: + path = args.nonlinear_dir / f"nonlinear_probe_{family}_folds.csv" + rows = read_csv(path) + nonlinear.extend({**row, "method": "nonlinear"} for row in rows) + expected_linear = 1320 * len(FAMILIES) + expected_nonlinear = 2160 * len(FAMILIES) + if len(linear) != expected_linear: + raise ValueError(f"Expected {expected_linear} linear rows, found {len(linear)}") + if len(nonlinear) != expected_nonlinear: + raise ValueError(f"Expected {expected_nonlinear} nonlinear rows, found {len(nonlinear)}") + unified = linear + nonlinear + prompt_rows = prompt_metric_rows(unified) + gains = gain_rows(prompt_rows) + metrics = summarize_metrics(prompt_rows) + gain_summary = summarize_gains(gains) + layer_gain_summary = summarize_layer_gains(gains) + write_csv(args.output_dir / "probe_folds_unified.csv", unified) + write_csv(args.output_dir / "probe_prompt_averaged.csv", prompt_rows) + write_csv(args.output_dir / "probe_metrics_summary.csv", metrics) + write_csv(args.output_dir / "probe_gain_by_prompt.csv", gains) + write_csv(args.output_dir / "probe_gain_summary.csv", gain_summary) + write_csv(args.output_dir / "probe_layer_gain_summary.csv", layer_gain_summary) + plot_primary(gains, args.output_dir / "conditional_gain_by_layer.png") + plot_controls(gains, args.output_dir / "conditional_control_comparison.png") + config = { + "linear_rows": len(linear), + "nonlinear_rows": len(nonlinear), + "families": list(FAMILIES), + "chunk_pairing": args.chunk_pairing, + "prompt_averaged_rows": len(prompt_rows), + "gain_rows": len(gains), + "prompt_count": 10, + "target_chunks": [2, 3], + "target_steps": [1, 2, 3], + "linear_baseline": "within_affine", + "nonlinear_baseline": "step_only", + "nonlinear_seed_aggregation": "mean within held-out prompt", + "outer_split": "held-out prompt plus its cyclic other-video donor", + } + (args.output_dir / "config.json").write_text(json.dumps(config, indent=2) + "\n", encoding="utf-8") + build_report(metrics, gain_summary, args.output_dir / "REPORT.md", config) + print(f"[complete] {args.output_dir} unified={len(unified)} gains={len(gains)}", flush=True) + + +if __name__ == "__main__": + main() diff --git a/scripts/summarize_confidence_stage1_experiments.py b/scripts/summarize_confidence_stage1_experiments.py new file mode 100644 index 0000000000000000000000000000000000000000..1d0dc1c375909867214f194a140f9d4b9c6ac2d1 --- /dev/null +++ b/scripts/summarize_confidence_stage1_experiments.py @@ -0,0 +1,67 @@ +#!/usr/bin/env python3 +"""Create a compact report for the Stage-1 confidence-head rerun.""" + +from __future__ import annotations + +import argparse +import csv +import json +from pathlib import Path + + +def csv_rows(path: Path) -> list[dict[str, str]]: + with path.open(encoding="utf-8") as handle: + return list(csv.DictReader(handle)) + + +def selected(root: Path) -> list[dict]: + return json.loads((root / "validation/selected.json").read_text())["selected_dynamic"] + + +def main() -> None: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--root", type=Path, required=True) + args = parser.parse_args() + root = args.root.resolve() + step12_manifest = json.loads( + (root / "confidence_step12/manifest.json").read_text() + ) + report = { + "status": "complete", + "predictor": step12_manifest["predictor_weights"], + } + for name in ("step12", "step123"): + report[f"confidence_{name}"] = json.loads( + (root / f"confidence_{name}/metrics.json").read_text() + ) + dynamic_root = root / f"dynamic_{name}" + test = {row["config_name"]: row for row in csv_rows(dynamic_root / "test/summary.csv")} + chosen = [] + for row in selected(dynamic_root): + item = test[str(row["config_name"])].copy() + item["validation_beta"] = row["beta"] + item["validation_threshold"] = row["threshold"] + chosen.append(item) + report[f"dynamic_{name}"] = chosen + report[f"baselines_{name}"] = { + key: test[key] for key in ("ffff", "fppf" if name == "step12" else "fppp") + } + vbench_path = root / "vbench/summary.csv" + if vbench_path.exists(): + vbench_rows = csv_rows(vbench_path) + high_speed_path = root / "vbench_high_speed/summary.csv" + if high_speed_path.exists(): + by_condition = {row["condition"]: row for row in vbench_rows} + by_condition.update( + {row["condition"]: row for row in csv_rows(high_speed_path)} + ) + vbench_rows = [by_condition[name] for name in sorted(by_condition)] + report["vbench"] = vbench_rows + (root / "summary.json").write_text( + json.dumps(report, indent=2, ensure_ascii=False) + "\n", encoding="utf-8" + ) + print(json.dumps(report, indent=2, ensure_ascii=False)) + + +if __name__ == "__main__": + main() diff --git a/scripts/summarize_layer17_dynamic_gate.py b/scripts/summarize_layer17_dynamic_gate.py new file mode 100644 index 0000000000000000000000000000000000000000..8f4cfbcd7d296bcb7b85e2e09da789928b82cc6e --- /dev/null +++ b/scripts/summarize_layer17_dynamic_gate.py @@ -0,0 +1,185 @@ +#!/usr/bin/env python3 +"""Summarize validation-selected Layer-17 dynamic-gating test results.""" + +from __future__ import annotations + +import argparse +import csv +import json +from pathlib import Path +from typing import Any + +import matplotlib.pyplot as plt + + +TARGETS = (4, 6, 8, 10) + + +def read_csv(path: Path) -> list[dict[str, str]]: + with path.open(encoding="utf-8") as handle: + return list(csv.DictReader(handle)) + + +def write_csv(path: Path, rows: list[dict[str, Any]]) -> None: + fields = list(rows[0]) + with path.open("w", encoding="utf-8", newline="") as handle: + writer = csv.DictWriter(handle, fieldnames=fields) + writer.writeheader() + writer.writerows(rows) + + +def main() -> None: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument( + "--root", type=Path, + default=Path("outputs/layer17_dynamic_gate_20260830"), + ) + args = parser.parse_args() + root = args.root.resolve() + test_dir = root / "test" + summary_rows = read_csv(test_dir / "summary.csv") + summary = {row["config_name"]: row for row in summary_rows} + selected = json.loads((root / "validation" / "selected.json").read_text()) + selected_by_target = { + int(row["target_accepts"]): row for row in selected["selected_dynamic"] + } + ffff_time = float(summary["ffff"]["generation_time_s"]) + + comparisons: list[dict[str, Any]] = [] + acceptance_rows: list[dict[str, Any]] = [] + for target in TARGETS: + selected_row = selected_by_target[target] + dynamic_name = selected_row["config_name"] + dynamic = summary[dynamic_name] + static = summary[f"static_late_k{target:02d}"] + dynamic_lpips = float(dynamic["tail_lpips"]) + static_lpips = float(static["tail_lpips"]) + comparisons.append( + { + "target_accepts": target, + "dynamic_config": dynamic_name, + "beta": float(dynamic["beta"]), + "threshold": float(dynamic["threshold"]), + "dynamic_actual_accepts": float(dynamic["accepted_predictor_calls"]), + "static_actual_accepts": float(static["accepted_predictor_calls"]), + "dynamic_full_calls": float(dynamic["full_calls"]), + "static_full_calls": float(static["full_calls"]), + "dynamic_tail_lpips": dynamic_lpips, + "static_tail_lpips": static_lpips, + "tail_lpips_reduction_percent": 100.0 * (static_lpips - dynamic_lpips) / static_lpips, + "dynamic_generation_time_s": float(dynamic["generation_time_s"]), + "static_generation_time_s": float(static["generation_time_s"]), + "dynamic_overhead_vs_static_percent": 100.0 * ( + float(dynamic["generation_time_s"]) + / float(static["generation_time_s"]) + - 1.0 + ), + "dynamic_speedup_vs_ffff_percent": 100.0 * ( + 1.0 - float(dynamic["generation_time_s"]) / ffff_time + ), + "dynamic_lpips": float(dynamic["lpips"]), + "static_lpips": float(static["lpips"]), + "dynamic_latent_tail_nrmse": float(dynamic["latent_tail_nrmse"]), + "static_latent_tail_nrmse": float(static["latent_tail_nrmse"]), + } + ) + + decision_files = sorted((test_dir / "per_run" / dynamic_name).glob("*.json")) + decisions = [ + decision + for path in decision_files + for decision in json.loads(path.read_text())["decisions"] + ] + for chunk in range(1, 7): + for step in (1, 2): + cell = [ + row for row in decisions + if int(row["chunk"]) == chunk and int(row["step"]) == step + ] + acceptance_rows.append( + { + "target_accepts": target, + "dynamic_config": dynamic_name, + "chunk": chunk, + "step": step, + "acceptance_ratio": sum(bool(row["accepted"]) for row in cell) + / len(cell), + } + ) + write_csv(test_dir / "dynamic_vs_static.csv", comparisons) + write_csv(test_dir / "acceptance_by_chunk_step.csv", acceptance_rows) + + fig, axes = plt.subplots(1, 2, figsize=(11, 4.2)) + dynamic_rows = [summary[selected_by_target[target]["config_name"]] for target in TARGETS] + static_rows = [summary[f"static_late_k{target:02d}"] for target in TARGETS] + for axis, x_field, label in ( + (axes[0], "full_calls", "Mean Full calls"), + (axes[1], "generation_time_s", "Generation time (s)"), + ): + axis.plot( + [float(row[x_field]) for row in dynamic_rows], + [float(row["tail_lpips"]) for row in dynamic_rows], + "o-", label="Dynamic confidence", color="#d64b40", linewidth=2, + ) + axis.plot( + [float(row[x_field]) for row in static_rows], + [float(row["tail_lpips"]) for row in static_rows], + "s--", label="Static late-first", color="#3977b8", linewidth=2, + ) + axis.scatter( + [float(summary["ffff"][x_field])], + [float(summary["ffff"]["tail_lpips"])], + marker="*", s=100, color="#333333", label="FFFF", + ) + axis.scatter( + [float(summary["fppf"][x_field])], + [float(summary["fppf"]["tail_lpips"])], + marker="X", s=80, color="#777777", label="FPPF", + ) + axis.set_xlabel(label) + axis.set_ylabel("Tail LPIPS") + axis.grid(alpha=0.25) + axes[0].legend(frameon=False) + fig.suptitle("Layer-17 Predictor: quality-compute frontier on prompts 90–99") + fig.tight_layout() + fig.savefig(root / "quality_compute_pareto.png", dpi=180) + plt.close(fig) + + report = [ + "# Layer-17 dynamic confidence gating", + "", + "Thresholds and beta were selected only on prompts 80–89. The table below " + "reports the frozen configurations on prompts 90–99.", + "", + "| Target P | Beta | Actual P | Full | Tail LPIPS dynamic | Static | Reduction | Gen speedup vs FFFF |", + "|---:|---:|---:|---:|---:|---:|---:|---:|", + ] + for row in comparisons: + report.append( + f"| {row['target_accepts']} | {row['beta']:.1f} | " + f"{row['dynamic_actual_accepts']:.1f} | {row['dynamic_full_calls']:.1f} | " + f"{row['dynamic_tail_lpips']:.5f} | {row['static_tail_lpips']:.5f} | " + f"{row['tail_lpips_reduction_percent']:.1f}% | " + f"{row['dynamic_speedup_vs_ffff_percent']:.1f}% |" + ) + report.extend( + [ + "", + f"FFFF generation time: {float(summary['ffff']['generation_time_s']):.3f}s. " + f"FPPF generation time: {float(summary['fppf']['generation_time_s']):.3f}s; " + f"tail LPIPS: {float(summary['fppf']['tail_lpips']):.5f}.", + "", + "Dynamic gating evaluates the Predictor at all 12 candidate decisions, " + "including rejected calls. Its generation-time overhead relative to the " + "budget-matched static policies is 0–4.4%, and is included in the table/plot.", + "", + "The K≈6 point is the recommended balanced operating point: beta=1.0, " + "threshold=0.333097, 5.8 accepted Predictor calls, 22.2 Full calls, " + "tail LPIPS 0.03449, and 16.3% generation speedup versus FFFF.", + ] + ) + (root / "REPORT.md").write_text("\n".join(report) + "\n", encoding="utf-8") + + +if __name__ == "__main__": + main() diff --git a/scripts/summarize_layer17_dynamic_vbench.py b/scripts/summarize_layer17_dynamic_vbench.py new file mode 100644 index 0000000000000000000000000000000000000000..4eef26be6ce7bd33f31def7abc09d6e761f19926 --- /dev/null +++ b/scripts/summarize_layer17_dynamic_vbench.py @@ -0,0 +1,99 @@ +#!/usr/bin/env python3 +"""Split a combined VBench result into per-condition averages.""" + +from __future__ import annotations + +import argparse +import csv +import json +from collections import defaultdict +from pathlib import Path + + +def main() -> None: + parser = argparse.ArgumentParser(description=__doc__) + source = parser.add_mutually_exclusive_group(required=True) + source.add_argument("--results", type=Path) + source.add_argument( + "--result_root", type=Path, + help="Directory of per-condition subdirectories containing *_eval_results.json.", + ) + parser.add_argument("--output", type=Path, required=True) + args = parser.parse_args() + + if args.result_root is not None: + rows = [] + dimensions = None + for condition_dir in sorted(path for path in args.result_root.iterdir() if path.is_dir()): + matches = list(condition_dir.glob("*_eval_results.json")) + if len(matches) != 1: + raise ValueError( + f"Expected one eval result in {condition_dir}, got {matches}" + ) + raw_condition = json.loads(matches[0].read_text(encoding="utf-8")) + if dimensions is None: + dimensions = list(raw_condition) + elif list(raw_condition) != dimensions: + raise ValueError(f"Dimension mismatch in {matches[0]}") + item: dict[str, str | float | int] = { + "condition": condition_dir.name, + "num_videos": len(next(iter(raw_condition.values()))[1]), + } + for dimension, value in raw_condition.items(): + item[dimension] = float(value[0]) + rows.append(item) + if not rows or dimensions is None: + raise ValueError(f"No result directories under {args.result_root}") + args.output.parent.mkdir(parents=True, exist_ok=True) + with args.output.open("w", encoding="utf-8", newline="") as handle: + writer = csv.DictWriter( + handle, fieldnames=["condition", "num_videos", *dimensions] + ) + writer.writeheader() + writer.writerows(rows) + print(json.dumps(rows, indent=2)) + return + + raw = json.loads(args.results.read_text(encoding="utf-8")) + by_condition: dict[str, dict[str, list[float]]] = defaultdict( + lambda: defaultdict(list) + ) + for dimension, (_, video_rows) in raw.items(): + for row in video_rows: + filename = Path(row["video_path"]).name + if "__" not in filename: + raise ValueError(f"Missing condition prefix in {filename}") + condition = filename.split("__", 1)[0] + value = float(row["video_results"]) + if dimension == "imaging_quality": + value /= 100.0 + by_condition[condition][dimension].append(value) + + dimensions = list(raw) + rows = [] + for condition in sorted(by_condition): + item: dict[str, str | float | int] = {"condition": condition} + counts = set() + for dimension in dimensions: + values = by_condition[condition][dimension] + if not values: + raise ValueError(f"No {dimension} rows for {condition}") + counts.add(len(values)) + item[dimension] = sum(values) / len(values) + if len(counts) != 1: + raise ValueError(f"Dimension counts differ for {condition}: {counts}") + item["num_videos"] = counts.pop() + rows.append(item) + + args.output.parent.mkdir(parents=True, exist_ok=True) + with args.output.open("w", encoding="utf-8", newline="") as handle: + writer = csv.DictWriter( + handle, fieldnames=["condition", "num_videos", *dimensions] + ) + writer.writeheader() + writer.writerows(rows) + print(json.dumps(rows, indent=2)) + + +if __name__ == "__main__": + main() diff --git a/scripts/summarize_layer17_moviebench_step2000.py b/scripts/summarize_layer17_moviebench_step2000.py new file mode 100644 index 0000000000000000000000000000000000000000..96e79ae10822c12c43ae942d81aee5a31e99d673 --- /dev/null +++ b/scripts/summarize_layer17_moviebench_step2000.py @@ -0,0 +1,110 @@ +#!/usr/bin/env python3 +"""Aggregate MovieBench pixel metrics and prepare/read VBench-5 results.""" + +from __future__ import annotations + +import argparse +import csv +import json +import math +import os +from pathlib import Path + + +DIMENSIONS = [ + "subject_consistency", + "background_consistency", + "motion_smoothness", + "aesthetic_quality", + "imaging_quality", +] +METHODS = ("ffff", "fppf_step2000") + + +def atomic_json(path: Path, value) -> 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) + "\n", encoding="utf-8") + os.replace(temporary, path) + + +def main() -> None: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--output_dir", type=Path, required=True) + args = parser.parse_args() + rows = [json.loads(path.read_text()) for path in sorted((args.output_dir / "per_prompt").glob("*.json"))] + rows.sort(key=lambda row: row["prompt_id"]) + ids = [row["prompt_id"] for row in rows] + if ids != list(range(100)): + raise RuntimeError(f"Expected MovieBench IDs 0..99, got {ids}") + + mse = [value for row in rows for value in row["mse_per_frame"]] + ssim = [value for row in rows for value in row["ssim_per_frame"]] + lpips = [value for row in rows for value in row["lpips_per_frame"]] + rollout_mse = [value for row in rows for value in row["mse_per_frame"][9:]] + rollout_ssim = [value for row in rows for value in row["ssim_per_frame"][9:]] + rollout_lpips = [value for row in rows for value in row["lpips_per_frame"][9:]] + pixel_metrics = { + "status": "complete", + "checkpoint_step": 2000, + "source_layer": 17, + "num_prompts": 100, + "num_frames": len(mse), + "prompt_ids": ids, + "generation_prompts": "MovieGenVideoBench_extended.txt first 100", + "evaluation_prompts": "MovieGenVideoBench.txt first 100", + "schedule": "chunk0=FFFF; chunks1-6=FPPF", + "psnr": -10 * math.log10(sum(mse) / len(mse)), + "ssim": sum(ssim) / len(ssim), + "lpips": sum(lpips) / len(lpips), + "rollout_psnr": -10 * math.log10(sum(rollout_mse) / len(rollout_mse)), + "rollout_ssim": sum(rollout_ssim) / len(rollout_ssim), + "rollout_lpips": sum(rollout_lpips) / len(rollout_lpips), + "mean_ffff_generation_time_s": sum(row["ffff"]["generation_time_s"] for row in rows) / 100, + "mean_fppf_generation_time_s": sum(row["fppf"]["generation_time_s"] for row in rows) / 100, + } + atomic_json(args.output_dir / "pixel_metrics.json", pixel_metrics) + + for method in METHODS: + video_dir = args.output_dir / "vbench_inputs" / method + video_dir.mkdir(parents=True, exist_ok=True) + info = [] + for row in rows: + prompt_id = row["prompt_id"] + source = (args.output_dir / "videos" / method / f"{prompt_id:05d}.mp4").resolve() + if not source.exists(): + raise FileNotFoundError(source) + target = video_dir / source.name + if target.exists() or target.is_symlink(): + target.unlink() + target.symlink_to(source) + info.append( + { + "video_list": [target.name], + "prompt_en": row["original_prompt"], + "dimension": DIMENSIONS, + } + ) + atomic_json(video_dir / "full_info.json", info) + + vbench_rows = [] + for method in METHODS: + matches = list((args.output_dir / "vbench_results" / method).glob("*_eval_results.json")) + if not matches: + continue + if len(matches) != 1: + raise RuntimeError(f"Expected one VBench result for {method}: {matches}") + result = json.loads(matches[0].read_text()) + scores = {dimension: float(result[dimension][0]) for dimension in DIMENSIONS} + vbench_rows.append({"method": method, **scores, "vbench5_mean": sum(scores.values()) / 5}) + if vbench_rows: + atomic_json(args.output_dir / "vbench_summary.json", vbench_rows) + with (args.output_dir / "vbench_summary.csv").open("w", newline="", encoding="utf-8") as handle: + writer = csv.DictWriter(handle, fieldnames=list(vbench_rows[0])) + writer.writeheader() + writer.writerows(vbench_rows) + print(json.dumps({"pixel_metrics": pixel_metrics, "vbench": vbench_rows}, indent=2)) + + +if __name__ == "__main__": + main() diff --git a/scripts/summarize_long_video_eval.py b/scripts/summarize_long_video_eval.py new file mode 100644 index 0000000000000000000000000000000000000000..e82269b0c262a25fe464390d859a18cb5a69e73b --- /dev/null +++ b/scripts/summarize_long_video_eval.py @@ -0,0 +1,127 @@ +#!/usr/bin/env python3 +"""Merge long-video shards and prepare VBench custom-input directories.""" + +from __future__ import annotations + +import csv +import json +import math +import os +from pathlib import Path + + +ROOT = Path(__file__).resolve().parents[1] +OUTPUT = ROOT / "outputs/long_video_2x4x_eval" +LENGTHS = (42, 84) +PROMPTS = range(80, 100) +DIMENSIONS = [ + "subject_consistency", "background_consistency", "motion_smoothness", + "aesthetic_quality", "imaging_quality", +] + + +def atomic_json(path: Path, value) -> 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) + "\n") + os.replace(temporary, path) + + +def find_run(length: int, prompt: int) -> Path: + matches = list(OUTPUT.glob(f"shard_gpu*/latent_{length}/prompt_{prompt:04d}")) + if len(matches) != 1: + raise RuntimeError(f"Expected one run for latent={length}, prompt={prompt}: {matches}") + return matches[0] + + +def main() -> None: + summary = [] + per_prompt = [] + for length in LENGTHS: + results = [] + for prompt in PROMPTS: + run = find_run(length, prompt) + value = json.loads((run / "metrics.json").read_text()) + if value.get("status") != "complete": + raise RuntimeError(f"Incomplete: {run}") + results.append(value) + per_prompt.append({ + "latent_length": length, "prompt_id": prompt, + "decoded_frames": value["decoded_frames"], + "psnr": value["psnr"], "ssim": value["ssim"], + "lpips": value["lpips"], + "ffff_generation_time_s": value["ffff"]["generation_time_s"], + "fppf_generation_time_s": value["fppf"]["generation_time_s"], + }) + for method, filename in (("ffff", "ffff.mp4"), ("fppf", "fppf_layer17.mp4")): + target_dir = OUTPUT / "vbench_inputs" / f"latent_{length}" / method + target_dir.mkdir(parents=True, exist_ok=True) + target = target_dir / f"prompt_{prompt:04d}.mp4" + if target.is_symlink() or target.exists(): + target.unlink() + target.symlink_to((run / filename).resolve()) + + mse = [x for result in results for x in result["mse_per_frame"]] + ssim = [x for result in results for x in result["ssim_per_frame"]] + lpips = [x for result in results for x in result["lpips_per_frame"]] + row = { + "latent_length": length, + "decoded_frames_per_video": results[0]["decoded_frames"], + "num_prompts": len(results), + "psnr": -10.0 * math.log10(sum(mse) / len(mse)), + "ssim": sum(ssim) / len(ssim), + "lpips": sum(lpips) / len(lpips), + "mean_ffff_generation_time_s": sum(r["ffff"]["generation_time_s"] for r in results) / len(results), + "mean_fppf_generation_time_s": sum(r["fppf"]["generation_time_s"] for r in results) / len(results), + "ffff_calls": results[0]["ffff"]["full_calls"], + "fppf_full_calls": results[0]["fppf"]["full_calls"], + "fppf_predictor_calls": results[0]["fppf"]["predictor_calls"], + } + summary.append(row) + + for method in ("ffff", "fppf"): + video_dir = OUTPUT / "vbench_inputs" / f"latent_{length}" / method + info = [{ + "video_list": [f"prompt_{result['prompt_id']:04d}.mp4"], + "prompt_en": result["prompt"], "dimension": DIMENSIONS, + } for result in results] + atomic_json(video_dir / "full_info.json", info) + + for name, rows in (("summary.csv", summary), ("per_prompt.csv", per_prompt)): + destination = OUTPUT / name + with destination.open("w", newline="", encoding="utf-8") as handle: + writer = csv.DictWriter(handle, fieldnames=list(rows[0])) + writer.writeheader() + writer.writerows(rows) + atomic_json(OUTPUT / "summary.json", summary) + + vbench_rows = [] + for length in LENGTHS: + for method in ("ffff", "fppf"): + result_path = ( + OUTPUT / "vbench_results" / f"latent_{length}" / method + / f"latent_{length}_{method}_eval_results.json" + ) + if not result_path.exists(): + continue + result = json.loads(result_path.read_text()) + scores = {dimension: result[dimension][0] for dimension in DIMENSIONS} + vbench_rows.append({ + "latent_length": length, + "method": method.upper(), + **scores, + "vbench5_mean": sum(scores.values()) / len(scores), + }) + if vbench_rows: + with (OUTPUT / "vbench_summary.csv").open( + "w", newline="", encoding="utf-8" + ) as handle: + writer = csv.DictWriter(handle, fieldnames=list(vbench_rows[0])) + writer.writeheader() + writer.writerows(vbench_rows) + atomic_json(OUTPUT / "vbench_summary.json", vbench_rows) + print(json.dumps(summary, indent=2)) + + +if __name__ == "__main__": + main() diff --git a/scripts/summarize_probe_outlier_sensitivity.py b/scripts/summarize_probe_outlier_sensitivity.py new file mode 100644 index 0000000000000000000000000000000000000000..5c376476e4b3734c14e29f715503224d0ede528c --- /dev/null +++ b/scripts/summarize_probe_outlier_sensitivity.py @@ -0,0 +1,142 @@ +#!/usr/bin/env python3 +"""Prompt-level sensitivity summary for explicitly identified probe outliers.""" + +from __future__ import annotations + +import argparse +import csv +from collections import defaultdict +from pathlib import Path + +import numpy as np + + +def bootstrap(values: list[float], seed: int, rounds: int = 10_000) -> tuple[float, float, float]: + array = np.asarray(values, dtype=np.float64) + rng = np.random.default_rng(seed) + indices = rng.integers(0, array.size, size=(rounds, array.size)) + means = array[indices].mean(axis=1) + return ( + float(array.mean()), + float(np.quantile(means, 0.025)), + float(np.quantile(means, 0.975)), + ) + + +def exact_signflip(values: list[float]) -> float: + array = np.asarray(values, dtype=np.float64) + observed = abs(float(array.mean())) + exceed = 0 + total = 1 << int(array.size) + for mask in range(total): + signs = np.asarray([1.0 if (mask >> index) & 1 else -1.0 for index in range(array.size)]) + if abs(float((array * signs).mean())) >= observed - 1e-15: + exceed += 1 + return float((exceed + 1) / (total + 1)) + + +def parse_pairs(value: str) -> set[tuple[int, int]]: + result: set[tuple[int, int]] = set() + for item in value.split(","): + if not item.strip(): + continue + prompt, step = item.split(":", maxsplit=1) + result.add((int(prompt), int(step))) + return result + + +def summarize( + rows: list[dict[str, str]], + scenario: str, + excluded_folds: set[tuple[int, int]], + excluded_prompts: set[int], + seed: int, +) -> tuple[dict[str, object], list[dict[str, object]]]: + grouped: dict[int, list[float]] = defaultdict(list) + for row in rows: + prompt = int(row["held_out_prompt"]) + step = int(row["target_step"]) + if prompt in excluded_prompts or (prompt, step) in excluded_folds: + continue + grouped[prompt].append(float(row["gain"])) + + prompt_rows = [ + { + "scenario": scenario, + "held_out_prompt": prompt, + "retained_step_count": len(grouped[prompt]), + "gain": float(np.mean(grouped[prompt])), + } + for prompt in sorted(grouped) + ] + values = [float(row["gain"]) for row in prompt_rows] + mean, low, high = bootstrap(values, seed) + summary = { + "scenario": scenario, + "prompt_count": len(values), + "fold_count": sum(int(row["retained_step_count"]) for row in prompt_rows), + "gain_mean": mean, + "gain_ci95_low": low, + "gain_ci95_high": high, + "wins": sum(value > 0 for value in values), + "signflip_p": exact_signflip(values), + "excluded_folds": ";".join(f"{prompt}:{step}" for prompt, step in sorted(excluded_folds)), + "excluded_prompts": ";".join(map(str, sorted(excluded_prompts))), + } + return summary, prompt_rows + + +def write_csv(path: Path, rows: list[dict[str, object]]) -> None: + path.parent.mkdir(parents=True, exist_ok=True) + with path.open("w", newline="", encoding="utf-8") as handle: + writer = csv.DictWriter(handle, fieldnames=list(rows[0])) + writer.writeheader() + writer.writerows(rows) + + +def main() -> None: + parser = argparse.ArgumentParser() + parser.add_argument("--gain_csv", type=Path, required=True) + parser.add_argument("--output_dir", type=Path, required=True) + parser.add_argument("--method", default="linear") + parser.add_argument("--family", default="causal_forcing") + parser.add_argument("--role", default="late") + parser.add_argument("--probe", default="fusion_same") + parser.add_argument("--exclude_folds", default="6:1,9:2") + args = parser.parse_args() + + with args.gain_csv.open(newline="", encoding="utf-8") as handle: + rows = [ + row + for row in csv.DictReader(handle) + if row["method"] == args.method + and row["model_family"] == args.family + and row["layer_role"] == args.role + and row["probe"] == args.probe + ] + if not rows: + raise ValueError("No rows matched the requested probe endpoint") + + excluded_folds = parse_pairs(args.exclude_folds) + excluded_prompt_ids = {prompt for prompt, _step in excluded_folds} + seed_key = (args.method, args.family, args.role, args.probe, rows[0]["baseline_probe"]) + seed = 17_000 + sum(ord(character) for character in str(seed_key)) + scenarios = [ + ("all_folds", set(), set()), + ("exclude_two_folds", excluded_folds, set()), + ("exclude_two_prompts", set(), excluded_prompt_ids), + ] + summaries: list[dict[str, object]] = [] + prompt_rows: list[dict[str, object]] = [] + for scenario, folds, prompts in scenarios: + summary, per_prompt = summarize(rows, scenario, folds, prompts, seed) + summaries.append(summary) + prompt_rows.extend(per_prompt) + + write_csv(args.output_dir / "causal_late_sensitivity_summary.csv", summaries) + write_csv(args.output_dir / "causal_late_sensitivity_by_prompt.csv", prompt_rows) + print(f"[complete] {args.output_dir.resolve()}") + + +if __name__ == "__main__": + main() diff --git a/scripts/summarize_timestep_chunk_cosine.py b/scripts/summarize_timestep_chunk_cosine.py new file mode 100644 index 0000000000000000000000000000000000000000..41b4f57239747948d978a8668ef4f5110987b453 --- /dev/null +++ b/scripts/summarize_timestep_chunk_cosine.py @@ -0,0 +1,674 @@ +#!/usr/bin/env python3 +"""Merge and visualize the six 50-step/4-step cosine experiments.""" + +from __future__ import annotations + +import argparse +import csv +import json +import math +import re +from collections import defaultdict +from pathlib import Path +from typing import Any, Iterable + +import matplotlib + +matplotlib.use("Agg") +import matplotlib.pyplot as plt +import numpy as np + + +ROLE_BY_VARIANT = { + ("self_forcing", "wan14b50"): {9: "early", 19: "middle", 29: "late", 39: "final"}, + ("self_forcing", "dmd4"): {7: "early", 14: "middle", 22: "late", 29: "final"}, + ("causal_forcing", "ar50"): {7: "early", 14: "middle", 22: "late", 29: "final"}, + ("causal_forcing", "dmd4"): {7: "early", 14: "middle", 22: "late", 29: "final"}, + ("hy_worldplay", "ar50"): {13: "early", 26: "middle", 40: "late", 53: "final"}, + ("hy_worldplay", "ar4"): {13: "early", 26: "middle", 40: "late", 53: "final"}, +} +ROLE_ORDER = ["early", "middle", "late", "final"] +FAMILIES = ["self_forcing", "causal_forcing", "hy_worldplay"] +VARIANT_ORDER = { + "self_forcing": ["wan14b50", "dmd4"], + "causal_forcing": ["ar50", "dmd4"], + "hy_worldplay": ["ar50", "ar4"], +} +STEP_COUNT = {"wan14b50": 50, "ar50": 50, "dmd4": 4, "ar4": 4} +EXPECTED_ROWS = { + ("self_forcing", "wan14b50"): 26560, + ("self_forcing", "dmd4"): 1800, + ("causal_forcing", "ar50"): 26560, + ("causal_forcing", "dmd4"): 1800, + ("hy_worldplay", "ar50"): 10240, + ("hy_worldplay", "ar4"): 680, +} + + +def read_csv(path: Path) -> list[dict[str, str]]: + with path.open(newline="", encoding="utf-8") as handle: + return list(csv.DictReader(handle)) + + +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) + path.parent.mkdir(parents=True, exist_ok=True) + with path.open("w", newline="", encoding="utf-8") as handle: + writer = csv.DictWriter(handle, fieldnames=fields, extrasaction="ignore") + writer.writeheader() + writer.writerows(rows) + + +def normalize( + family: str, + variant: str, + prompt_id: int, + layer: int, + comparison: str, + reference_chunk: int, + target_chunk: int, + reference_step: int, + target_step: int, + reference_timestep: float, + target_timestep: float, + cosine: float, + p10: float, + p50: float, + p90: float, + chunk_semantics: str, +) -> dict[str, Any]: + role = ROLE_BY_VARIANT[(family, variant)][layer] + return { + "model_family": family, + "model_variant": variant, + "step_count": STEP_COUNT[variant], + "layer_role": role, + "layer_index": layer, + "prompt_id": prompt_id, + "comparison": comparison, + "reference_chunk": reference_chunk, + "target_chunk": target_chunk, + "reference_step": reference_step, + "target_step": target_step, + "reference_timestep": reference_timestep, + "target_timestep": target_timestep, + "normalized_progress": target_step / max(1, STEP_COUNT[variant] - 1), + "token_cosine_mean": cosine, + "token_cosine_p10": p10, + "token_cosine_p50": p50, + "token_cosine_p90": p90, + "chunk_semantics": chunk_semantics, + } + + +def load_self_dmd(path: Path) -> list[dict[str, Any]]: + timesteps = [1000.0, 937.5, 833.3333129882812, 625.0] + result = [] + for row in read_csv(path): + match = re.fullmatch(r"block_(\d+)_hidden", row["stage"]) + if match is None or row["comparison"] not in {"within_adjacent", "cross_same"}: + continue + reference_step = int(row["reference_step"]) + target_step = int(row["target_step"]) + result.append(normalize( + "self_forcing", "dmd4", int(row["run"]), int(match.group(1)), + row["comparison"], int(row["reference_chunk"]), int(row["target_chunk"]), + reference_step, target_step, timesteps[reference_step], timesteps[target_step], + float(row["token_cosine_mean"]), float(row["token_cosine_p10"]), + float(row["token_cosine_p50"]), float(row["token_cosine_p90"]), + "ar_previous_chunk", + )) + return result + + +def load_direct(paths: Iterable[Path], family: str, variant: str) -> list[dict[str, Any]]: + result = [] + timestep_by_step: dict[int, float] = {} + raw_rows: list[dict[str, str]] = [] + for path in paths: + raw_rows.extend(read_csv(path)) + for row in raw_rows: + timestep_by_step[int(row["target_step"])] = float(row["target_timestep"]) + for row in raw_rows: + comparison = row["comparison"] + if comparison not in {"within_adjacent", "within_matched_gap", "cross_same"}: + continue + reference_step = int(row["reference_step"]) + target_step = int(row["target_step"]) + reference_timestep = ( + float(row["reference_timestep"]) + if row.get("reference_timestep") not in {None, ""} + else timestep_by_step[reference_step] + ) + result.append(normalize( + family, variant, int(row["prompt_id"]), int(row["layer_index"]), comparison, + int(row["reference_chunk"]), int(row["chunk"]), reference_step, target_step, + reference_timestep, float(row["target_timestep"]), + float(row["token_cosine_mean"]), float(row["token_cosine_p10"]), + float(row["token_cosine_p50"]), float(row["token_cosine_p90"]), + row.get("chunk_semantics") or "ar_previous_chunk", + )) + return result + + +def load_hy(paths: Iterable[Path], variant: str) -> list[dict[str, Any]]: + result = [] + for path in paths: + for row in read_csv(path): + match = re.fullmatch(r"block_(\d+)", row["stage"]) + if match is None: + continue + reference_step = int(row["reference_step"]) + target_step = int(row["target_step"]) + if row["comparison"] == "cross_chunk_same_timestep": + comparison = "cross_same" + elif row["comparison"] == "intra_chunk": + comparison = ( + "within_adjacent" + if target_step == reference_step + 1 + else "within_matched_gap" + ) + else: + continue + prompt_match = re.search(r"prompt_(\d+)", row["case"]) + if prompt_match is None: + raise ValueError(f"Cannot parse prompt id from {row['case']}") + result.append(normalize( + "hy_worldplay", variant, int(prompt_match.group(1)), int(match.group(1)), + comparison, int(row["reference_chunk"]), int(row["target_chunk"]), + reference_step, target_step, float(row["reference_timestep"]), + float(row["target_timestep"]), float(row["token_cosine_mean"]), + float(row["token_cosine_p10"]), float(row["token_cosine_p50"]), + float(row["token_cosine_p90"]), "ar_previous_chunk", + )) + return result + + +def group_mean(rows: list[dict[str, Any]], keys: list[str]) -> list[dict[str, Any]]: + groups: dict[tuple[Any, ...], list[float]] = defaultdict(list) + for row in rows: + groups[tuple(row[key] for key in keys)].append(float(row["token_cosine_mean"])) + result = [] + for key, values in groups.items(): + result.append({ + **dict(zip(keys, key)), + "count": len(values), + "token_cosine_mean": float(np.mean(values)), + "token_cosine_std": float(np.std(values)), + }) + return result + + +def bootstrap_summary( + prompt_rows: list[dict[str, Any]], keys: list[str] | None = None +) -> list[dict[str, Any]]: + if keys is None: + keys = ["model_family", "model_variant", "layer_role", "comparison"] + groups: dict[tuple[Any, ...], list[float]] = defaultdict(list) + for row in prompt_rows: + groups[tuple(row[key] for key in keys)].append(float(row["token_cosine_mean"])) + rng = np.random.default_rng(20260826) + result = [] + for key, values_list in groups.items(): + values = np.asarray(values_list, dtype=np.float64) + draws = values[rng.integers(0, len(values), size=(10000, len(values)))].mean(axis=1) + result.append({ + **dict(zip(keys, key)), + "prompt_count": len(values), + "mean": float(values.mean()), + "std_across_prompts": float(values.std(ddof=1)) if len(values) > 1 else 0.0, + "ci95_low": float(np.quantile(draws, 0.025)), + "ci95_high": float(np.quantile(draws, 0.975)), + }) + return result + + +def plot_overview(summary: list[dict[str, Any]], path: Path) -> None: + lookup = { + (row["model_family"], row["model_variant"], row["layer_role"], row["comparison"]): row + for row in summary + if row["comparison"] in {"within_adjacent", "cross_same"} + } + fig, axes = plt.subplots(3, 4, figsize=(17, 10), sharey=True) + colors = {"within_adjacent": "#3B82F6", "cross_same": "#F59E0B"} + for family_index, family in enumerate(FAMILIES): + variants = VARIANT_ORDER[family] + for role_index, role in enumerate(ROLE_ORDER): + ax = axes[family_index, role_index] + positions = np.arange(2) + width = 0.34 + for comparison_index, comparison in enumerate(["within_adjacent", "cross_same"]): + means, lows, highs = [], [], [] + for variant in variants: + row = lookup[(family, variant, role, comparison)] + means.append(row["mean"]) + lows.append(row["mean"] - row["ci95_low"]) + highs.append(row["ci95_high"] - row["mean"]) + ax.bar( + positions + (comparison_index - 0.5) * width, + means, + width, + color=colors[comparison], + label=comparison if family_index == 0 and role_index == 0 else None, + yerr=np.asarray([lows, highs]), + capsize=3, + ) + ax.set_xticks(positions, variants, rotation=15) + ax.set_ylim(0, 1.02) + ax.grid(axis="y", alpha=0.25) + if family_index == 0: + ax.set_title(role) + if role_index == 0: + ax.set_ylabel(family.replace("_", " ") + "\ncosine") + fig.legend(loc="upper center", ncol=2) + fig.suptitle("Native adjacent timestep vs previous chunk at the same timestep", y=1.01) + fig.tight_layout() + fig.savefig(path, dpi=180, bbox_inches="tight") + plt.close(fig) + + +def plot_matched(rows: list[dict[str, Any]], path: Path) -> None: + selected = [ + row for row in rows + if (row["step_count"] == 4 and row["comparison"] == "within_adjacent") + or (row["step_count"] == 50 and row["comparison"] == "within_matched_gap") + ] + prompt_transition = group_mean( + selected, + ["model_family", "model_variant", "layer_role", "prompt_id", "reference_step", "target_step"], + ) + transition_groups: dict[tuple[str, str, str, int], list[float]] = defaultdict(list) + for row in prompt_transition: + transition_id = ( + int(row["target_step"]) - 1 + if STEP_COUNT[row["model_variant"]] == 4 + else {12: 0, 25: 1, 37: 2}[int(row["target_step"])] + ) + transition_groups[(row["model_family"], row["model_variant"], row["layer_role"], transition_id)].append( + float(row["token_cosine_mean"]) + ) + fig, axes = plt.subplots(3, 4, figsize=(17, 10), sharex=True, sharey=True) + for family_index, family in enumerate(FAMILIES): + for role_index, role in enumerate(ROLE_ORDER): + ax = axes[family_index, role_index] + for variant, marker in zip(VARIANT_ORDER[family], ["o", "s"]): + means = [ + np.mean(transition_groups[(family, variant, role, transition)]) + for transition in range(3) + ] + ax.plot(range(3), means, marker=marker, linewidth=2, label=variant) + ax.set_xticks(range(3), ["high", "middle", "low"]) + ax.set_ylim(0, 1.02) + ax.grid(alpha=0.25) + if family_index == 0: + ax.set_title(role) + if role_index == 0: + ax.set_ylabel(family.replace("_", " ") + "\ncosine") + if family_index == 0 and role_index == 0: + ax.legend() + fig.suptitle("Matched noise-gap timestep cosine: 4-step transitions vs 50-step endpoints") + fig.tight_layout() + fig.savefig(path, dpi=180, bbox_inches="tight") + plt.close(fig) + + +def plot_layer_curves(summary: list[dict[str, Any]], path: Path) -> None: + lookup = { + (row["model_family"], row["model_variant"], row["layer_role"], row["comparison"]): row + for row in summary + } + fig, axes = plt.subplots(1, 3, figsize=(16, 4.6), sharey=True) + colors = {"within_adjacent": "#2563EB", "cross_same": "#D97706"} + for ax, family in zip(axes, FAMILIES): + for variant, marker in zip(VARIANT_ORDER[family], ["o", "s"]): + for comparison, linestyle in zip(["within_adjacent", "cross_same"], ["-", "--"]): + points = [lookup[(family, variant, role, comparison)] for role in ROLE_ORDER] + ax.plot( + ROLE_ORDER, + [point["mean"] for point in points], + marker=marker, + linestyle=linestyle, + linewidth=2, + color=colors[comparison], + label=f"{variant} / {comparison}", + ) + ax.set_title(family.replace("_", " ")) + ax.set_ylim(0, 1.02) + ax.grid(alpha=0.25) + ax.tick_params(axis="x", rotation=20) + axes[0].set_ylabel("token-wise cosine") + axes[-1].legend(fontsize=8, loc="lower right") + fig.suptitle("Layer-depth cosine curves on common chunk/step support") + fig.tight_layout() + fig.savefig(path, dpi=180, bbox_inches="tight") + plt.close(fig) + + +def plot_timestep_curves(by_step: list[dict[str, Any]], path: Path) -> None: + fig, axes = plt.subplots(3, 2, figsize=(15, 11), sharey=True) + role_colors = dict(zip(ROLE_ORDER, ["#2563EB", "#059669", "#D97706", "#DC2626"])) + for family_index, family in enumerate(FAMILIES): + for variant_index, variant in enumerate(VARIANT_ORDER[family]): + ax = axes[family_index, variant_index] + for role in ROLE_ORDER: + for comparison, linestyle in zip(["within_adjacent", "cross_same"], ["-", "--"]): + points = sorted( + [ + row for row in by_step + if row["model_family"] == family + and row["model_variant"] == variant + and row["layer_role"] == role + and row["comparison"] == comparison + ], + key=lambda row: int(row["target_step"]), + ) + ax.plot( + [int(point["target_step"]) / (STEP_COUNT[variant] - 1) for point in points], + [float(point["token_cosine_mean"]) for point in points], + color=role_colors[role], + linestyle=linestyle, + linewidth=1.6, + label=f"{role} / {comparison}" if family_index == 0 and variant_index == 0 else None, + ) + ax.set_title(f"{family.replace('_', ' ')} / {variant}") + ax.set_xlim(0, 1) + ax.set_ylim(0, 1.02) + ax.grid(alpha=0.25) + ax.set_xlabel("normalized denoising-step index") + if variant_index == 0: + ax.set_ylabel("token-wise cosine") + axes[0, 0].legend(fontsize=7, ncol=2, loc="lower left") + fig.suptitle("Cosine evolution over denoising steps (solid: within, dashed: cross chunk)") + fig.tight_layout() + fig.savefig(path, dpi=180, bbox_inches="tight") + plt.close(fig) + + +def plot_heatmaps(rows: list[dict[str, Any]], output_dir: Path) -> None: + for family in FAMILIES: + fig, axes = plt.subplots(4, 4, figsize=(19, 11)) + variants = VARIANT_ORDER[family] + columns = [ + (variants[0], "within_adjacent"), + (variants[0], "cross_same"), + (variants[1], "within_adjacent"), + (variants[1], "cross_same"), + ] + image = None + for role_index, role in enumerate(ROLE_ORDER): + for column_index, (variant, comparison) in enumerate(columns): + ax = axes[role_index, column_index] + subset = [ + row for row in rows + if row["model_family"] == family + and row["model_variant"] == variant + and row["layer_role"] == role + and row["comparison"] == comparison + ] + chunks = sorted({int(row["target_chunk"]) for row in subset}) + steps = sorted({int(row["target_step"]) for row in subset}) + values: dict[tuple[int, int], list[float]] = defaultdict(list) + for row in subset: + values[(int(row["target_chunk"]), int(row["target_step"]))].append( + float(row["token_cosine_mean"]) + ) + matrix = np.full((len(chunks), len(steps)), np.nan) + for i, chunk in enumerate(chunks): + for j, step in enumerate(steps): + if values[(chunk, step)]: + matrix[i, j] = np.mean(values[(chunk, step)]) + image = ax.imshow(matrix, aspect="auto", vmin=0, vmax=1, cmap="viridis") + if role_index == 0: + ax.set_title(f"{variant}\n{comparison}") + if column_index == 0: + ax.set_ylabel(f"{role}\nchunk") + if role_index == 3: + ax.set_xlabel("denoising step") + if len(steps) <= 6: + ax.set_xticks(range(len(steps)), steps) + else: + tick_positions = np.linspace(0, len(steps) - 1, 6).round().astype(int) + ax.set_xticks(tick_positions, [steps[index] for index in tick_positions]) + ax.set_yticks(range(len(chunks)), chunks) + if image is not None: + fig.colorbar(image, ax=axes, shrink=0.72, label="token-wise cosine") + fig.suptitle(f"{family}: chunk × timestep redundancy maps") + fig.subplots_adjust(left=0.06, right=0.91, top=0.90, bottom=0.07, wspace=0.25, hspace=0.35) + fig.savefig(output_dir / f"{family}_chunk_timestep_heatmaps.png", dpi=180) + plt.close(fig) + + +def validate(rows: list[dict[str, Any]]) -> dict[str, Any]: + result = {} + for family in FAMILIES: + for variant in VARIANT_ORDER[family]: + subset = [row for row in rows if row["model_family"] == family and row["model_variant"] == variant] + prompt_ids = sorted({int(row["prompt_id"]) for row in subset}) + roles = sorted({row["layer_role"] for row in subset}, key=ROLE_ORDER.index) + result[f"{family}/{variant}"] = { + "rows": len(subset), + "prompt_ids": prompt_ids, + "roles": roles, + "comparisons": sorted({row["comparison"] for row in subset}), + "chunks": sorted({int(row["target_chunk"]) for row in subset}), + "steps": sorted({int(row["target_step"]) for row in subset}), + "finite_cosine": all( + math.isfinite(float(row["token_cosine_mean"])) for row in subset + ), + } + expected_comparisons = {"within_adjacent", "cross_same"} + if STEP_COUNT[variant] == 50: + expected_comparisons.add("within_matched_gap") + details = result[f"{family}/{variant}"] + if ( + prompt_ids != list(range(10)) + or roles != ROLE_ORDER + or len(subset) != EXPECTED_ROWS[(family, variant)] + or set(details["comparisons"]) != expected_comparisons + or details["steps"] != list(range(STEP_COUNT[variant])) + or not details["finite_cosine"] + ): + raise ValueError(f"Incomplete data for {family}/{variant}: {result[f'{family}/{variant}']}") + return result + + +def main() -> None: + parser = argparse.ArgumentParser() + parser.add_argument( + "--input_root", + type=Path, + required=True, + help="Feature-metric root on fast local storage.", + ) + parser.add_argument("--output_dir", type=Path, required=True) + args = parser.parse_args() + root = args.input_root.resolve() + output_dir = args.output_dir.resolve() + output_dir.mkdir(parents=True, exist_ok=True) + rows: list[dict[str, Any]] = [] + rows.extend(load_self_dmd(root / "self_forcing/dmd4/feature_pair_metrics.csv")) + rows.extend(load_direct( + sorted((root / "self_forcing/wan14b50").glob("shard_gpu*/feature_pair_metrics.csv")), + "self_forcing", "wan14b50", + )) + rows.extend(load_direct( + [root / "causal_forcing/dmd4/feature_pair_metrics.csv"], + "causal_forcing", "dmd4", + )) + rows.extend(load_direct( + sorted((root / "causal_forcing/ar50").glob("shard_gpu*/feature_pair_metrics.csv")), + "causal_forcing", "ar50", + )) + rows.extend(load_hy( + [root / "hy_worldplay/ar4/aggregate/feature_pair_metrics.csv"], "ar4" + )) + rows.extend(load_hy( + sorted((root / "hy_worldplay/ar50").glob("shard_gpu*/aggregate/feature_pair_metrics.csv")), + "ar50", + )) + validation = validate(rows) + write_csv(output_dir / "feature_pair_metrics_unified.csv", rows) + common_support_rows = [ + row for row in rows + if row["comparison"] in {"within_adjacent", "cross_same"} + and int(row["target_chunk"]) >= 1 + and int(row["target_step"]) >= 1 + ] + prompt_rows = group_mean( + common_support_rows, + ["model_family", "model_variant", "layer_role", "comparison", "prompt_id"], + ) + write_csv(output_dir / "feature_pair_prompt_summary.csv", prompt_rows) + summary = bootstrap_summary(prompt_rows) + write_csv(output_dir / "feature_pair_summary.csv", summary) + prompt_lookup = { + ( + row["model_family"], row["model_variant"], row["layer_role"], + int(row["prompt_id"]), row["comparison"], + ): float(row["token_cosine_mean"]) + for row in prompt_rows + } + delta_rows = [] + for family in FAMILIES: + for variant in VARIANT_ORDER[family]: + for role in ROLE_ORDER: + for prompt_id in range(10): + prefix = (family, variant, role, prompt_id) + delta_rows.append({ + "model_family": family, + "model_variant": variant, + "layer_role": role, + "prompt_id": prompt_id, + "token_cosine_mean": ( + prompt_lookup[(*prefix, "within_adjacent")] + - prompt_lookup[(*prefix, "cross_same")] + ), + }) + write_csv(output_dir / "within_minus_cross_by_prompt.csv", delta_rows) + delta_summary = bootstrap_summary( + delta_rows, ["model_family", "model_variant", "layer_role"] + ) + write_csv(output_dir / "within_minus_cross_summary.csv", delta_summary) + by_step = group_mean( + rows, + ["model_family", "model_variant", "layer_role", "comparison", "target_step", "target_timestep"], + ) + write_csv(output_dir / "feature_pair_by_step.csv", by_step) + matched = group_mean( + [ + row for row in rows + if (row["step_count"] == 4 and row["comparison"] == "within_adjacent") + or (row["step_count"] == 50 and row["comparison"] == "within_matched_gap") + ], + [ + "model_family", "model_variant", "layer_role", "prompt_id", + "reference_step", "target_step", "reference_timestep", "target_timestep", + ], + ) + write_csv(output_dir / "matched_gap_prompt_summary.csv", matched) + matched_summary = bootstrap_summary( + matched, + [ + "model_family", "model_variant", "layer_role", + "reference_step", "target_step", "reference_timestep", "target_timestep", + ], + ) + write_csv(output_dir / "matched_gap_summary.csv", matched_summary) + plot_overview(summary, output_dir / "cosine_overview.png") + plot_matched(rows, output_dir / "matched_gap_comparison.png") + plot_layer_curves(summary, output_dir / "layer_depth_curves.png") + plot_timestep_curves(by_step, output_dir / "timestep_curves.png") + plot_heatmaps(rows, output_dir) + (output_dir / "validation.json").write_text( + json.dumps(validation, indent=2, ensure_ascii=False) + "\n", encoding="utf-8" + ) + (output_dir / "analysis_config.json").write_text( + json.dumps( + { + "input_root": str(root), + "prompt_ids": list(range(10)), + "layer_roles": { + f"{family}/{variant}": mapping + for (family, variant), mapping in ROLE_BY_VARIANT.items() + }, + "primary_common_support": { + "minimum_target_chunk": 1, + "minimum_target_step": 1, + }, + "matched_50_step_pairs": [[0, 12], [12, 25], [25, 37]], + "bootstrap_seed": 20260826, + "bootstrap_draws": 10000, + "token_sample_count": 240, + "weights": { + "self_forcing/wan14b50": "Wan2.1-T2V-14B", + "self_forcing/dmd4": "Self-Forcing DMD checkpoint in active config", + "causal_forcing/ar50": { + "repository": "zhuhz22/Causal-Forcing", + "file": "chunkwise/ar_diffusion.pt", + "sha256": "dc9abfd679263b6a22236b9468603e7eaf3df50893f83700ff597a2ab5a6275d", + }, + "causal_forcing/dmd4": "checkpoints/chunkwise/causal_forcing.pt", + "hy_worldplay/ar50": "HY-WorldPlay/ar_model", + "hy_worldplay/ar4": "HY-WorldPlay/ar_distilled_action_model", + }, + }, + indent=2, + ensure_ascii=False, + ) + "\n", + encoding="utf-8", + ) + lines = [ + "# 50-step vs 4-step timestep/chunk cosine analysis", + "", + "Primary values use the common support (target chunk >= 1 and target step >= 1),", + "then average per prompt and report prompt-bootstrap 95% intervals.", + "Wan14B cross-chunk values are adjacent temporal slices from one full-video forward,", + "not previously computed autoregressive chunks.", + "The `final` role is the hidden state after the final transformer block.", + "Interpret 50-step vs 4-step within each family; absolute values across families are", + "confounded by architecture, feature dimension, and T2V/I2V conditioning differences.", + "", + "| family | variant | layer | comparison | cosine | 95% CI |", + "|---|---|---|---|---:|---:|", + ] + for family in FAMILIES: + for variant in VARIANT_ORDER[family]: + for role in ROLE_ORDER: + for comparison in ["within_adjacent", "cross_same"]: + row = next( + item for item in summary + if item["model_family"] == family + and item["model_variant"] == variant + and item["layer_role"] == role + and item["comparison"] == comparison + ) + lines.append( + f"| {family} | {variant} | {role} | {comparison} | " + f"{row['mean']:.4f} | [{row['ci95_low']:.4f}, {row['ci95_high']:.4f}] |" + ) + lines.extend([ + "", + "## Paired within-step minus cross-chunk difference", + "", + "Positive values mean the current chunk's previous denoising step is more similar.", + "", + "| family | variant | layer | delta cosine | 95% CI |", + "|---|---|---|---:|---:|", + ]) + for row in delta_summary: + lines.append( + f"| {row['model_family']} | {row['model_variant']} | {row['layer_role']} | " + f"{row['mean']:.4f} | [{row['ci95_low']:.4f}, {row['ci95_high']:.4f}] |" + ) + (output_dir / "REPORT.md").write_text("\n".join(lines) + "\n", encoding="utf-8") + print(f"[complete] {output_dir}") + + +if __name__ == "__main__": + main() diff --git a/scripts/summarize_trained_long_predictor_eval.py b/scripts/summarize_trained_long_predictor_eval.py new file mode 100644 index 0000000000000000000000000000000000000000..d6288ef99643be27a6235e79c18bd5ac60af8c59 --- /dev/null +++ b/scripts/summarize_trained_long_predictor_eval.py @@ -0,0 +1,126 @@ +#!/usr/bin/env python3 +"""Aggregate trained long-predictor evaluation and prepare VBench inputs.""" + +from __future__ import annotations + +import csv +import json +import math +import os +from pathlib import Path + + +ROOT = Path(__file__).resolve().parents[1] +OUTPUT = ROOT / "outputs/layer17_long_training_eval" +LENGTHS = (21, 42, 84) +PROMPTS = range(80, 100) +METHOD_FILES = { + "ffff": "ffff.mp4", + "trained_2x": "trained_2x.mp4", + "trained_4x": "trained_4x.mp4", +} +DIMENSIONS = [ + "subject_consistency", "background_consistency", "motion_smoothness", + "aesthetic_quality", "imaging_quality", +] + + +def atomic_json(path: Path, value) -> 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) + "\n") + os.replace(temporary, path) + + +def find_run(length: int, prompt: int) -> Path: + matches = list(OUTPUT.glob(f"shard_gpu*/latent_{length}/prompt_{prompt:04d}")) + if len(matches) != 1: + raise RuntimeError(f"Expected one run for {length}/{prompt}: {matches}") + return matches[0] + + +def main() -> None: + summary = [] + per_prompt = [] + for length in LENGTHS: + runs = [] + for prompt_id in PROMPTS: + run_dir = find_run(length, prompt_id) + value = json.loads((run_dir / "metrics.json").read_text()) + if value.get("status") != "complete": + raise RuntimeError(f"Incomplete: {run_dir}") + runs.append(value) + for method, filename in METHOD_FILES.items(): + video_dir = OUTPUT / "vbench_inputs" / f"latent_{length}" / method + video_dir.mkdir(parents=True, exist_ok=True) + target = video_dir / f"prompt_{prompt_id:04d}.mp4" + if target.exists() or target.is_symlink(): + target.unlink() + target.symlink_to((run_dir / filename).resolve()) + + for model in ("trained_2x", "trained_4x"): + results = [run["predictors"][model] for run in runs] + mse = [x for result in results for x in result["mse_per_frame"]] + ssim = [x for result in results for x in result["ssim_per_frame"]] + lpips = [x for result in results for x in result["lpips_per_frame"]] + summary.append({ + "latent_length": length, + "model": model, + "decoded_frames_per_video": runs[0]["decoded_frames"], + "num_prompts": len(runs), + "psnr": -10.0 * math.log10(sum(mse) / len(mse)), + "ssim": sum(ssim) / len(ssim), + "lpips": sum(lpips) / len(lpips), + "mean_ffff_generation_time_s": sum(r["ffff"]["generation_time_s"] for r in runs) / len(runs), + "mean_fppf_generation_time_s": sum(x["fppf"]["generation_time_s"] for x in results) / len(results), + }) + for run, result in zip(runs, results): + per_prompt.append({ + "latent_length": length, "model": model, + "prompt_id": run["prompt_id"], + "psnr": result["psnr"], "ssim": result["ssim"], + "lpips": result["lpips"], + }) + + for method in METHOD_FILES: + video_dir = OUTPUT / "vbench_inputs" / f"latent_{length}" / method + atomic_json(video_dir / "full_info.json", [ + { + "video_list": [f"prompt_{run['prompt_id']:04d}.mp4"], + "prompt_en": run["prompt"], "dimension": DIMENSIONS, + } + for run in runs + ]) + + for filename, rows in (("summary.csv", summary), ("per_prompt.csv", per_prompt)): + with (OUTPUT / filename).open("w", newline="", encoding="utf-8") as handle: + writer = csv.DictWriter(handle, fieldnames=list(rows[0])) + writer.writeheader(); writer.writerows(rows) + atomic_json(OUTPUT / "summary.json", summary) + + vbench_rows = [] + for length in LENGTHS: + for method in METHOD_FILES: + result_dir = OUTPUT / "vbench_results" / f"latent_{length}" / method + matches = list(result_dir.glob("*_eval_results.json")) + if not matches: + continue + if len(matches) != 1: + raise RuntimeError(f"Expected one VBench result in {result_dir}: {matches}") + path = matches[0] + result = json.loads(path.read_text()) + scores = {dimension: result[dimension][0] for dimension in DIMENSIONS} + vbench_rows.append({ + "latent_length": length, "method": method, + **scores, "vbench5_mean": sum(scores.values()) / len(scores), + }) + if vbench_rows: + with (OUTPUT / "vbench_summary.csv").open("w", newline="", encoding="utf-8") as handle: + writer = csv.DictWriter(handle, fieldnames=list(vbench_rows[0])) + writer.writeheader(); writer.writerows(vbench_rows) + atomic_json(OUTPUT / "vbench_summary.json", vbench_rows) + print(json.dumps(summary, indent=2)) + + +if __name__ == "__main__": + main() diff --git a/scripts/summarize_vbench8_extended.py b/scripts/summarize_vbench8_extended.py new file mode 100644 index 0000000000000000000000000000000000000000..51ae139cff11effb8a2e1dd1d4ecb2989d92fc9f --- /dev/null +++ b/scripts/summarize_vbench8_extended.py @@ -0,0 +1,167 @@ +#!/usr/bin/env python3 +"""Merge VBench-8 and FFFF-relative pixel metrics for all strategies.""" + +from __future__ import annotations + +import argparse +import csv +import hashlib +import json +import sys +from pathlib import Path +from typing import Any + +REPO_ROOT = Path(__file__).resolve().parents[1] +if str(REPO_ROOT) not in sys.path: + sys.path.insert(0, str(REPO_ROOT)) + +from scripts.summarize_vbench8_generation import STRATEGIES +from scripts.vbench8_protocol import DIMENSIONS, PROTOCOL_NAME + + +def sha256(path: Path) -> str: + digest = hashlib.sha256() + with path.open("rb") as handle: + for block in iter(lambda: handle.read(1024 * 1024), b""): + digest.update(block) + return digest.hexdigest() + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--output-root", type=Path, required=True) + parser.add_argument("--mapping", type=Path, required=True) + parser.add_argument("--extended-prompts", type=Path, required=True) + parser.add_argument("--vbench-info", type=Path, required=True) + parser.add_argument( + "--strategies", + nargs="+", + default=list(STRATEGIES), + help="Strategies to merge; must match the pixel summary.", + ) + return parser.parse_args() + + +def main() -> None: + args = parse_args() + root = args.output_root.resolve() + strategies = tuple(args.strategies) + if len(set(strategies)) != len(strategies): + raise ValueError("--strategies values must be unique") + if "ffff" not in strategies: + raise ValueError("--strategies must include ffff") + pixel_rows = { + row["strategy"]: row + for row in json.loads((root / "summaries/pixel_metrics.json").read_text()) + } + rows: list[dict[str, Any]] = [] + for strategy in strategies: + score_path = root / "vbench/scores" / f"{strategy}.json" + if not score_path.is_file(): + raise FileNotFoundError(score_path) + score = json.loads(score_path.read_text(encoding="utf-8")) + if strategy not in pixel_rows: + raise KeyError(f"Missing pixel summary for {strategy}") + pixel = pixel_rows[strategy] + row: dict[str, Any] = { + "strategy": strategy, + "num_videos": score["num_videos"], + "psnr": pixel["psnr"], + "ssim": pixel["ssim"], + "lpips": pixel["lpips"], + "pixel_metric_scope": pixel["pixel_metric_scope"], + "pixel_metric_input": pixel.get("pixel_metric_input"), + "accepted_predictor_calls": pixel["accepted_predictor_calls"], + "full_calls": pixel["full_calls"], + "predictor_calls": pixel["predictor_calls"], + "reuse_calls": pixel.get("reuse_calls", 0.0), + "latency_ms": pixel["mean_policy_latency_ms"], + "speedup_percent_vs_ffff": pixel[ + "policy_latency_speedup_percent_vs_ffff" + ], + "denoise_dit_latency_ms": pixel["mean_denoise_dit_latency_ms"], + "denoise_dit_speedup_percent_vs_ffff": pixel[ + "denoise_dit_speedup_percent_vs_ffff" + ], + "confidence_head_time_ms": pixel["mean_confidence_head_time_ms"], + "excluded_context_dit_time_ms": pixel[ + "mean_excluded_context_dit_time_ms" + ], + "quality_score": score["quality_score"], + "semantic_score": score["semantic_score"], + "selected_vbench_score": score["selected_vbench_score"], + "selected_vbench_percent": score["selected_vbench_percent"], + } + for dimension in DIMENSIONS: + row[f"raw_{dimension}"] = score["raw_scores"][dimension] + row[f"normalized_{dimension}"] = score["normalized_scores"][dimension] + rows.append(row) + + summary_dir = root / "summaries" + fields = list(rows[0].keys()) + with (summary_dir / "final_summary.csv").open("w", encoding="utf-8", newline="") as handle: + writer = csv.DictWriter(handle, fieldnames=fields) + writer.writeheader() + writer.writerows(rows) + final = { + "protocol": PROTOCOL_NAME, + "protocol_version": "1.0", + "benchmark": "standard VBench", + "vbench_long": False, + "vbench_version": "0.1.5", + "prompt_source": str(args.extended_prompts.resolve()), + "prompt_source_sha256": sha256(args.extended_prompts.resolve()), + "mapping": str(args.mapping.resolve()), + "mapping_sha256": sha256(args.mapping.resolve()), + "vbench_info": str(args.vbench_info.resolve()), + "vbench_info_sha256": sha256(args.vbench_info.resolve()), + "num_unique_prompts": 251, + "videos_per_prompt": 1, + "dimensions": list(DIMENSIONS), + "strategies": rows, + "normalization": { + "source": f"{PROTOCOL_NAME} protocol", + "formula": "(raw-min)/(max-min), no clipping", + }, + "aggregation": { + "quality_semantic_weight": "4:1", + "formula": "(4*quality_score + semantic_score)/5", + }, + "pixel_metrics": { + "reference": "same extended prompt and seed; FFFF strategy", + "scope": "all decoded frames 0:81 (81 frames)", + "psnr_aggregation": "-10*log10(mean pixel MSE)", + "ssim_lpips_aggregation": "arithmetic mean over frames and prompts", + }, + "latency": { + "primary": "policy_latency_ms", + "formula": ( + "full_dit_time_ms + predictor_time_ms + " + "confidence_head_time_ms" + ), + "excluded": [ + "context_dit_time_ms (KV-cache update DiT)", + "text encoding", + "VAE decoding", + "scheduler/noise overhead", + "video I/O", + "metric evaluation", + ], + "speedup_formula": ( + "100 * (1 - mean(strategy_policy_latency) / " + "mean(matched_ffff_policy_latency))" + ), + }, + "scene_note": ( + "In vbench==0.1.5 scene uses official auxiliary scene keywords; " + "overall_consistency uses the extended prompt_en text." + ), + } + (summary_dir / "final_summary.json").write_text( + json.dumps(final, ensure_ascii=False, indent=2) + "\n", encoding="utf-8" + ) + print(f"[complete] final summary={summary_dir / 'final_summary.csv'}") + + +if __name__ == "__main__": + main() diff --git a/scripts/summarize_vbench8_generation.py b/scripts/summarize_vbench8_generation.py new file mode 100644 index 0000000000000000000000000000000000000000..f5cc039d77b60548d4eb81f72db39ef4deb95966 --- /dev/null +++ b/scripts/summarize_vbench8_generation.py @@ -0,0 +1,229 @@ +#!/usr/bin/env python3 +"""Aggregate generation diagnostics and FFFF-relative pixel metrics.""" + +from __future__ import annotations + +import argparse +import csv +import json +import math +import statistics +import sys +from pathlib import Path +from typing import Any + +REPO_ROOT = Path(__file__).resolve().parents[1] +if str(REPO_ROOT) not in sys.path: + sys.path.insert(0, str(REPO_ROOT)) + +from scripts.vbench8_protocol import SUITE_COUNTS + + +STRATEGIES = ( + "ffff", + "fppf", + "step12_k06", + "step12_k08", + "step12_k10", + "step123_k06", + "step123_k09", + "step123_k12", + "step123_k15", +) + + +def read_records(root: Path, strategy: str) -> list[dict[str, Any]]: + directory = root / "generation_metrics/per_prompt" / strategy + paths = sorted(directory.glob("global_*.json")) + if len(paths) != 251: + raise ValueError( + f"Expected 251 records for {strategy}, found {len(paths)} in {directory}" + ) + records = [json.loads(path.read_text(encoding="utf-8")) for path in paths] + globals_seen = {int(record["global_index"]) for record in records} + if len(globals_seen) != 251: + raise ValueError(f"Duplicate global indices for {strategy}") + suite_counts = {suite: 0 for suite in SUITE_COUNTS} + for record in records: + suite_counts[str(record["prompt_suite"])] += 1 + if suite_counts != SUITE_COUNTS: + raise ValueError(f"Unexpected suite counts for {strategy}: {suite_counts}") + return sorted(records, key=lambda record: int(record["global_index"])) + + +def mean(values: list[float]) -> float: + return sum(values) / len(values) + + +def std(values: list[float]) -> float: + return statistics.pstdev(values) if len(values) > 1 else 0.0 + + +def denoise_dit_latency_ms(generation: dict[str, Any]) -> float: + """Timed denoise model path, excluding the context/KV-cache DiT pass.""" + return float(generation["full_dit_time_ms"]) + float( + generation["predictor_time_ms"] + ) + + +def policy_latency_ms(generation: dict[str, Any]) -> float: + """Primary latency: denoise model path plus Confidence Head overhead.""" + return denoise_dit_latency_ms(generation) + float( + generation["confidence_head_time_ms"] + ) + + +def speedup_percent(times: list[float], reference_times: list[float]) -> float: + """Ratio-of-means speedup for matched prompt/seed measurements.""" + if len(times) != len(reference_times) or not times: + raise ValueError("Latency and reference latency lists must be non-empty and matched") + return 100.0 * (1.0 - mean(times) / mean(reference_times)) + + +def aggregate_strategy( + strategy: str, + records: list[dict[str, Any]], + reference_by_global: dict[int, dict[str, Any]], +) -> dict[str, Any]: + generations = [record["generation"] for record in records] + metrics = [record["pixel_metrics_vs_ffff"] for record in records] + denoise_dit_times = [denoise_dit_latency_ms(item) for item in generations] + policy_times = [policy_latency_ms(item) for item in generations] + reference_denoise_dit_times = [ + denoise_dit_latency_ms( + reference_by_global[int(record["global_index"])]["generation"] + ) + for record in records + ] + reference_policy_times = [ + policy_latency_ms( + reference_by_global[int(record["global_index"])]["generation"] + ) + for record in records + ] + if strategy == "ffff": + policy_speedup = 0.0 + denoise_dit_speedup = 0.0 + else: + policy_speedup = speedup_percent(policy_times, reference_policy_times) + denoise_dit_speedup = speedup_percent( + denoise_dit_times, reference_denoise_dit_times + ) + pixel_mse = mean([float(item["pixel_mse"]) for item in metrics]) + psnr_values = [float(item["psnr"]) for item in metrics] + ssim_values = [float(item["ssim"]) for item in metrics] + lpips_values = [float(item["lpips"]) for item in metrics] + accepted = [float(item["accepted_predictor_calls"]) for item in generations] + full_calls = [float(item["full_calls"]) for item in generations] + predictor_calls = [float(item["predictor_calls"]) for item in generations] + reuse_calls = [float(item.get("reuse_calls", 0.0)) for item in generations] + pixel_inputs = { + str( + record.get( + "pixel_metric_input", + "pre-MP4 uint8 RGB tensors, all 81 frames on both sides", + ) + ) + for record in records + } + if len(pixel_inputs) != 1: + raise ValueError(f"Inconsistent pixel metric inputs for {strategy}: {pixel_inputs}") + aggregate = { + "strategy": strategy, + "num_prompts": len(records), + "num_videos": len(records), + "accepted_predictor_calls": mean(accepted), + "full_calls": mean(full_calls), + "predictor_calls": mean(predictor_calls), + "reuse_calls": mean(reuse_calls), + "mean_generation_time_s": mean( + [float(item["generation_time_s"]) for item in generations] + ), + "mean_policy_latency_ms": mean(policy_times), + "policy_latency_speedup_percent_vs_ffff": policy_speedup, + "mean_denoise_dit_latency_ms": mean(denoise_dit_times), + "denoise_dit_speedup_percent_vs_ffff": denoise_dit_speedup, + "mean_full_dit_time_ms": mean( + [float(item["full_dit_time_ms"]) for item in generations] + ), + "mean_predictor_time_ms": mean( + [float(item["predictor_time_ms"]) for item in generations] + ), + "mean_confidence_head_time_ms": mean( + [float(item["confidence_head_time_ms"]) for item in generations] + ), + "mean_excluded_context_dit_time_ms": mean( + [float(item["context_dit_time_ms"]) for item in generations] + ), + "latency_definition": ( + "full_dit_time_ms + predictor_time_ms + confidence_head_time_ms; " + "context_dit_time_ms excluded" + ), + "pixel_mse": pixel_mse, + "psnr": -10.0 * math.log10(max(pixel_mse, 1e-12)), + "psnr_prompt_mean": mean(psnr_values), + "psnr_prompt_std": std(psnr_values), + "ssim": mean(ssim_values), + "ssim_prompt_std": std(ssim_values), + "lpips": mean(lpips_values), + "lpips_prompt_std": std(lpips_values), + "pixel_metric_scope": "all decoded frames 0:81 (81 frames)", + "pixel_metric_input": next(iter(pixel_inputs)), + "reference": "same extended prompt and seed; FFFF strategy", + } + return aggregate + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--output-root", type=Path, required=True) + parser.add_argument( + "--strategies", + nargs="+", + default=list(STRATEGIES), + help="Strategies to aggregate. FFFF must be included as the reference.", + ) + return parser.parse_args() + + +def main() -> None: + args = parse_args() + root = args.output_root.resolve() + strategies = tuple(args.strategies) + if len(set(strategies)) != len(strategies): + raise ValueError("--strategies values must be unique") + if "ffff" not in strategies: + raise ValueError("--strategies must include ffff") + records_by_strategy = { + strategy: read_records(root, strategy) for strategy in strategies + } + reference_by_global = { + int(record["global_index"]): record for record in records_by_strategy["ffff"] + } + summary = [ + aggregate_strategy(strategy, records_by_strategy[strategy], reference_by_global) + for strategy in strategies + ] + summary_dir = root / "summaries" + summary_dir.mkdir(parents=True, exist_ok=True) + fields = list(summary[0].keys()) + with (summary_dir / "pixel_metrics.csv").open("w", encoding="utf-8", newline="") as handle: + writer = csv.DictWriter(handle, fieldnames=fields) + writer.writeheader() + writer.writerows(summary) + (summary_dir / "pixel_metrics.json").write_text( + json.dumps(summary, ensure_ascii=False, indent=2) + "\n", encoding="utf-8" + ) + print(f"[complete] pixel summary={summary_dir / 'pixel_metrics.csv'}") + for row in summary: + print( + f"{row['strategy']}: policy_speedup=" + f"{row['policy_latency_speedup_percent_vs_ffff']:.3f}% " + f"psnr={row['psnr']:.4f} " + f"ssim={row['ssim']:.6f} " + f"lpips={row['lpips']:.6f}" + ) + + +if __name__ == "__main__": + main() diff --git a/scripts/train_conditional_mlp_offline.py b/scripts/train_conditional_mlp_offline.py new file mode 100644 index 0000000000000000000000000000000000000000..9e576a3de642737b1c872d60137ed387e962dd4a --- /dev/null +++ b/scripts/train_conditional_mlp_offline.py @@ -0,0 +1,305 @@ +#!/usr/bin/env python3 +"""Train a capacity-matched token-wise nonlinear conditional probe offline.""" + +from __future__ import annotations + +import argparse +import csv +import json +import os +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) + + +_preparse_gpu() + +import numpy as np +import torch +from torch import nn + +from run_conditional_probe_offline import FAMILIES, ROLES, load_family, samples + + +VARIANTS = ( + "step_only", + "chunk_only", + "both_correct", + "step_duplicate", + "wrong_step", + "other_video", +) + + +def write_csv(path: Path, rows: list[dict[str, Any]]) -> None: + if not rows: + return + path.parent.mkdir(parents=True, exist_ok=True) + 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, extrasaction="ignore") + writer.writeheader() + writer.writerows(rows) + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser() + parser.add_argument("--gpu", default="0") + parser.add_argument("--dataset_root", type=Path, required=True) + parser.add_argument("--output_dir", type=Path, required=True) + parser.add_argument("--family", choices=FAMILIES, required=True) + parser.add_argument("--num_prompts", type=int, default=10) + parser.add_argument("--chunks", type=int, default=4) + parser.add_argument("--steps", type=int, default=4) + parser.add_argument("--hidden_dim", type=int, default=32) + parser.add_argument("--train_steps", type=int, default=160) + parser.add_argument("--batch_size", type=int, default=1024) + parser.add_argument("--lr", type=float, default=1e-3) + parser.add_argument("--seeds", default="0,1,2") + parser.add_argument( + "--chunk_pairing", + choices=("matched_slot", "boundary_to_all"), + default="matched_slot", + ) + return parser.parse_args() + + +class ConditionalMLP(nn.Module): + def __init__(self, dim: int, hidden_dim: int): + super().__init__() + self.step_norm = nn.LayerNorm(dim) + self.aux_norm = nn.LayerNorm(dim) + self.step_proj = nn.Linear(dim, hidden_dim) + self.aux_proj = nn.Linear(dim, hidden_dim) + self.trunk = nn.Sequential( + nn.Linear(2 * hidden_dim + 3, hidden_dim), + nn.SiLU(), + nn.Linear(hidden_dim, hidden_dim), + nn.SiLU(), + ) + self.out = nn.Linear(hidden_dim, dim) + + def forward(self, step: torch.Tensor, aux: torch.Tensor, t: torch.Tensor) -> torch.Tensor: + if t.ndim == 1: + t_onehot = torch.nn.functional.one_hot(t.long(), num_classes=3).float() + else: + t_onehot = t.float() + fused = torch.cat( + [self.step_proj(self.step_norm(step)), self.aux_proj(self.aux_norm(aux)), t_onehot], + dim=-1, + ) + return self.out(self.trunk(fused)) + + +def flatten_training(prepared, prompt_ids, role, chunks, steps, aux_kind): + step_rows, aux_rows, target_rows, t_rows = [], [], [], [] + for position, prompt_id in enumerate(prompt_ids): + for step in range(1, steps): + item = prepared[prompt_id][step] + step_rows.append(item["within"]) + target_rows.append(item["target"]) + t_rows.append(torch.full((item["target"].shape[0],), step - 1, dtype=torch.long)) + if aux_kind == "correct": + aux_rows.append(item["cross"]) + elif aux_kind == "wrong": + aux_rows.append(item["wrong"]) + elif aux_kind == "distant": + aux_rows.append(item["distant"]) + elif aux_kind == "duplicate": + aux_rows.append(item["within"]) + elif aux_kind == "zero": + aux_rows.append(torch.zeros_like(item["within"])) + elif aux_kind == "other": + donor = prompt_ids[(position + 1) % len(prompt_ids)] + aux_rows.append(prepared[donor][step]["cross"]) + else: + raise ValueError(aux_kind) + return ( + torch.cat(step_rows, dim=0), + torch.cat(aux_rows, dim=0), + torch.cat(target_rows, dim=0), + torch.cat(t_rows, dim=0), + ) + + +def aux_for_variant(variant: str) -> str: + return { + "step_only": "zero", + "chunk_only": "correct", + "both_correct": "correct", + "step_duplicate": "duplicate", + "wrong_step": "wrong", + "other_video": "other", + }[variant] + + +def target_metrics(pred: torch.Tensor, target: torch.Tensor) -> dict[str, float]: + pred, target = pred.float(), target.float() + error = pred - target + mse = error.square().mean() + variance = (target - target.mean()).square().mean().clamp_min(1e-12) + nmse = mse / variance + cosine = torch.nn.functional.cosine_similarity( + pred.reshape(1, -1), target.reshape(1, -1), dim=1, eps=1e-8 + )[0] + return { + "mse": float(mse), + "nMSE": float(nmse), + "nRMSE": float(torch.sqrt(nmse)), + "r2": float(1.0 - nmse), + "cosine": float(cosine), + } + + +def train_one( + model: nn.Module, + step: torch.Tensor, + aux: torch.Tensor, + target: torch.Tensor, + t: torch.Tensor, + steps: int, + train_steps: int, + batch_size: int, + lr: float, +) -> tuple[torch.Tensor, torch.Tensor]: + device = next(model.parameters()).device + target_mean = target.mean(dim=0, keepdim=True) + target_std = target.std(dim=0, keepdim=True).clamp_min(1e-3) + target_norm = (target - target_mean) / target_std + optimizer = torch.optim.AdamW(model.parameters(), lr=lr, weight_decay=1e-4) + model.train() + count = target.shape[0] + for _ in range(train_steps): + indices = torch.randint(0, count, (min(batch_size, count),), device=device) + prediction = model(step[indices], aux[indices], t[indices]) + loss = torch.nn.functional.mse_loss(prediction, target_norm[indices]) + optimizer.zero_grad(set_to_none=True) + loss.backward() + optimizer.step() + return target_mean, target_std + + +def main() -> None: + args = parse_args() + torch.set_float32_matmul_precision("high") + torch.set_num_threads(4) + device = torch.device("cuda" if torch.cuda.is_available() else "cpu") + runs = load_family(args.dataset_root.resolve(), args.family, args.num_prompts) + seeds = [int(value) for value in args.seeds.split(",") if value.strip()] + rows: list[dict[str, Any]] = [] + for role in ROLES: + prepared = [ + { + step: samples(run, role, args.chunks, step, args.chunk_pairing) + for step in range(1, args.steps) + } + for run in runs + ] + # A single layer-specific model sees all three target timesteps; results + # are still reported separately by target timestep. + for held_out in range(args.num_prompts): + donor = (held_out + 1) % args.num_prompts + train_ids = [i for i in range(args.num_prompts) if i not in {held_out, donor}] + for seed in seeds: + torch.manual_seed(10000 * seed + 100 * held_out + ROLES.index(role)) + test_by_step = {step: prepared[held_out][step] for step in range(1, args.steps)} + for variant in VARIANTS: + train_step, train_aux, train_target, train_t = flatten_training( + prepared, + train_ids, + role, + args.chunks, + args.steps, + aux_for_variant(variant), + ) + if variant == "chunk_only": + train_step = torch.zeros_like(train_step) + dim = int(train_target.shape[1]) + model = ConditionalMLP(dim, args.hidden_dim).to(device) + train_step = train_step.to(device) + train_aux = train_aux.to(device) + train_target = train_target.to(device) + train_t = train_t.to(device) + target_mean, target_std = train_one( + model, + train_step, + train_aux, + train_target, + train_t, + args.steps, + args.train_steps, + args.batch_size, + args.lr, + ) + model.eval() + with torch.no_grad(): + for step in range(1, args.steps): + item = test_by_step[step] + test_step = item["within"].to(device) + if variant == "chunk_only": + test_step = torch.zeros_like(test_step) + if aux_for_variant(variant) == "correct": + test_aux = item["cross"].to(device) + elif aux_for_variant(variant) == "wrong": + test_aux = item["wrong"].to(device) + elif aux_for_variant(variant) == "distant": + test_aux = item["distant"].to(device) + elif aux_for_variant(variant) == "duplicate": + test_aux = item["within"].to(device) + elif aux_for_variant(variant) == "zero": + test_aux = torch.zeros_like(test_step) + else: + test_aux = prepared[donor][step]["cross"].to(device) + test_t = torch.full( + (test_step.shape[0],), step - 1, dtype=torch.long, device=device + ) + prediction = model(test_step, test_aux, test_t) + prediction = prediction * target_std.to(device) + target_mean.to(device) + values = target_metrics(prediction, item["target"].to(device)) + rows.append({ + "model_family": args.family, + "layer_role": role, + "layer_index": ( + {"early": 13, "middle": 26, "late": 40, "final": 53}[role] + if args.family == "hy_worldplay" + else {"early": 7, "middle": 14, "late": 22, "final": 29}[role] + ), + "target_step": step, + "held_out_prompt": held_out, + "other_video_prompt": donor, + "seed": seed, + "train_prompts": len(train_ids), + "test_tokens": int(item["target"].shape[0]), + "probe": variant, + **values, + }) + del model, train_step, train_aux, train_target, train_t + torch.cuda.empty_cache() + print( + f"[progress] {args.family} {role} heldout={held_out}", flush=True + ) + + args.output_dir.mkdir(parents=True, exist_ok=True) + write_csv(args.output_dir / f"nonlinear_probe_{args.family}_folds.csv", rows) + config = vars(args).copy() + config["device"] = str(device) + config["rows"] = len(rows) + (args.output_dir / f"nonlinear_probe_{args.family}_config.json").write_text( + json.dumps(config, indent=2, default=str) + "\n", encoding="utf-8" + ) + print(f"[complete] {args.family} rows={len(rows)}", flush=True) + + +if __name__ == "__main__": + main() diff --git a/scripts/train_confidence_token_lazy_ddp.py b/scripts/train_confidence_token_lazy_ddp.py new file mode 100644 index 0000000000000000000000000000000000000000..ece0b254079dad28715c767d70cc3c0711afd4f7 --- /dev/null +++ b/scripts/train_confidence_token_lazy_ddp.py @@ -0,0 +1,816 @@ +#!/usr/bin/env python3 +"""Train the Layer-17 Confidence-token head with frozen Predictor features.""" + +from __future__ import annotations + +import argparse +import csv +import json +import math +import os +import random +import sys +import time +from pathlib import Path +from typing import Any + +import torch +import torch.distributed as dist +import torch.nn.functional as F +from safetensors import safe_open +from safetensors.torch import load_file, save_file +from torch.nn.parallel import DistributedDataParallel as DDP +from torch.optim import AdamW + +ROOT = Path(__file__).resolve().parents[1] +if str(ROOT) not in sys.path: + sys.path.insert(0, str(ROOT)) + +from predictor_training.confidence import ConfidenceTokenHead +from predictor_training.lazy_offline_data import ( + LazyLayer17Dataset, + collate_lazy_samples, +) +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 frozen_inputs, load_teacher +from utils.misc import set_seed +from wan.modules.causal_model import causal_rope_apply + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--dataset_root", type=Path, required=True) + parser.add_argument("--output_dir", type=Path, required=True) + parser.add_argument("--predictor_weights", type=Path, required=True) + parser.add_argument( + "--predictor_input_variant", + choices=("auto", "self_forcing", "disca", "atc"), + default="auto", + help=( + "Expected Predictor input variant. Legacy concat checkpoints have no " + "metadata and therefore require --predictor_input_variant self_forcing." + ), + ) + parser.add_argument( + "--checkpoint_path", + type=Path, + default=Path("checkpoints/self_forcing_dmd.pt"), + ) + parser.add_argument( + "--config_path", type=Path, default=Path("configs/self_forcing_sid.yaml") + ) + parser.add_argument("--epochs", type=int, default=20) + parser.add_argument("--per_device_batch_size", type=int, default=64) + parser.add_argument("--learning_rate", type=float, default=3e-4) + parser.add_argument( + "--lr_schedule", + choices=("constant", "warmup_cosine"), + default="constant", + ) + parser.add_argument("--warmup_steps", type=int, default=0) + parser.add_argument("--warmup_start_lr", type=float, default=1e-5) + parser.add_argument("--min_learning_rate", type=float, default=1e-6) + parser.add_argument("--weight_decay", type=float, default=0.01) + parser.add_argument("--dropout", type=float, default=0.1) + parser.add_argument("--grad_clip", type=float, default=1.0) + parser.add_argument("--history_projection_batch_size", type=int, default=8) + parser.add_argument("--patience", type=int, default=5) + parser.add_argument("--min_delta", type=float, default=1e-4) + parser.add_argument("--seed", type=int, default=0) + args = parser.parse_args() + if args.epochs < 1 or args.per_device_batch_size < 1: + parser.error("epochs and batch size must be positive") + if args.learning_rate <= 0 or args.grad_clip <= 0: + parser.error("learning rate and grad clip must be positive") + if args.warmup_steps < 0: + parser.error("warmup steps must be non-negative") + if args.lr_schedule == "warmup_cosine" and args.warmup_steps < 1: + parser.error("warmup_cosine requires at least one warmup step") + if not 0 < args.warmup_start_lr <= args.learning_rate: + parser.error("warmup start LR must be in (0, learning_rate]") + if not 0 < args.min_learning_rate <= args.learning_rate: + parser.error("minimum LR must be in (0, learning_rate]") + if args.history_projection_batch_size < 1: + parser.error("history projection batch size must be positive") + return args + + +def resolve(path: Path) -> Path: + path = path.expanduser() + return path.resolve() if path.is_absolute() else (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 append_jsonl(path: Path, value: dict[str, Any]) -> None: + with path.open("a", encoding="utf-8") as handle: + handle.write(json.dumps(value, sort_keys=True) + "\n") + + +def save_weights( + model: torch.nn.Module, + path: Path, + *, + metadata: dict[str, str], +) -> None: + tensors = { + key: value.detach().cpu().contiguous() + for key, value in model.state_dict().items() + } + temporary = path.with_suffix(path.suffix + ".tmp") + save_file(tensors, temporary, metadata=metadata) + os.replace(temporary, path) + + +def write_predictions(path: Path, rows: list[dict[str, float | int]]) -> None: + temporary = path.with_suffix(path.suffix + ".tmp") + with temporary.open("w", newline="", encoding="utf-8") as handle: + writer = csv.DictWriter(handle, fieldnames=list(rows[0])) + writer.writeheader() + writer.writerows(rows) + os.replace(temporary, path) + + +def read_predictor_config( + path: Path, expected_input_variant: str = "auto" +) -> dict[str, Any]: + with safe_open(path, framework="pt", device="cpu") as handle: + metadata = handle.metadata() or {} + raw = metadata.get("predictor_config") + if raw is None: + if expected_input_variant != "self_forcing": + raise ValueError( + f"Missing predictor_config metadata: {path}; legacy concat " + "checkpoints require --predictor_input_variant self_forcing" + ) + return { + "source_layer": 17, + "input_variant": "self_forcing", + "gate_mode": "baseline", + "metadata_source": "explicit_legacy_concat_override", + } + config = json.loads(raw) + actual = str(config.get("input_variant", "self_forcing")) + if expected_input_variant != "auto" and actual != expected_input_variant: + raise ValueError( + f"Predictor input variant mismatch: expected={expected_input_variant} " + f"actual={actual}" + ) + return config + + +def load_predictor( + teacher: torch.nn.Module, + weights: Path, + device: torch.device, + expected_input_variant: str = "auto", +) -> tuple[SingleBlockPredictor, dict[str, Any]]: + config = read_predictor_config(weights, expected_input_variant) + source_layer = int(config.get("source_layer", 17)) + input_variant = config.get("input_variant", "self_forcing") + predictor = SingleBlockPredictor( + block=initialize_predictor_block( + teacher.blocks[source_layer], "teacher_full" + ), + dim=teacher.dim, + gradient_checkpointing=False, + input_variant=input_variant, + atc_previous_scope=config.get("atc_previous_scope", "chunk"), + atc_freq_dim=int(config.get("atc_freq_dim", 256)), + atc_mlp_hidden_dim=int(config.get("atc_mlp_hidden_dim", 3072)), + atc_gate_hidden_dim=int(config.get("atc_gate_hidden_dim", 512)), + atc_transport_residual_scale=float( + config.get("atc_transport_residual_scale", 0.1) + ), + atc_gate_initial_probability=float( + config.get("atc_gate_initial_probability", 0.3) + ), + atc_collect_diagnostics=False, + ) + predictor.load_state_dict(load_file(str(weights), device="cpu"), strict=True) + predictor.to(device=device, dtype=torch.bfloat16) + predictor.eval().requires_grad_(False) + return predictor, config + + +@torch.inference_mode() +def project_history_chunked( + prefeature: torch.Tensor, + chunk: int, + teacher: torch.nn.Module, + device: torch.device, + micro_batch_size: int, +) -> tuple[torch.Tensor, torch.Tensor]: + """Project history K/V without a full-batch float64 RoPE temporary.""" + batch, sequence, _ = prefeature.shape + block = teacher.blocks[17] + heads = block.num_heads + head_dim = block.dim // heads + key_output = torch.empty( + batch, + sequence, + heads, + head_dim, + dtype=torch.bfloat16, + device=device, + ) + value_output = torch.empty_like(key_output) + for start in range(0, batch, micro_batch_size): + end = min(start + micro_batch_size, batch) + value_in = prefeature[start:end].to( + device=device, dtype=torch.bfloat16 + ) + with torch.autocast(device_type="cuda", dtype=torch.bfloat16): + key = block.self_attn.norm_k(block.self_attn.k(value_in)).view( + end - start, sequence, heads, head_dim + ) + value = block.self_attn.v(value_in).view( + end - start, sequence, heads, head_dim + ) + grid_sizes = torch.tensor( + [[chunk * 3, 30, 52]] * (end - start), + dtype=torch.long, + device="cpu", + ) + key = causal_rope_apply( + key, grid_sizes, teacher.freqs, start_frame=0 + ) + key_output[start:end].copy_(key) + value_output[start:end].copy_(value) + del value_in, key, value + return key_output, value_output + + +def move_batch( + cpu_batch: dict[str, Any], + teacher: torch.nn.Module, + device: torch.device, + history_projection_batch_size: int, +) -> dict[str, Any]: + history_k, history_v = project_history_chunked( + cpu_batch.pop("history_prefeature"), + cpu_batch["chunk"], + teacher, + device, + history_projection_batch_size, + ) + required = ( + "noisy_latent", + "timestep", + "anchor_timestep", + "anchor_distance", + "anchor_hidden", + "previous_hidden", + "target_hidden", + "cross_k", + "cross_v", + ) + batch = { + key: cpu_batch[key].to( + device=device, + dtype=( + torch.bfloat16 + if cpu_batch[key].is_floating_point() and key != "timestep" + else cpu_batch[key].dtype + ), + ) + for key in required + } + batch.update( + prompt_ids=cpu_batch["prompt_ids"], + chunk=cpu_batch["chunk"], + target_step=cpu_batch["target_step"], + history_k=history_k, + history_v=history_v, + ) + return batch + + +@torch.no_grad() +def predictor_features( + predictor: SingleBlockPredictor, + batch: dict[str, Any], + teacher: torch.nn.Module, + device: torch.device, +) -> tuple[torch.Tensor, torch.Tensor]: + frozen = frozen_inputs(batch, teacher, device) + with torch.autocast(device_type="cuda", dtype=torch.bfloat16): + output = predictor( + current_tokens=frozen["current_tokens"], + anchor_hidden=batch["anchor_hidden"], + previous_hidden=batch["previous_hidden"], + timestep_modulation=frozen["timestep_modulation"], + grid_sizes=frozen["grid_sizes"], + freqs=frozen["freqs"], + history_k=batch["history_k"], + history_v=batch["history_v"], + cross_k=batch["cross_k"], + cross_v=batch["cross_v"], + current_start=batch["chunk"] * TOKENS_PER_CHUNK, + return_features=True, + condition_tokens=frozen["condition_tokens"], + anchor_distance=batch["anchor_distance"], + ) + if not isinstance(output, tuple): + raise RuntimeError("Predictor did not return confidence features") + return output + + +def hidden_nrmse(predicted: torch.Tensor, target: torch.Tensor) -> torch.Tensor: + error_energy = (predicted.float() - target.float()).square().sum(dim=(1, 2)) + target_energy = target.float().square().sum(dim=(1, 2)) + return torch.sqrt(error_energy / target_energy.clamp_min(1e-8)) + + +def ranks(values: list[float]) -> list[float]: + order = sorted(range(len(values)), key=values.__getitem__) + result = [0.0] * len(values) + offset = 0 + while offset < len(order): + end = offset + 1 + while end < len(order) and values[order[end]] == values[order[offset]]: + end += 1 + rank = 0.5 * (offset + end - 1) + for index in order[offset:end]: + result[index] = rank + offset = end + return result + + +def correlation(left: list[float], right: list[float]) -> float: + left_mean = sum(left) / len(left) + right_mean = sum(right) / len(right) + numerator = sum((a - left_mean) * (b - right_mean) for a, b in zip(left, right)) + denominator = math.sqrt( + sum((a - left_mean) ** 2 for a in left) + * sum((b - right_mean) ** 2 for b in right) + ) + return numerator / denominator if denominator > 0 else float("nan") + + +def learning_rate_at_step( + *, + step: int, + total_steps: int, + schedule: str, + peak_lr: float, + warmup_steps: int, + warmup_start_lr: float, + min_lr: float, +) -> float: + """Learning rate used for a zero-based optimizer step.""" + if step < 0 or total_steps < 1 or step >= total_steps: + raise ValueError("step must be in [0, total_steps)") + if schedule == "constant": + return float(peak_lr) + if schedule != "warmup_cosine": + raise ValueError(f"Unsupported LR schedule: {schedule}") + if step < warmup_steps: + if warmup_steps == 1: + return float(peak_lr) + fraction = step / (warmup_steps - 1) + return float(warmup_start_lr + fraction * (peak_lr - warmup_start_lr)) + decay_steps = max(1, total_steps - warmup_steps) + progress = (step - warmup_steps + 1) / decay_steps + progress = min(max(progress, 0.0), 1.0) + cosine = 0.5 * (1.0 + math.cos(math.pi * progress)) + return float(min_lr + (peak_lr - min_lr) * cosine) + + +def validation_metrics(rows: list[dict[str, float | int]]) -> dict[str, float]: + target_log = [float(row["target_log_error"]) for row in rows] + predicted_log = [float(row["predicted_log_error"]) for row in rows] + target = [float(row["target_hidden_nrmse"]) for row in rows] + predicted = [float(row["predicted_hidden_nrmse"]) for row in rows] + huber = [ + float(F.smooth_l1_loss(torch.tensor(a), torch.tensor(b))) + for a, b in zip(predicted_log, target_log) + ] + return { + "huber_log": sum(huber) / len(huber), + "mae_log": sum(abs(a - b) for a, b in zip(predicted_log, target_log)) + / len(rows), + "mae_nrmse": sum(abs(a - b) for a, b in zip(predicted, target)) + / len(rows), + "pearson_nrmse": correlation(predicted, target), + "spearman_nrmse": correlation(ranks(predicted), ranks(target)), + "num_samples": float(len(rows)), + } + + +def epoch_batches( + *, + epoch: int, + rank: int, + world_size: int, + per_device_batch_size: int, + seed: int, +) -> list[tuple[int, int, list[int]]]: + rng = random.Random(seed + epoch) + groups = [(chunk, step) for chunk in range(1, 7) for step in range(1, 4)] + rng.shuffle(groups) + result = [] + global_batch_size = world_size * per_device_batch_size + for chunk, target_step in groups: + prompt_ids = list(range(900)) + rng.shuffle(prompt_ids) + for start in range(0, len(prompt_ids), global_batch_size): + selected = prompt_ids[start : start + global_batch_size] + if len(selected) % world_size: + raise RuntimeError("Training split must divide evenly across ranks") + local_count = len(selected) // world_size + offset = rank * local_count + local_ids = selected[offset : offset + local_count] + result.append((chunk, target_step, local_ids)) + return result + + +def run_validation( + *, + ddp: DDP, + predictor: SingleBlockPredictor, + dataset: LazyLayer17Dataset, + teacher: torch.nn.Module, + device: torch.device, + rank: int, + world_size: int, + history_projection_batch_size: int, +) -> tuple[dict[str, float] | None, list[dict[str, float | int]] | None]: + ddp.eval() + validation_ids = list(range(900 + rank, 1000, world_size)) + local_rows: list[dict[str, float | int]] = [] + for chunk in range(1, 7): + for target_step in range(1, 4): + cpu_batch = collate_lazy_samples( + [dataset[(prompt_id, chunk, target_step)] for prompt_id in validation_ids] + ) + batch = move_batch( + cpu_batch, + teacher, + device, + history_projection_batch_size, + ) + del cpu_batch + pred_hidden, transformed = predictor_features( + predictor, batch, teacher, device + ) + target_error = hidden_nrmse(pred_hidden, batch["target_hidden"]) + chunk_position = torch.full( + (len(validation_ids),), + (chunk - 1) / 5.0, + device=device, + ) + step_id = torch.full( + (len(validation_ids),), + target_step, + dtype=torch.long, + device=device, + ) + with torch.no_grad(), torch.autocast( + device_type="cuda", dtype=torch.bfloat16 + ): + predicted_log = ddp( + transformed_hidden=transformed, + pred_hidden=pred_hidden, + anchor_hidden=batch["anchor_hidden"], + chunk_position=chunk_position, + step_id=step_id, + ) + target_log = torch.log(target_error + 1e-6) + predicted_error = predicted_log.exp() + for index, prompt_id in enumerate(validation_ids): + local_rows.append( + { + "prompt_id": prompt_id, + "chunk": chunk, + "target_step": target_step, + "target_hidden_nrmse": float(target_error[index]), + "predicted_hidden_nrmse": float(predicted_error[index]), + "target_log_error": float(target_log[index]), + "predicted_log_error": float(predicted_log[index]), + } + ) + del batch, pred_hidden, transformed, target_error, predicted_log + del target_log, predicted_error, chunk_position, step_id + torch.cuda.empty_cache() + + gathered: list[list[dict[str, float | int]] | None] | None = ( + [None] * world_size if rank == 0 else None + ) + dist.gather_object(local_rows, gathered, dst=0) + if rank != 0: + return None, None + assert gathered is not None + rows = [row for shard in gathered if shard is not None for row in shard] + rows.sort(key=lambda row: (int(row["chunk"]), int(row["target_step"]), int(row["prompt_id"]))) + return validation_metrics(rows), rows + + +def main() -> None: + args = parse_args() + args.dataset_root = resolve(args.dataset_root) + args.output_dir = resolve(args.output_dir) + args.predictor_weights = resolve(args.predictor_weights) + args.checkpoint_path = resolve(args.checkpoint_path) + args.config_path = resolve(args.config_path) + + local_rank = int(os.environ["LOCAL_RANK"]) + torch.cuda.set_device(local_rank) + dist.init_process_group("nccl", device_id=torch.device("cuda", local_rank)) + rank = dist.get_rank() + world_size = dist.get_world_size() + if world_size != 4: + raise ValueError(f"This training split expects four ranks, got {world_size}") + device = torch.device("cuda", local_rank) + is_main = rank == 0 + if is_main: + args.output_dir.mkdir(parents=True, exist_ok=True) + dist.barrier() + + set_seed(args.seed + rank) + torch.set_num_threads(4) + torch.set_num_interop_threads(1) + torch.backends.cuda.matmul.allow_tf32 = True + torch.set_float32_matmul_precision("high") + + print(f"[rank {rank}] loading frozen Teacher", flush=True) + teacher = load_teacher(args.checkpoint_path, args.config_path, device) + print(f"[rank {rank}] loading frozen Predictor", flush=True) + predictor, predictor_config = load_predictor( + teacher, + args.predictor_weights, + device, + expected_input_variant=args.predictor_input_variant, + ) + + head = ConfidenceTokenHead(dropout=args.dropout, num_steps=3).to(device=device) + ddp = DDP( + head, + device_ids=[local_rank], + output_device=local_rank, + broadcast_buffers=False, + ) + optimizer = AdamW( + ddp.parameters(), + lr=( + args.warmup_start_lr + if args.lr_schedule == "warmup_cosine" + else args.learning_rate + ), + betas=(0.9, 0.999), + weight_decay=args.weight_decay, + ) + dataset = LazyLayer17Dataset(args.dataset_root, layer_id=17) + parameter_count = sum(parameter.numel() for parameter in ddp.module.parameters()) + if parameter_count != 4_872_065: + raise RuntimeError(f"Unexpected Confidence Head size: {parameter_count}") + + if is_main: + atomic_json( + args.output_dir / "config.json", + { + **{ + key: str(value) if isinstance(value, Path) else value + for key, value in vars(args).items() + }, + "world_size": world_size, + "effective_global_batch_size": world_size + * args.per_device_batch_size, + "train_prompt_ids": "0..899", + "validation_prompt_ids": "900..999", + "candidate_steps": [1, 2, 3], + "chunks": [1, 2, 3, 4, 5, 6], + "samples_per_epoch": 900 * 6 * 3, + "optimizer_steps_per_epoch": 72, + "maximum_optimizer_steps": args.epochs * 72, + "optimizer": { + "name": "AdamW", + "learning_rate": args.learning_rate, + "betas": [0.9, 0.999], + "weight_decay": args.weight_decay, + "scheduler": args.lr_schedule, + "warmup_steps": args.warmup_steps, + "warmup_start_lr": args.warmup_start_lr, + "min_learning_rate": args.min_learning_rate, + "grad_clip": args.grad_clip, + }, + "loss": "SmoothL1(predicted_log_hidden_nrmse, target_log_hidden_nrmse)", + "head_parameters": parameter_count, + "predictor_config": predictor_config, + "from_scratch": True, + }, + ) + + best_loss = float("inf") + patience_loss = float("inf") + best_epoch = 0 + stale_epochs = 0 + total_steps = 0 + started = time.perf_counter() + train_log = args.output_dir / "train_log.jsonl" + best_path = args.output_dir / "confidence_best.safetensors" + latest_path = args.output_dir / "confidence_latest.safetensors" + metadata = { + "head_config": json.dumps( + { + "architecture": "ConfidenceTokenHead", + "dim": 1536, + "token_dim": 512, + "context_dim": 64, + "num_heads": 8, + "ffn_dim": 2048, + "dropout": args.dropout, + "num_steps": 3, + }, + sort_keys=True, + ), + "predictor_weights": str(args.predictor_weights), + "train_prompt_ids": "0..899", + "validation_prompt_ids": "900..999", + "learning_rate_schedule": json.dumps( + { + "name": args.lr_schedule, + "peak_lr": args.learning_rate, + "warmup_steps": args.warmup_steps, + "warmup_start_lr": args.warmup_start_lr, + "min_lr": args.min_learning_rate, + "planned_total_steps": args.epochs * 72, + }, + sort_keys=True, + ), + } + + for epoch in range(1, args.epochs + 1): + ddp.train() + epoch_started = time.perf_counter() + loss_sum = torch.zeros(2, dtype=torch.float64, device=device) + schedule = epoch_batches( + epoch=epoch, + rank=rank, + world_size=world_size, + per_device_batch_size=args.per_device_batch_size, + seed=args.seed, + ) + for chunk, target_step, prompt_ids in schedule: + current_lr = learning_rate_at_step( + step=total_steps, + total_steps=args.epochs * 72, + schedule=args.lr_schedule, + peak_lr=args.learning_rate, + warmup_steps=args.warmup_steps, + warmup_start_lr=args.warmup_start_lr, + min_lr=args.min_learning_rate, + ) + optimizer.param_groups[0]["lr"] = current_lr + optimizer.zero_grad(set_to_none=True) + cpu_batch = collate_lazy_samples( + [dataset[(prompt_id, chunk, target_step)] for prompt_id in prompt_ids] + ) + batch = move_batch( + cpu_batch, + teacher, + device, + args.history_projection_batch_size, + ) + del cpu_batch + pred_hidden, transformed = predictor_features( + predictor, batch, teacher, device + ) + target_log = torch.log( + hidden_nrmse(pred_hidden, batch["target_hidden"]) + 1e-6 + ).detach() + chunk_position = torch.full( + (len(prompt_ids),), (chunk - 1) / 5.0, device=device + ) + step_id = torch.full( + (len(prompt_ids),), target_step, dtype=torch.long, device=device + ) + with torch.autocast(device_type="cuda", dtype=torch.bfloat16): + predicted_log = ddp( + transformed_hidden=transformed, + pred_hidden=pred_hidden, + anchor_hidden=batch["anchor_hidden"], + chunk_position=chunk_position, + step_id=step_id, + ) + loss = F.smooth_l1_loss(predicted_log, target_log) + loss.backward() + torch.nn.utils.clip_grad_norm_(ddp.module.parameters(), args.grad_clip) + optimizer.step() + local_count = len(prompt_ids) + loss_sum += torch.tensor( + [float(loss.detach()) * local_count, local_count], + dtype=torch.float64, + device=device, + ) + total_steps += 1 + del batch, pred_hidden, transformed, target_log, predicted_log, loss + del chunk_position, step_id + + dist.all_reduce(loss_sum) + train_loss = float(loss_sum[0] / loss_sum[1]) + validation, validation_rows = run_validation( + ddp=ddp, + predictor=predictor, + dataset=dataset, + teacher=teacher, + device=device, + rank=rank, + world_size=world_size, + history_projection_batch_size=args.history_projection_batch_size, + ) + decision = torch.zeros(4, dtype=torch.float64, device=device) + if is_main: + assert validation is not None and validation_rows is not None + validation_loss = float(validation["huber_log"]) + if validation_loss < best_loss: + best_loss = validation["huber_log"] + best_epoch = epoch + save_weights(ddp.module, best_path, metadata=metadata) + write_predictions( + args.output_dir / "validation_best_predictions.csv", + validation_rows, + ) + if validation_loss < patience_loss - args.min_delta: + patience_loss = validation_loss + stale_epochs = 0 + else: + stale_epochs += 1 + save_weights(ddp.module, latest_path, metadata=metadata) + record = { + "epoch": epoch, + "optimizer_step": total_steps, + "learning_rate": current_lr, + "train_huber_log": train_loss, + "validation": validation, + "best_epoch": best_epoch, + "best_validation_huber_log": best_loss, + "early_stop_reference_huber_log": patience_loss, + "stale_epochs": stale_epochs, + "epoch_time_s": time.perf_counter() - epoch_started, + "peak_gpu_gib": torch.cuda.max_memory_allocated() / 2**30, + } + append_jsonl(train_log, record) + atomic_json( + args.output_dir / "status.json", + { + "status": "running", + **record, + "elapsed_s": time.perf_counter() - started, + }, + ) + print( + f"[epoch] {epoch}/{args.epochs} step={total_steps} " + f"train={train_loss:.6f} val={validation['huber_log']:.6f} " + f"rho={validation['spearman_nrmse']:.4f} " + f"best={best_epoch} stale={stale_epochs} " + f"time={record['epoch_time_s']:.1f}s", + flush=True, + ) + decision[:] = torch.tensor( + [best_loss, patience_loss, best_epoch, stale_epochs], + dtype=torch.float64, + device=device, + ) + dist.broadcast(decision, src=0) + best_loss = float(decision[0]) + patience_loss = float(decision[1]) + best_epoch = int(decision[2]) + stale_epochs = int(decision[3]) + if stale_epochs >= args.patience: + if is_main: + print(f"[early-stop] stale epochs={stale_epochs}", flush=True) + break + + dist.barrier() + if is_main: + atomic_json( + args.output_dir / "status.json", + { + "status": "complete", + "best_epoch": best_epoch, + "best_validation_huber_log": best_loss, + "completed_epochs": epoch, + "optimizer_steps": total_steps, + "training_time_s": time.perf_counter() - started, + "selected_checkpoint": str(best_path), + }, + ) + print(f"[complete] best_epoch={best_epoch} path={best_path}", flush=True) + dist.destroy_process_group() + + +if __name__ == "__main__": + main() diff --git a/scripts/train_layer17_confidence.py b/scripts/train_layer17_confidence.py new file mode 100644 index 0000000000000000000000000000000000000000..b42c6cd99c86dbd6039f0677b5a0865c9497eaf2 --- /dev/null +++ b/scripts/train_layer17_confidence.py @@ -0,0 +1,487 @@ +#!/usr/bin/env python3 +"""Train a teacher-forced confidence head for the frozen Layer-17 Predictor.""" + +from __future__ import annotations + +import argparse +import csv +import json +import math +import os +import random +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 torch +import torch.nn.functional as F +from safetensors.torch import load_file, save_file +from torch.optim import AdamW + +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 predictor_training.offline_data import OfflinePredictorStore, TOKENS_PER_CHUNK +from predictor_training.single_block import SingleBlockPredictor, initialize_predictor_block +from scripts.run_single_block_init_sweep import frozen_inputs, load_teacher, move_batch +from utils.misc import set_seed + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--gpu", default=PHYSICAL_GPU) + parser.add_argument( + "--dataset_root", type=Path, + default=Path("outputs/predictor_offline_100_all_blocks"), + ) + parser.add_argument( + "--checkpoint_path", type=Path, + default=Path("checkpoints/self_forcing_dmd.pt"), + ) + parser.add_argument( + "--config_path", type=Path, default=Path("configs/self_forcing_sid.yaml") + ) + parser.add_argument( + "--predictor_weights", type=Path, + default=Path( + "outputs/single_block_init_sweep/teacher_layer_17/" + "predictor_final.safetensors" + ), + ) + parser.add_argument( + "--output_dir", type=Path, + default=Path("outputs/layer17_confidence_teacher_forced"), + ) + parser.add_argument("--epochs", type=int, default=20) + parser.add_argument("--batch_size", type=int, default=8) + parser.add_argument("--eval_batch_size", type=int, default=10) + parser.add_argument("--learning_rate", type=float, default=3e-4) + parser.add_argument("--weight_decay", type=float, default=0.01) + parser.add_argument("--dropout", type=float, default=0.1) + parser.add_argument( + "--candidate_steps", type=int, nargs="+", choices=(1, 2, 3), + default=[1, 2], + ) + parser.add_argument("--init_confidence_weights", type=Path, default=None) + parser.add_argument("--patience", type=int, default=5) + parser.add_argument("--min_delta", type=float, default=1e-4) + parser.add_argument("--grad_clip", type=float, default=1.0) + parser.add_argument("--seed", type=int, default=0) + args = parser.parse_args() + for name in ( + "dataset_root", "checkpoint_path", "config_path", "predictor_weights", + "output_dir", "init_confidence_weights", + ): + 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()) + if args.epochs < 1 or args.batch_size < 1 or args.eval_batch_size < 1: + parser.error("epochs and batch sizes must be positive") + return args + + +def atomic_json(path: Path, value: Any) -> None: + temporary = path.with_suffix(path.suffix + ".tmp") + temporary.write_text(json.dumps(value, indent=2) + "\n", encoding="utf-8") + os.replace(temporary, path) + + +def save_weights(model: torch.nn.Module, path: Path) -> None: + tensors = { + key: value.detach().cpu().contiguous() + for key, value in model.state_dict().items() + } + temporary = path.with_suffix(path.suffix + ".tmp") + save_file(tensors, temporary) + os.replace(temporary, path) + + +def load_predictor( + teacher: torch.nn.Module, path: Path, device: torch.device +) -> SingleBlockPredictor: + predictor = SingleBlockPredictor( + initialize_predictor_block(teacher.blocks[17], "teacher_full"), + dim=teacher.dim, + gradient_checkpointing=False, + ) + predictor.load_state_dict(load_file(str(path), device="cpu"), strict=True) + return predictor.to(device=device).eval().requires_grad_(False) + + +@torch.no_grad() +def predictor_features( + predictor: SingleBlockPredictor, + batch: dict[str, Any], + teacher: torch.nn.Module, + device: torch.device, +) -> tuple[torch.Tensor, torch.Tensor]: + frozen = frozen_inputs(batch, teacher, device) + with torch.autocast(device_type="cuda", dtype=torch.bfloat16): + output = predictor( + current_tokens=frozen["current_tokens"], + anchor_hidden=batch["anchor_hidden"], + previous_hidden=batch["previous_hidden"], + timestep_modulation=frozen["timestep_modulation"], + grid_sizes=frozen["grid_sizes"], + freqs=frozen["freqs"], + history_k=batch["history_k"], + history_v=batch["history_v"], + cross_k=batch["cross_k"], + cross_v=batch["cross_v"], + current_start=batch["chunk"] * TOKENS_PER_CHUNK, + return_features=True, + ) + if not isinstance(output, tuple): + raise RuntimeError("Predictor did not return confidence features") + return output + + +def hidden_nrmse(predicted: torch.Tensor, target: torch.Tensor) -> torch.Tensor: + error_energy = (predicted.float() - target.float()).square().sum(dim=(1, 2)) + target_energy = target.float().square().sum(dim=(1, 2)) + return torch.sqrt(error_energy / target_energy.clamp_min(1e-8)) + + +def ranks(values: list[float]) -> list[float]: + order = sorted(range(len(values)), key=values.__getitem__) + result = [0.0] * len(values) + offset = 0 + while offset < len(order): + end = offset + 1 + while end < len(order) and values[order[end]] == values[order[offset]]: + end += 1 + rank = 0.5 * (offset + end - 1) + for index in order[offset:end]: + result[index] = rank + offset = end + return result + + +def correlation(left: list[float], right: list[float]) -> float: + left_mean = sum(left) / len(left) + right_mean = sum(right) / len(right) + numerator = sum((a - left_mean) * (b - right_mean) for a, b in zip(left, right)) + denominator = math.sqrt( + sum((a - left_mean) ** 2 for a in left) + * sum((b - right_mean) ** 2 for b in right) + ) + return numerator / denominator if denominator > 0 else float("nan") + + +def evaluate( + *, + head: PredictorConfidenceHead, + predictor: SingleBlockPredictor, + store: OfflinePredictorStore, + prompt_ids: list[int], + batch_size: int, + teacher: torch.nn.Module, + device: torch.device, + candidate_steps: list[int], +) -> tuple[dict[str, float], list[dict[str, float | int]]]: + head.eval() + rows: list[dict[str, float | int]] = [] + losses: list[float] = [] + for chunk in range(1, 7): + for target_step in candidate_steps: + for start in range(0, len(prompt_ids), batch_size): + selected_ids = prompt_ids[start : start + batch_size] + batch = move_batch(store.batch(selected_ids, chunk, target_step), device) + pred_hidden, transformed = predictor_features( + predictor, batch, teacher, device + ) + target_error = hidden_nrmse(pred_hidden, batch["target_hidden"]) + count = len(selected_ids) + chunk_position = torch.full( + (count,), (chunk - 1) / 5.0, device=device + ) + step_id = torch.full( + (count,), target_step, dtype=torch.long, device=device + ) + with torch.no_grad(), torch.autocast( + device_type="cuda", dtype=torch.bfloat16 + ): + predicted_log = head( + transformed_hidden=transformed, + pred_hidden=pred_hidden, + anchor_hidden=batch["anchor_hidden"], + chunk_position=chunk_position, + step_id=step_id, + ) + target_log = torch.log(target_error + 1e-6) + losses.extend( + F.smooth_l1_loss(predicted_log, target_log, reduction="none") + .cpu().tolist() + ) + predicted_error = predicted_log.exp() + for offset, prompt_id in enumerate(selected_ids): + rows.append( + { + "prompt_id": prompt_id, + "chunk": chunk, + "target_step": target_step, + "target_hidden_nrmse": float(target_error[offset]), + "predicted_hidden_nrmse": float(predicted_error[offset]), + "target_log_error": float(target_log[offset]), + "predicted_log_error": float(predicted_log[offset]), + } + ) + del batch, pred_hidden, transformed, target_error, predicted_log + target = [float(row["target_hidden_nrmse"]) for row in rows] + predicted = [float(row["predicted_hidden_nrmse"]) for row in rows] + log_target = [float(row["target_log_error"]) for row in rows] + log_predicted = [float(row["predicted_log_error"]) for row in rows] + metrics = { + "huber_log": sum(losses) / len(losses), + "mae_log": sum(abs(a - b) for a, b in zip(log_target, log_predicted)) / len(rows), + "mae_nrmse": sum(abs(a - b) for a, b in zip(target, predicted)) / len(rows), + "pearson_nrmse": correlation(target, predicted), + "spearman_nrmse": correlation(ranks(target), ranks(predicted)), + "target_mean": sum(target) / len(target), + "prediction_mean": sum(predicted) / len(predicted), + "num_samples": float(len(rows)), + } + grouped: dict[tuple[int, int], list[dict[str, float | int]]] = {} + for row in rows: + key = (int(row["chunk"]), int(row["target_step"])) + grouped.setdefault(key, []).append(row) + target_residual: list[float] = [] + prediction_residual: list[float] = [] + cell_spearman: list[float] = [] + for selected in grouped.values(): + cell_target = [float(row["target_hidden_nrmse"]) for row in selected] + cell_prediction = [ + float(row["predicted_hidden_nrmse"]) for row in selected + ] + target_mean = sum(cell_target) / len(cell_target) + prediction_mean = sum(cell_prediction) / len(cell_prediction) + target_residual.extend(value - target_mean for value in cell_target) + prediction_residual.extend(value - prediction_mean for value in cell_prediction) + cell_spearman.append( + correlation(ranks(cell_target), ranks(cell_prediction)) + ) + metrics.update( + { + "controlled_chunk_step_spearman": correlation( + ranks(target_residual), ranks(prediction_residual) + ), + "mean_within_chunk_step_spearman": sum(cell_spearman) + / len(cell_spearman), + "positive_chunk_step_cells": float( + sum(value > 0 for value in cell_spearman) + ), + "num_chunk_step_cells": float(len(cell_spearman)), + } + ) + return metrics, rows + + +def write_predictions(path: Path, rows: list[dict[str, float | int]]) -> None: + fields = list(rows[0]) + temporary = path.with_suffix(path.suffix + ".tmp") + with temporary.open("w", newline="", encoding="utf-8") as handle: + writer = csv.DictWriter(handle, fieldnames=fields) + writer.writeheader() + writer.writerows(rows) + os.replace(temporary, path) + + +def training_batches( + prompt_ids: list[int], batch_size: int, seed: int, + candidate_steps: list[int], +) -> list[tuple[int, int, list[int]]]: + rng = random.Random(seed) + groups = [ + (chunk, step) for chunk in range(1, 7) for step in candidate_steps + ] + rng.shuffle(groups) + batches = [] + for chunk, target_step in groups: + shuffled = prompt_ids.copy() + rng.shuffle(shuffled) + for start in range(0, len(shuffled), batch_size): + batches.append((chunk, target_step, shuffled[start : start + batch_size])) + return batches + + +def main() -> None: + args = parse_args() + args.output_dir.mkdir(parents=True, exist_ok=True) + set_seed(args.seed) + device = torch.device("cuda") + train_ids = list(range(0, 80)) + val_ids = list(range(80, 90)) + test_ids = list(range(90, 100)) + manifest = { + "status": "running", + "physical_gpu": str(args.gpu), + "predictor_weights": str(args.predictor_weights), + "train_prompt_ids": train_ids, + "val_prompt_ids": val_ids, + "test_prompt_ids": test_ids, + "teacher_forced_only": True, + "target": "log_hidden_nrmse", + "chunk_risk_applied_after_head": True, + "candidate_steps": args.candidate_steps, + "args": {key: str(value) if isinstance(value, Path) else value for key, value in vars(args).items()}, + } + atomic_json(args.output_dir / "manifest.json", manifest) + + print("[setup] loading offline data", flush=True) + store = OfflinePredictorStore(args.dataset_root, range(100)) + print("[setup] loading frozen Teacher", flush=True) + teacher = load_teacher(args.checkpoint_path, args.config_path, device) + print("[setup] loading Layer-17 cache", flush=True) + store.load_layer_cache(17, teacher, device) + print("[setup] loading frozen Predictor", flush=True) + predictor = load_predictor(teacher, args.predictor_weights, device) + head = PredictorConfidenceHead( + dropout=args.dropout, num_steps=max(args.candidate_steps) + ).to(device=device) + if args.init_confidence_weights is not None: + initial = load_file(str(args.init_confidence_weights), device="cpu") + current = head.state_dict() + for key, value in initial.items(): + if key == "step_embedding.weight": + rows = min(value.shape[0], current[key].shape[0]) + current[key][:rows].copy_(value[:rows]) + elif key in current and value.shape == current[key].shape: + current[key].copy_(value) + head.load_state_dict(current, strict=True) + print(f"[setup] warm-started Head from {args.init_confidence_weights}", flush=True) + parameter_count = sum(parameter.numel() for parameter in head.parameters()) + print(f"[setup] confidence parameters={parameter_count:,}", flush=True) + expected_parameters = 1_107_650 + 64 * (max(args.candidate_steps) - 2) + if parameter_count != expected_parameters: + raise RuntimeError(f"Unexpected confidence parameter count: {parameter_count}") + optimizer = AdamW( + head.parameters(), lr=args.learning_rate, weight_decay=args.weight_decay + ) + + history: list[dict[str, Any]] = [] + best_loss = float("inf") + best_epoch = 0 + stale_epochs = 0 + best_path = args.output_dir / "confidence_best.safetensors" + started = time.perf_counter() + for epoch in range(1, args.epochs + 1): + head.train() + epoch_losses: list[float] = [] + epoch_started = time.perf_counter() + for chunk, target_step, selected_ids in training_batches( + train_ids, args.batch_size, args.seed + epoch, args.candidate_steps + ): + batch = move_batch(store.batch(selected_ids, chunk, target_step), device) + pred_hidden, transformed = predictor_features( + predictor, batch, teacher, device + ) + target_error = hidden_nrmse(pred_hidden, batch["target_hidden"]) + target_log = torch.log(target_error + 1e-6) + count = len(selected_ids) + chunk_position = torch.full( + (count,), (chunk - 1) / 5.0, device=device + ) + step_id = torch.full( + (count,), target_step, dtype=torch.long, device=device + ) + with torch.autocast(device_type="cuda", dtype=torch.bfloat16): + predicted_log = head( + transformed_hidden=transformed, + pred_hidden=pred_hidden, + anchor_hidden=batch["anchor_hidden"], + chunk_position=chunk_position, + step_id=step_id, + ) + loss = F.smooth_l1_loss(predicted_log, target_log) + optimizer.zero_grad(set_to_none=True) + loss.backward() + torch.nn.utils.clip_grad_norm_(head.parameters(), args.grad_clip) + optimizer.step() + epoch_losses.append(float(loss.detach())) + del batch, pred_hidden, transformed, target_error, predicted_log, loss + + validation, val_rows = evaluate( + head=head, predictor=predictor, store=store, prompt_ids=val_ids, + batch_size=args.eval_batch_size, teacher=teacher, device=device, + candidate_steps=args.candidate_steps, + ) + record = { + "epoch": epoch, + "train_huber_log": sum(epoch_losses) / len(epoch_losses), + "validation": validation, + "epoch_time_s": time.perf_counter() - epoch_started, + } + history.append(record) + improved = validation["huber_log"] < best_loss - args.min_delta + if improved: + best_loss = validation["huber_log"] + best_epoch = epoch + stale_epochs = 0 + save_weights(head, best_path) + write_predictions(args.output_dir / "validation_best_predictions.csv", val_rows) + else: + stale_epochs += 1 + atomic_json(args.output_dir / "train_history.json", history) + print( + f"[epoch] {epoch}/{args.epochs} train={record['train_huber_log']:.6f} " + f"val={validation['huber_log']:.6f} " + f"rho={validation['spearman_nrmse']:.4f} " + f"best={best_epoch} time={record['epoch_time_s']:.1f}s", + flush=True, + ) + if stale_epochs >= args.patience: + print(f"[early-stop] no improvement for {stale_epochs} epochs", flush=True) + break + + head.load_state_dict(load_file(str(best_path), device="cpu"), strict=True) + head.to(device=device) + validation, val_rows = evaluate( + head=head, predictor=predictor, store=store, prompt_ids=val_ids, + batch_size=args.eval_batch_size, teacher=teacher, device=device, + candidate_steps=args.candidate_steps, + ) + test, test_rows = evaluate( + head=head, predictor=predictor, store=store, prompt_ids=test_ids, + batch_size=args.eval_batch_size, teacher=teacher, device=device, + candidate_steps=args.candidate_steps, + ) + write_predictions(args.output_dir / "validation_predictions.csv", val_rows) + write_predictions(args.output_dir / "test_predictions.csv", test_rows) + result = { + "status": "complete", + "parameter_count": parameter_count, + "best_epoch": best_epoch, + "validation": validation, + "test": test, + "training_time_s": time.perf_counter() - started, + } + atomic_json(args.output_dir / "metrics.json", result) + manifest["status"] = "complete" + manifest["result"] = result + atomic_json(args.output_dir / "manifest.json", manifest) + print( + f"[complete] best_epoch={best_epoch} " + f"val_rho={validation['spearman_nrmse']:.4f} " + f"test_rho={test['spearman_nrmse']:.4f} -> {args.output_dir}", + flush=True, + ) + + +if __name__ == "__main__": + main() diff --git a/scripts/train_layer17_predictor_stage2_dmd.py b/scripts/train_layer17_predictor_stage2_dmd.py new file mode 100644 index 0000000000000000000000000000000000000000..bfec8e53a7e97418a92fbb2ac63b0d0aad0a09a8 --- /dev/null +++ b/scripts/train_layer17_predictor_stage2_dmd.py @@ -0,0 +1,1234 @@ +#!/usr/bin/env python3 +"""Stage-2 random-exit DMD training for a Stage-1 Layer-17 Predictor. + +The frozen Self-Forcing generator supplies Full calls and clean causal history. +For each student update one exit in P1/P2/P3 is sampled and shared across all +distributed ranks. Chunk 0 remains Full, while chunks 1..6 use the Predictor +after step 0. Only the Predictor call at the sampled exit retains a graph. + +Stage-1 Predictor metadata is used to reconstruct both the original concat +model and ATC variants. The saved raw/EMA checkpoints preserve that metadata +and can therefore be consumed directly by the ConfidenceTokenHead trainer. +""" + +from __future__ import annotations + +import argparse +import contextlib +import json +import math +import os +import random +import sys +import time +from pathlib import Path +from typing import Any + +import torch +import torch.distributed as dist +import torch.nn.functional as F +from omegaconf import OmegaConf +from safetensors import safe_open +from safetensors.torch import load_file, save_file +from torch.nn.parallel import DistributedDataParallel as DDP + +ROOT = Path(__file__).resolve().parents[1] +if str(ROOT) not in sys.path: + sys.path.insert(0, str(ROOT)) + +from pipeline import CausalInferencePipeline # noqa: E402 +from predictor_training.offline_data import TOKENS_PER_CHUNK # noqa: E402 +from predictor_training.single_block import ( # noqa: E402 + SingleBlockPredictor, + initialize_predictor_block, +) +from scripts.run_single_block_init_sweep import forward_predictor # noqa: E402 +from utils.distributed import fsdp_wrap # noqa: E402 +from utils.misc import set_seed # noqa: E402 +from utils.wan_wrapper import WanDiffusionWrapper, WanTextEncoder # noqa: E402 + + +NUM_CHUNKS = 7 +FRAMES_PER_CHUNK = 3 +NUM_STEPS = 4 +LATENT_SHAPE = (16, 60, 104) + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument( + "--config_path", + type=Path, + default=Path("configs/self_forcing_dmd.yaml"), + ) + parser.add_argument( + "--generator_ckpt", + type=Path, + default=Path("checkpoints/self_forcing_dmd.pt"), + ) + parser.add_argument("--predictor_init", type=Path, required=True) + parser.add_argument( + "--predictor_input_variant", + choices=("auto", "self_forcing", "atc"), + default="auto", + help="Validate the Stage-1 checkpoint variant; auto trusts metadata.", + ) + parser.add_argument( + "--prompt_path", + type=Path, + default=Path("prompts/vidprom_filtered_extended.txt"), + ) + parser.add_argument("--output_dir", type=Path, required=True) + parser.add_argument( + "--real_score_name", + default="/data1/chenzhuo/Wan2.1/Wan2.1-T2V-14B", + help="Wan real-score model name or absolute checkpoint directory.", + ) + parser.add_argument("--fake_score_name", default="Wan2.1-T2V-1.3B") + parser.add_argument("--student_steps", type=int, default=2000) + parser.add_argument("--critic_updates_per_student", type=int, default=5) + parser.add_argument("--gradient_accumulation_steps", type=int, default=2) + parser.add_argument( + "--prompt_count", + type=int, + default=0, + help="Number of shuffled prompts to use; 0 uses the complete prompt file.", + ) + parser.add_argument( + "--prompt_seed", + type=int, + default=0, + help="Seed for the shared, deterministic prompt permutation.", + ) + parser.add_argument("--fusion_lr", type=float, default=1.0e-5) + parser.add_argument("--block_lr", type=float, default=1.0e-6) + parser.add_argument("--critic_lr", type=float, default=4.0e-7) + parser.add_argument("--weight_decay", type=float, default=0.01) + parser.add_argument("--warmup_steps", type=int, default=100) + parser.add_argument("--predictor_grad_clip", type=float, default=1.0) + parser.add_argument("--critic_grad_clip", type=float, default=10.0) + parser.add_argument("--guidance_scale", type=float, default=3.0) + parser.add_argument("--timestep_shift", type=float, default=5.0) + parser.add_argument("--ema_start", type=int, default=200) + parser.add_argument("--ema_decay", type=float, default=0.99) + parser.add_argument("--seed", type=int, default=0) + parser.add_argument("--save_every", type=int, default=100) + parser.add_argument("--state_save_every", type=int, default=500) + parser.add_argument("--log_every", type=int, default=1) + parser.add_argument("--expected_world_size", type=int, default=4) + parser.add_argument("--swanlab_project", default="Self-Forcing-Predictor-DMD") + parser.add_argument("--swanlab_name", default="layer17-random-exit-dmd") + parser.add_argument("--swanlab_id", default=None) + parser.add_argument("--disable_swanlab", action="store_true") + parser.add_argument( + "--resume_from", + type=Path, + default=None, + help="training_latest.pt containing Predictor, EMA, fake score and optimizers.", + ) + parser.add_argument("--smoke_steps", type=int, default=None) + parser.add_argument( + "--skip_real_score", + action="store_true", + help="Graph smoke test only; forbidden for formal training.", + ) + args = parser.parse_args() + positive = ( + "student_steps", + "critic_updates_per_student", + "gradient_accumulation_steps", + "save_every", + "state_save_every", + "expected_world_size", + ) + if any(getattr(args, name) < 1 for name in positive): + parser.error("steps, counts, intervals and world size must be positive") + if args.prompt_count < 0: + parser.error("prompt_count must be non-negative (0 means all prompts)") + if not 0.0 < args.ema_decay < 1.0: + parser.error("ema_decay must be in (0, 1)") + return args + + +def resolve(path: Path) -> Path: + path = path.expanduser() + return path.resolve() if path.is_absolute() else (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) + "\n", + encoding="utf-8", + ) + os.replace(temporary, path) + + +def append_jsonl(path: Path, value: Any) -> None: + with path.open("a", encoding="utf-8") as handle: + handle.write(json.dumps(value, ensure_ascii=False) + "\n") + + +def load_prompt_pool(path: Path, prompt_count: int, prompt_seed: int) -> list[str]: + """Build one rank-independent prompt permutation, optionally truncated.""" + prompts = [ + line.strip() + for line in path.read_text(encoding="utf-8").splitlines() + if line.strip() + ] + if not prompts: + raise ValueError(f"No non-empty prompts found in {path}") + if prompt_count > len(prompts): + raise ValueError(f"Requested {prompt_count} prompts, found {len(prompts)}") + + random.Random(prompt_seed).shuffle(prompts) + return prompts if prompt_count == 0 else prompts[:prompt_count] + + +def prompt_sample_id( + *, + student_step: int, + batch_slot: int, + global_batch_size: int, + batches_per_step: int, + accumulation_index: int, + world_size: int, + rank: int, + prompt_pool_size: int, +) -> int: + """Assign consecutive draws so prompts repeat only after a complete pool pass.""" + draw_id = ( + student_step * batches_per_step * global_batch_size + + batch_slot * global_batch_size + + accumulation_index * world_size + + rank + ) + return draw_id % prompt_pool_size + + +def read_predictor_config( + weights: Path, expected_input_variant: str = "auto" +) -> dict[str, Any]: + with safe_open(weights, framework="pt", device="cpu") as handle: + metadata = handle.metadata() or {} + raw = metadata.get("predictor_config") + if raw is None: + raise ValueError( + f"Stage-2 requires predictor_config metadata in the Stage-1 checkpoint: {weights}" + ) + config = json.loads(raw) + actual = str(config.get("input_variant", "self_forcing")) + if actual not in {"self_forcing", "atc"}: + raise ValueError(f"Stage-2 does not support input_variant={actual!r}") + if expected_input_variant != "auto" and actual != expected_input_variant: + raise ValueError( + "Predictor input variant mismatch: " + f"expected={expected_input_variant} actual={actual}" + ) + return config + + +class FinalHiddenCapture: + """Capture the frozen Full DiT hidden immediately before its output head.""" + + 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 _hook(self, _module: torch.nn.Module, inputs: tuple[torch.Tensor, ...]) -> None: + if self.enabled: + if self.value is not None: + raise RuntimeError("Teacher head 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 close(self) -> None: + self.handle.remove() + + +def reset_caches( + pipeline: CausalInferencePipeline, device: torch.device, batch_size: int +) -> None: + if batch_size != 1: + raise ValueError( + "The current causal KV-cache implementation requires batch size 1" + ) + if pipeline.kv_cache1 is None: + pipeline._initialize_kv_cache(batch_size, torch.bfloat16, device) + pipeline._initialize_crossattn_cache(batch_size, torch.bfloat16, device) + 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 + + +def shifted_timestep(value: torch.Tensor, shift: float) -> torch.Tensor: + value = value.float() + shifted = ( + shift * (value / 1000.0) / (1.0 + (shift - 1.0) * (value / 1000.0)) * 1000.0 + ) + return shifted.clamp(20.0, 980.0) + + +def lr_factor(step: int, total: int, warmup: int) -> float: + if step < warmup: + return float(step + 1) / max(1, warmup) + progress = float(step - warmup) / max(1, total - warmup) + return 0.5 * (1.0 + math.cos(math.pi * min(1.0, progress))) + + +def set_optimizer_lrs( + optimizer: torch.optim.Optimizer, bases: list[float], factor: float +) -> None: + if len(optimizer.param_groups) != len(bases): + raise ValueError("Learning-rate bases do not match optimizer groups") + for group, base in zip(optimizer.param_groups, bases): + group["lr"] = base * factor + + +def build_pipeline( + config: Any, checkpoint: Path, 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=torch.nn.Identity(), + ) + state = torch.load(checkpoint, map_location="cpu", weights_only=False, mmap=True) + if "generator_ema" not in state: + raise KeyError( + f"Self-Forcing checkpoint has no generator_ema: keys={sorted(state)}" + ) + pipeline.generator.load_state_dict(state["generator_ema"], strict=True) + del state + pipeline.generator.to(device=device, dtype=torch.bfloat16) + pipeline.generator.eval().requires_grad_(False) + pipeline.scheduler.timesteps = pipeline.scheduler.timesteps.to(device) + return pipeline + + +def build_predictor( + teacher: torch.nn.Module, + weights: Path, + config: dict[str, Any], + device: torch.device, +) -> SingleBlockPredictor: + source_layer = int(config.get("source_layer", 17)) + predictor = SingleBlockPredictor( + block=initialize_predictor_block(teacher.blocks[source_layer], "teacher_full"), + dim=teacher.dim, + gradient_checkpointing=True, + input_variant=str(config.get("input_variant", "self_forcing")), + gate_mode=str(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)), + atc_previous_scope=str(config.get("atc_previous_scope", "chunk")), + atc_freq_dim=int(config.get("atc_freq_dim", 256)), + atc_mlp_hidden_dim=int(config.get("atc_mlp_hidden_dim", 3072)), + atc_gate_hidden_dim=int(config.get("atc_gate_hidden_dim", 512)), + atc_transport_residual_scale=float( + config.get("atc_transport_residual_scale", 0.1) + ), + atc_gate_initial_probability=float( + config.get("atc_gate_initial_probability", 0.3) + ), + atc_collect_diagnostics=False, + ) + predictor.load_state_dict(load_file(str(weights), device="cpu"), strict=True) + + # Cross-attention K/V are supplied from the Full DiT cache. These modules + # are bypassed and must not be included in a find_unused_parameters=False DDP. + for module in ( + predictor.block.cross_attn.k, + predictor.block.cross_attn.v, + predictor.block.cross_attn.norm_k, + ): + module.requires_grad_(False) + return predictor.float().to(device).train() + + +def online_predictor_step( + *, + predictor: torch.nn.Module, + teacher: torch.nn.Module, + noisy_input: torch.Tensor, + timestep: torch.Tensor, + anchor_timestep: torch.Tensor, + anchor_hidden: torch.Tensor, + previous_hidden: torch.Tensor, + history_cache: dict[str, torch.Tensor], + cross_cache: dict[str, torch.Tensor], + chunk: int, +) -> tuple[torch.Tensor, torch.Tensor]: + """Run the exact Stage-1 training forward on an online rollout state.""" + + # A Full call at step 0 has already written the current chunk into the + # pipeline cache. Stage-1 Predictor training, however, exposes only clean + # history from earlier chunks. Slice by causal position instead of the + # cache's current local_end_index so Stage-2 sees the same history contract. + history_end = chunk * TOKENS_PER_CHUNK + available_history = int(history_cache["local_end_index"].item()) + if available_history < history_end: + raise RuntimeError( + f"KV history is incomplete for chunk={chunk}: " + f"available={available_history}, required={history_end}" + ) + batch = { + "noisy_latent": noisy_input, + "timestep": timestep, + "anchor_timestep": anchor_timestep, + "anchor_distance": (timestep.float() - anchor_timestep.float()) + .abs() + .mean(dim=1), + "anchor_hidden": anchor_hidden.detach(), + "previous_hidden": previous_hidden.detach(), + "history_k": history_cache["k"][:, :history_end].detach(), + "history_v": history_cache["v"][:, :history_end].detach(), + "cross_k": cross_cache["k"].detach(), + "cross_v": cross_cache["v"].detach(), + "chunk": chunk, + } + return forward_predictor(predictor, batch, teacher, noisy_input.device) + + +def random_exit_rollout( + *, + pipeline: CausalInferencePipeline, + predictor: torch.nn.Module, + predictor_config: dict[str, Any], + conditional: dict[str, torch.Tensor], + exit_step: int, + noise: torch.Tensor, + predictor_grad: bool, +) -> torch.Tensor: + """Generate 21 latent frames without a graph crossing step/chunk boundaries.""" + + if exit_step not in (1, 2, 3): + raise ValueError(f"exit_step must be 1, 2, or 3, got {exit_step}") + reset_caches(pipeline, noise.device, noise.shape[0]) + teacher = pipeline.generator.model + source_layer = int(predictor_config.get("source_layer", 17)) + timesteps = pipeline.denoising_step_list.to(noise.device) + capture = FinalHiddenCapture(teacher) + previous_hidden: list[torch.Tensor | None] | None = None + outputs: list[torch.Tensor] = [] + try: + for chunk in range(NUM_CHUNKS): + token_start = chunk * TOKENS_PER_CHUNK + noisy_input = noise[ + :, + chunk * FRAMES_PER_CHUNK : (chunk + 1) * FRAMES_PER_CHUNK, + ] + current_hidden: list[torch.Tensor | None] = [None] * NUM_STEPS + denoised: torch.Tensor | None = None + timestep: torch.Tensor | None = None + for step in range(exit_step + 1): + timestep = ( + torch.ones( + (noise.shape[0], FRAMES_PER_CHUNK), + dtype=torch.long, + device=noise.device, + ) + * timesteps[step] + ) + use_predictor = chunk > 0 and step > 0 + if use_predictor: + anchor_hidden = current_hidden[step - 1] + previous_step_hidden = ( + None if previous_hidden is None else previous_hidden[step] + ) + if anchor_hidden is None or previous_step_hidden is None: + raise RuntimeError( + f"Missing Predictor feature at chunk={chunk}, step={step}" + ) + anchor_timestep = torch.ones_like(timestep) * timesteps[step - 1] + grad_here = predictor_grad and step == exit_step + grad_context = torch.enable_grad() if grad_here else torch.no_grad() + with grad_context: + hidden, flow = online_predictor_step( + predictor=predictor, + teacher=teacher, + noisy_input=noisy_input, + timestep=timestep, + anchor_timestep=anchor_timestep, + anchor_hidden=anchor_hidden, + previous_hidden=previous_step_hidden, + history_cache=pipeline.kv_cache1[source_layer], + cross_cache=pipeline.crossattn_cache[source_layer], + chunk=chunk, + ) + denoised = 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] = hidden.detach() + else: + with ( + torch.no_grad(), + torch.autocast(device_type="cuda", dtype=torch.bfloat16), + ): + capture.start() + _, denoised = pipeline.generator( + noisy_image_or_video=noisy_input, + conditional_dict=conditional, + timestep=timestep, + kv_cache=pipeline.kv_cache1, + crossattn_cache=pipeline.crossattn_cache, + current_start=token_start, + ) + current_hidden[step] = capture.finish() + if step < exit_step: + assert denoised is not None + flat = denoised.detach().flatten(0, 1) + noisy_input = pipeline.scheduler.add_noise( + flat, + torch.randn_like(flat), + timesteps[step + 1] + * torch.ones( + flat.shape[0], dtype=torch.long, device=noise.device + ), + ).unflatten(0, denoised.shape[:2]) + + assert denoised is not None and timestep is not None + outputs.append(denoised) + with ( + torch.no_grad(), + torch.autocast(device_type="cuda", dtype=torch.bfloat16), + ): + pipeline.generator( + noisy_image_or_video=denoised.detach(), + conditional_dict=conditional, + timestep=torch.ones_like(timestep) * pipeline.args.context_noise, + kv_cache=pipeline.kv_cache1, + crossattn_cache=pipeline.crossattn_cache, + current_start=token_start, + ) + previous_hidden = [ + value.detach() if value is not None else None + for value in current_hidden + ] + finally: + capture.close() + return torch.cat(outputs, dim=1) + + +@torch.no_grad() +def dmd_gradient( + generated: torch.Tensor, + conditional: dict[str, torch.Tensor], + unconditional: dict[str, torch.Tensor], + real_score: torch.nn.Module | None, + fake_score: torch.nn.Module, + scheduler: Any, + guidance: float, + shift: float, +) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + batch, frames = generated.shape[:2] + timestep = shifted_timestep( + torch.randint(0, 1000, (batch, 1), device=generated.device).repeat(1, frames), + shift, + ) + noise = torch.randn_like(generated) + noisy = scheduler.add_noise( + generated.detach().flatten(0, 1), + noise.flatten(0, 1), + timestep.flatten(0, 1), + ).unflatten(0, (batch, frames)) + with torch.autocast("cuda", dtype=torch.bfloat16): + _, fake_x0 = fake_score(noisy, conditional, timestep) + if real_score is None: + real_x0 = generated.detach() + else: + _, real_cond = real_score(noisy, conditional, timestep) + _, real_uncond = real_score(noisy, unconditional, timestep) + real_x0 = real_cond + (real_cond - real_uncond) * guidance + normalizer = ( + (generated.detach() - real_x0) + .abs() + .mean((1, 2, 3, 4), keepdim=True) + .clamp_min(1e-6) + ) + gradient = torch.nan_to_num((fake_x0 - real_x0) / normalizer) + return gradient, timestep, gradient.abs().mean() + + +def dmd_loss(generated: torch.Tensor, gradient: torch.Tensor) -> torch.Tensor: + prediction = generated[:, FRAMES_PER_CHUNK:].double() + target = ( + generated[:, FRAMES_PER_CHUNK:].double() + - gradient[:, FRAMES_PER_CHUNK:].double() + ).detach() + return 0.5 * F.mse_loss(prediction, target) + + +def fake_score_loss( + generated: torch.Tensor, + conditional: dict[str, torch.Tensor], + fake_score: torch.nn.Module, + scheduler: Any, + shift: float, +) -> tuple[torch.Tensor, torch.Tensor]: + generated = generated.detach() + batch, frames = generated.shape[:2] + timestep = shifted_timestep( + torch.randint(0, 1000, (batch, 1), device=generated.device).repeat(1, frames), + shift, + ) + noise = torch.randn_like(generated) + noisy = scheduler.add_noise( + generated.flatten(0, 1), + noise.flatten(0, 1), + timestep.flatten(0, 1), + ).unflatten(0, (batch, frames)) + with torch.autocast("cuda", dtype=torch.bfloat16): + _, predicted_x0 = fake_score(noisy, conditional, timestep) + predicted_flow = WanDiffusionWrapper._convert_x0_to_flow_pred( + scheduler, + predicted_x0.flatten(0, 1), + noisy.flatten(0, 1), + timestep.flatten(0, 1), + ) + target_flow = noise.flatten(0, 1) - generated.flatten(0, 1) + return ( + F.mse_loss(predicted_flow.float(), target_flow.float()), + timestep, + ) + + +class PredictorEMA: + def __init__(self, module: torch.nn.Module, decay: float) -> None: + self.decay = decay + self.shadow = { + key: value.detach().cpu().float().clone() + for key, value in module.state_dict().items() + } + self.started = False + + @torch.no_grad() + def update(self, module: torch.nn.Module) -> None: + if not self.started: + self.shadow = { + key: value.detach().cpu().float().clone() + for key, value in module.state_dict().items() + } + self.started = True + return + for name, value in module.state_dict().items(): + self.shadow[name].mul_(self.decay).add_( + value.detach().cpu().float(), alpha=1.0 - self.decay + ) + + def load_state_dict(self, state: dict[str, torch.Tensor]) -> None: + if set(state) != set(self.shadow): + missing = sorted(set(self.shadow) - set(state)) + unexpected = sorted(set(state) - set(self.shadow)) + raise RuntimeError( + f"EMA state mismatch: missing={missing}, unexpected={unexpected}" + ) + self.shadow = { + name: value.detach().cpu().float().clone() for name, value in state.items() + } + self.started = True + + +def cpu_state(module: torch.nn.Module) -> dict[str, torch.Tensor]: + return { + key: value.detach().cpu().contiguous() + for key, value in module.state_dict().items() + } + + +def optimizer_to(optimizer: torch.optim.Optimizer, device: torch.device) -> None: + for state in optimizer.state.values(): + for key, value in state.items(): + if torch.is_tensor(value): + state[key] = value.to(device=device) + + +def truncate_log_after(path: Path, step: int) -> None: + if not path.exists(): + return + retained = [] + for line in path.read_text(encoding="utf-8").splitlines(): + if not line.strip(): + continue + record = json.loads(line) + if int(record["student_step"]) <= step: + retained.append(json.dumps(record, ensure_ascii=False)) + temporary = path.with_suffix(path.suffix + ".tmp") + temporary.write_text( + "\n".join(retained) + ("\n" if retained else ""), encoding="utf-8" + ) + os.replace(temporary, path) + + +def predictor_metadata( + *, step: int, predictor_config: dict[str, Any], ema_decay: float | None = None +) -> dict[str, str]: + metadata = { + "step": str(step), + "stage": "random-exit DMD", + "predictor_config": json.dumps( + predictor_config, sort_keys=True, separators=(",", ":") + ), + "confidence_post_training": "compatible", + } + if ema_decay is not None: + metadata["ema_decay"] = str(ema_decay) + return metadata + + +def save_predictor( + output: Path, + predictor: DDP, + ema: PredictorEMA, + predictor_config: dict[str, Any], + step: int, +) -> None: + directory = output / f"checkpoint_step_{step}" + directory.mkdir(parents=True, exist_ok=True) + save_file( + cpu_state(predictor.module), + str(directory / "predictor.safetensors"), + metadata=predictor_metadata(step=step, predictor_config=predictor_config), + ) + save_file( + {key: value.contiguous() for key, value in ema.shadow.items()}, + str(directory / "predictor_ema.safetensors"), + metadata=predictor_metadata( + step=step, + predictor_config=predictor_config, + ema_decay=ema.decay, + ), + ) + + +def save_training_state( + output: Path, + predictor: DDP, + fake_score: DDP, + predictor_opt: torch.optim.Optimizer, + critic_opt: torch.optim.Optimizer, + ema: PredictorEMA, + predictor_config: dict[str, Any], + step: int, +) -> None: + temporary = output / "training_latest.pt.tmp" + torch.save( + { + "student_step": step, + "predictor": cpu_state(predictor.module), + "fake_score": cpu_state(fake_score.module), + "predictor_optimizer": predictor_opt.state_dict(), + "critic_optimizer": critic_opt.state_dict(), + "predictor_ema": ema.shadow, + "predictor_config": predictor_config, + }, + temporary, + ) + os.replace(temporary, output / "training_latest.pt") + + +def main() -> None: + args = parse_args() + dist.init_process_group("nccl") + rank, world_size = dist.get_rank(), dist.get_world_size() + local_rank = int(os.environ["LOCAL_RANK"]) + if world_size != args.expected_world_size and args.smoke_steps is None: + raise ValueError( + f"Formal setup requires {args.expected_world_size} GPUs, got {world_size}" + ) + torch.cuda.set_device(local_rank) + device = torch.device("cuda", local_rank) + torch.backends.cuda.matmul.allow_tf32 = True + torch.backends.cudnn.allow_tf32 = True + torch.set_float32_matmul_precision("high") + + for key in ( + "config_path", + "generator_ckpt", + "predictor_init", + "prompt_path", + "output_dir", + ): + setattr(args, key, resolve(getattr(args, key))) + if args.resume_from is not None: + args.resume_from = resolve(args.resume_from) + if rank == 0: + args.output_dir.mkdir(parents=True, exist_ok=True) + dist.barrier() + + total_steps = args.smoke_steps or args.student_steps + global_batch_size = world_size * args.gradient_accumulation_steps + set_seed(args.seed + rank) + random.seed(args.seed + rank) + prompts = load_prompt_pool( + args.prompt_path, + prompt_count=args.prompt_count, + prompt_seed=args.prompt_seed, + ) + + config = OmegaConf.merge( + OmegaConf.load(ROOT / "configs/default_config.yaml"), + OmegaConf.load(args.config_path), + ) + if ( + list(config.denoising_step_list) != [1000, 750, 500, 250] + or int(config.num_frame_per_block) != FRAMES_PER_CHUNK + ): + raise ValueError( + "Expected the four-step, three-latent-frame Self-Forcing schedule" + ) + predictor_config = read_predictor_config( + args.predictor_init, args.predictor_input_variant + ) + manifest = { + **{ + key: str(value) if isinstance(value, Path) else value + for key, value in vars(args).items() + }, + "world_size": world_size, + "global_batch_size": global_batch_size, + "prompt_sampling": { + "pool_size": len(prompts), + "seed": args.prompt_seed, + "order": "shared seeded permutation without replacement", + "repeat_policy": "cycle only after consuming the complete pool", + }, + "schedule": ( + "one random exit e in {1,2,3}; chunk0 Full-to-e; " + "chunks1..6 F then Predictor-to-e" + ), + "gradient_scope": ( + "only P_e in chunks1..6; no BPTT; detached clean KV history" + ), + "score_models": ( + f"real {args.real_score_name} CFG={args.guidance_scale}; " + f"fake {args.fake_score_name} CFG=0" + ), + "predictor_config": predictor_config, + "confidence_post_training": { + "trainer": "scripts/train_confidence_token_lazy_ddp.py", + "recommended_predictor": "checkpoint_step_*/predictor_ema.safetensors", + }, + } + if rank == 0: + atomic_json(args.output_dir / "config.json", manifest) + dist.barrier() + + print(f"[rank {rank}] loading frozen Self-Forcing generator", flush=True) + pipeline = build_pipeline(config, args.generator_ckpt, device) + predictor_module = build_predictor( + pipeline.generator.model, args.predictor_init, predictor_config, device + ) + predictor = DDP( + predictor_module, + device_ids=[local_rank], + output_device=local_rank, + broadcast_buffers=False, + find_unused_parameters=False, + ) + + print(f"[rank {rank}] loading text encoder", flush=True) + text_encoder = ( + WanTextEncoder() + .to(device=device, dtype=torch.bfloat16) + .eval() + .requires_grad_(False) + ) + real_score = None + if not args.skip_real_score: + print( + f"[rank {rank}] loading and sharding real score {args.real_score_name}", + flush=True, + ) + real_module = WanDiffusionWrapper( + model_name=args.real_score_name, + is_causal=False, + timestep_shift=args.timestep_shift, + ) + real_module.to(dtype=torch.bfloat16).eval().requires_grad_(False) + real_score = fsdp_wrap( + real_module, + sharding_strategy="full", + mixed_precision=True, + wrap_strategy="size", + ) + elif args.smoke_steps is None: + raise ValueError("--skip_real_score is forbidden for formal training") + + print( + f"[rank {rank}] loading trainable fake score {args.fake_score_name}", + flush=True, + ) + fake_module = WanDiffusionWrapper( + model_name=args.fake_score_name, + is_causal=False, + timestep_shift=args.timestep_shift, + ) + fake_module.enable_gradient_checkpointing() + fake_module.to(device=device).train().requires_grad_(True) + fake_score = DDP( + fake_module, + device_ids=[local_rank], + output_device=local_rank, + broadcast_buffers=False, + find_unused_parameters=False, + gradient_as_bucket_view=True, + ) + + fusion_parameters = [ + parameter + for parameter in predictor.module.stage1_input_parameters() + if parameter.requires_grad + ] + block_parameters = [ + parameter + for parameter in predictor.module.block_parameters() + if parameter.requires_grad + ] + predictor_opt = torch.optim.AdamW( + [ + {"params": fusion_parameters, "lr": args.fusion_lr}, + {"params": block_parameters, "lr": args.block_lr}, + ], + betas=(0.0, 0.999), + weight_decay=args.weight_decay, + ) + critic_opt = torch.optim.AdamW( + fake_score.parameters(), + lr=args.critic_lr, + betas=(0.0, 0.999), + weight_decay=args.weight_decay, + ) + ema = PredictorEMA(predictor.module, args.ema_decay) + + start_step = 0 + if args.resume_from is not None: + if args.smoke_steps is not None: + raise ValueError("Resume is not supported with --smoke_steps") + print( + f"[rank {rank}] restoring complete state from {args.resume_from}", + flush=True, + ) + state = torch.load( + args.resume_from, + map_location="cpu", + weights_only=False, + mmap=True, + ) + if state.get("predictor_config") != predictor_config: + raise RuntimeError("Resume state Predictor architecture differs") + start_step = int(state["student_step"]) + if not 0 < start_step < total_steps: + raise ValueError(f"Resume step {start_step} must be in (0, {total_steps})") + predictor.module.load_state_dict(state["predictor"], strict=True) + fake_score.module.load_state_dict(state["fake_score"], strict=True) + predictor_opt.load_state_dict(state["predictor_optimizer"]) + critic_opt.load_state_dict(state["critic_optimizer"]) + optimizer_to(predictor_opt, device) + optimizer_to(critic_opt, device) + ema.load_state_dict(state["predictor_ema"]) + del state + print(f"[rank {rank}] resumed at student step {start_step}", flush=True) + + log_path = args.output_dir / "train_log.jsonl" + if rank == 0 and start_step: + truncate_log_after(log_path, start_step) + dist.barrier() + + negative_prompt = str(config.negative_prompt) + with torch.no_grad(), torch.autocast("cuda", dtype=torch.bfloat16): + unconditional = text_encoder([negative_prompt]) + + swanlab_run = None + if rank == 0 and args.smoke_steps is None and not args.disable_swanlab: + import swanlab + + swanlab_run = swanlab.init( + project=args.swanlab_project, + name=args.swanlab_name, + mode="online", + id=args.swanlab_id, + resume="allow" if args.swanlab_id else None, + config=manifest, + ) + dist.barrier() + + train_started = time.perf_counter() + for student_step in range(start_step, total_steps): + step_started = time.perf_counter() + schedule_factor = lr_factor(student_step, args.student_steps, args.warmup_steps) + set_optimizer_lrs( + predictor_opt, + [args.fusion_lr, args.block_lr], + schedule_factor, + ) + set_optimizer_lrs(critic_opt, [args.critic_lr], schedule_factor) + critic_total = torch.zeros(2, dtype=torch.float64, device=device) + critic_grad_total = 0.0 + + for critic_index in range(args.critic_updates_per_student): + critic_opt.zero_grad(set_to_none=True) + for accumulation_index in range(args.gradient_accumulation_steps): + sample_id = prompt_sample_id( + student_step=student_step, + batch_slot=critic_index, + global_batch_size=global_batch_size, + batches_per_step=args.critic_updates_per_student + 1, + accumulation_index=accumulation_index, + world_size=world_size, + rank=rank, + prompt_pool_size=len(prompts), + ) + with torch.no_grad(), torch.autocast("cuda", dtype=torch.bfloat16): + conditional = text_encoder([prompts[sample_id]]) + exit_step = 1 + ( + (student_step + critic_index + accumulation_index + rank) % 3 + ) + noise = torch.randn( + (1, NUM_CHUNKS * FRAMES_PER_CHUNK, *LATENT_SHAPE), + device=device, + dtype=torch.bfloat16, + ) + with torch.no_grad(): + generated = random_exit_rollout( + pipeline=pipeline, + predictor=predictor, + predictor_config=predictor_config, + conditional=conditional, + exit_step=exit_step, + noise=noise, + predictor_grad=False, + ) + sync = ( + fake_score.no_sync() + if accumulation_index + 1 < args.gradient_accumulation_steps + else contextlib.nullcontext() + ) + with sync: + loss, timestep = fake_score_loss( + generated, + conditional, + fake_score, + pipeline.scheduler, + args.timestep_shift, + ) + (loss / args.gradient_accumulation_steps).backward() + critic_total += torch.tensor( + (float(loss.detach()), float(timestep.mean())), + device=device, + dtype=torch.float64, + ) + del conditional, noise, generated, loss + critic_grad_total += float( + torch.nn.utils.clip_grad_norm_( + fake_score.parameters(), args.critic_grad_clip + ) + ) + critic_opt.step() + critic_opt.zero_grad(set_to_none=True) + + sampled_exit = ( + torch.randint(1, 4, (1,), device=device) + if rank == 0 + else torch.zeros(1, dtype=torch.long, device=device) + ) + dist.broadcast(sampled_exit, src=0) + exit_step = int(sampled_exit.item()) + predictor_opt.zero_grad(set_to_none=True) + student_total = torch.zeros(3, dtype=torch.float64, device=device) + for accumulation_index in range(args.gradient_accumulation_steps): + sample_id = prompt_sample_id( + student_step=student_step, + batch_slot=args.critic_updates_per_student, + global_batch_size=global_batch_size, + batches_per_step=args.critic_updates_per_student + 1, + accumulation_index=accumulation_index, + world_size=world_size, + rank=rank, + prompt_pool_size=len(prompts), + ) + with torch.no_grad(), torch.autocast("cuda", dtype=torch.bfloat16): + conditional = text_encoder([prompts[sample_id]]) + noise = torch.randn( + (1, NUM_CHUNKS * FRAMES_PER_CHUNK, *LATENT_SHAPE), + device=device, + dtype=torch.bfloat16, + ) + sync = ( + predictor.no_sync() + if accumulation_index + 1 < args.gradient_accumulation_steps + else contextlib.nullcontext() + ) + with sync: + generated = random_exit_rollout( + pipeline=pipeline, + predictor=predictor, + predictor_config=predictor_config, + conditional=conditional, + exit_step=exit_step, + noise=noise, + predictor_grad=True, + ) + gradient, score_timestep, gradient_abs = dmd_gradient( + generated, + conditional, + unconditional, + real_score, + fake_score.module, + pipeline.scheduler, + args.guidance_scale, + args.timestep_shift, + ) + loss = dmd_loss(generated, gradient) + (loss / args.gradient_accumulation_steps).backward() + student_total += torch.tensor( + ( + float(loss.detach()), + float(score_timestep.mean()), + float(gradient_abs), + ), + device=device, + dtype=torch.float64, + ) + del conditional, noise, generated, gradient, loss + + predictor_grad_norm = float( + torch.nn.utils.clip_grad_norm_( + predictor.parameters(), args.predictor_grad_clip + ) + ) + predictor_opt.step() + completed = student_step + 1 + if completed >= args.ema_start: + ema.update(predictor.module) + + critic_total /= ( + args.critic_updates_per_student * args.gradient_accumulation_steps + ) + student_total /= args.gradient_accumulation_steps + dist.all_reduce(critic_total) + dist.all_reduce(student_total) + critic_total /= world_size + student_total /= world_size + + if completed == 1 or completed % args.log_every == 0: + torch.cuda.synchronize() + if rank == 0: + record = { + "student_step": completed, + "train/dmd_loss": float(student_total[0]), + "train/dmd_gradient_abs": float(student_total[2]), + "train/fake_score_loss": float(critic_total[0]), + "train/exit_step": exit_step, + "train/predictor_grad_norm": predictor_grad_norm, + "train/fake_score_grad_norm": critic_grad_total + / args.critic_updates_per_student, + "train/dmd_timestep": float(student_total[1]), + "train/fake_score_timestep": float(critic_total[1]), + "lr/fusion": predictor_opt.param_groups[0]["lr"], + "lr/block": predictor_opt.param_groups[1]["lr"], + "lr/fake_score": critic_opt.param_groups[0]["lr"], + "perf/student_step_s": time.perf_counter() - step_started, + "perf/peak_gpu_gib_rank0": torch.cuda.max_memory_allocated() + / 1024**3, + } + append_jsonl(log_path, record) + if swanlab_run is not None: + import swanlab + + swanlab.log(record, step=completed) + print( + f"[student] {completed}/{total_steps} exit=P{exit_step} " + f"dmd={record['train/dmd_loss']:.6f} " + f"fake={record['train/fake_score_loss']:.6f} " + f"time={record['perf/student_step_s']:.1f}s " + f"mem={record['perf/peak_gpu_gib_rank0']:.1f}G", + flush=True, + ) + + if args.smoke_steps is None and completed % args.save_every == 0: + dist.barrier() + if rank == 0: + save_predictor( + args.output_dir, + predictor, + ema, + predictor_config, + completed, + ) + dist.barrier() + if args.smoke_steps is None and completed % args.state_save_every == 0: + dist.barrier() + if rank == 0: + save_training_state( + args.output_dir, + predictor, + fake_score, + predictor_opt, + critic_opt, + ema, + predictor_config, + completed, + ) + dist.barrier() + + dist.barrier() + if rank == 0 and args.smoke_steps is None: + save_predictor( + args.output_dir, + predictor, + ema, + predictor_config, + total_steps, + ) + atomic_json( + args.output_dir / "metrics.json", + { + "status": "complete", + "student_steps": total_steps, + "fake_score_updates": total_steps * args.critic_updates_per_student, + "elapsed_s": time.perf_counter() - train_started, + "selected_predictor": str( + args.output_dir + / f"checkpoint_step_{total_steps}" + / "predictor_ema.safetensors" + ), + "confidence_post_training": ( + "Run scripts/train_confidence_token_lazy_ddp.py with " + "--predictor_weights set to selected_predictor" + ), + }, + ) + dist.barrier() + if rank == 0 and swanlab_run is not None: + import swanlab + + swanlab.finish() + dist.destroy_process_group() + + +if __name__ == "__main__": + main() diff --git a/scripts/train_layer17_stage1_lazy_ddp.py b/scripts/train_layer17_stage1_lazy_ddp.py new file mode 100644 index 0000000000000000000000000000000000000000..78626607c1153ea72cb94ac89601f0ec5e3bb898 --- /dev/null +++ b/scripts/train_layer17_stage1_lazy_ddp.py @@ -0,0 +1,555 @@ +#!/usr/bin/env python3 +"""Train a fresh Self-Forcing Layer-17 Predictor with lazy DDP input.""" + +from __future__ import annotations + +import argparse +import hashlib +import json +import os +import random +import sys +import time +from contextlib import nullcontext +from pathlib import Path +from typing import Any + +import torch +import torch.distributed as dist +import torch.nn.functional as F +from safetensors.torch import save_file +from torch.nn.parallel import DistributedDataParallel as DDP +from torch.optim import AdamW +from torch.utils.data import DataLoader + +ROOT = Path(__file__).resolve().parents[1] +if str(ROOT) not in sys.path: + sys.path.insert(0, str(ROOT)) + +from predictor_training.lazy_offline_data import ( + LazyLayer17Dataset, + ScheduledBatchSampler, + collate_lazy_samples, +) +from scripts.run_single_block_init_sweep import ( + atomic_json, + build_shared_nonblock_state, + forward_predictor, + gradient_norm, + load_teacher, + lr_values, + make_model, +) +from utils.misc import set_seed +from wan.modules.causal_model import causal_rope_apply + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--dataset_root", type=Path, required=True) + parser.add_argument("--output_dir", type=Path, required=True) + parser.add_argument( + "--checkpoint_path", + type=Path, + default=Path("checkpoints/self_forcing_dmd.pt"), + ) + parser.add_argument( + "--config_path", + type=Path, + default=Path("configs/self_forcing_sid.yaml"), + ) + parser.add_argument("--max_steps", type=int, default=2000) + parser.add_argument("--per_device_batch_size", type=int, default=16) + parser.add_argument("--gradient_accumulation_steps", type=int, default=1) + parser.add_argument("--num_workers", type=int, default=0) + parser.add_argument("--save_every", type=int, default=100) + parser.add_argument("--log_every", type=int, default=20) + parser.add_argument("--seed", type=int, default=0) + parser.add_argument("--fusion_lr", type=float, default=1e-4) + parser.add_argument("--block_lr", type=float, default=1e-5) + parser.add_argument("--weight_decay", type=float, default=0.01) + parser.add_argument("--hidden_weight", type=float, default=0.1) + parser.add_argument("--flow_weight", type=float, default=1.0) + parser.add_argument("--grad_clip", type=float, default=1.0) + parser.add_argument("--fusion_warmup_steps", type=int, default=100) + parser.add_argument("--block_freeze_steps", type=int, default=100) + parser.add_argument("--block_warmup_steps", type=int, default=100) + parser.add_argument( + "--input_variant", + choices=("self_forcing", "disca", "atc"), + default="self_forcing", + help=( + "disca keeps the same Predictor training recipe but removes the " + "previous-chunk hidden channel from input fusion" + ), + ) + parser.add_argument( + "--atc_previous_scope", + choices=("chunk", "last_frame"), + default="chunk", + ) + parser.add_argument("--atc_freq_dim", type=int, default=256) + parser.add_argument("--atc_mlp_hidden_dim", type=int, default=3072) + parser.add_argument("--atc_gate_hidden_dim", type=int, default=512) + parser.add_argument("--atc_transport_residual_scale", type=float, default=0.1) + parser.add_argument("--atc_gate_initial_probability", type=float, default=0.3) + args = parser.parse_args() + if args.max_steps < 1 or args.per_device_batch_size < 1: + parser.error("steps and batch size must be positive") + if args.gradient_accumulation_steps < 1 or args.save_every < 1: + parser.error("accumulation and save interval must be positive") + return args + + +def resolve(path: Path) -> Path: + path = path.expanduser() + return path.resolve() if path.is_absolute() else (ROOT / path).resolve() + + +def append_jsonl(path: Path, value: dict[str, Any]) -> None: + with path.open("a", encoding="utf-8") as handle: + handle.write(json.dumps(value, sort_keys=True) + "\n") + + +def build_rank_schedule( + *, + rank: int, + world_size: int, + batch_size: int, + accumulation: int, + steps: int, + seed: int, + start_step: int = 0, +) -> tuple[list[list[tuple[int, int, int]]], str]: + rng = random.Random(seed) + groups = [ + (chunk, target) + for chunk in range(1, 7) + for target in range(1, 4) + ] + all_steps: list[tuple[int, int, list[int]]] = [] + while len(all_steps) < steps: + shuffled = groups.copy() + rng.shuffle(shuffled) + for chunk, target in shuffled: + count = world_size * batch_size * accumulation + prompts = rng.sample(range(1000), count) + all_steps.append((chunk, target, prompts)) + if len(all_steps) == steps: + break + fingerprint = hashlib.sha256( + json.dumps(all_steps, separators=(",", ":")).encode() + ).hexdigest() + batches = [] + for chunk, target, prompts in all_steps[start_step:]: + for micro in range(accumulation): + offset = (micro * world_size + rank) * batch_size + local = prompts[offset : offset + batch_size] + batches.append([(prompt, chunk, target) for prompt in local]) + return batches, fingerprint + + +@torch.inference_mode() +def project_history( + prefeature: torch.Tensor, + chunk: int, + teacher: torch.nn.Module, + device: torch.device, +) -> tuple[torch.Tensor, torch.Tensor]: + value_in = prefeature.to(device=device, dtype=torch.bfloat16) + block = teacher.blocks[17] + batch, sequence, _ = value_in.shape + heads = block.num_heads + head_dim = block.dim // heads + with torch.autocast(device_type="cuda", dtype=torch.bfloat16): + key = block.self_attn.norm_k(block.self_attn.k(value_in)).view( + batch, sequence, heads, head_dim + ) + value = block.self_attn.v(value_in).view( + batch, sequence, heads, head_dim + ) + grid_sizes = torch.tensor( + [[chunk * 3, 30, 52]] * batch, + dtype=torch.long, + device="cpu", + ) + key = causal_rope_apply(key, grid_sizes, teacher.freqs, start_frame=0) + return key, value + + +def move_training_batch( + batch: dict[str, Any], + teacher: torch.nn.Module, + device: torch.device, +) -> dict[str, Any]: + history_k, history_v = project_history( + batch.pop("history_prefeature"), batch["chunk"], teacher, device + ) + output = { + key: value.to( + device=device, + dtype=( + torch.bfloat16 + if value.is_floating_point() and key != "timestep" + else value.dtype + ), + ) + for key, value in batch.items() + if isinstance(value, torch.Tensor) + } + output.update( + prompt_ids=batch["prompt_ids"], + chunk=batch["chunk"], + anchor_step=batch["anchor_step"], + target_step=batch["target_step"], + history_k=history_k, + history_v=history_v, + ) + return output + + +def save_checkpoint( + *, + model: torch.nn.Module, + optimizer: AdamW, + step: int, + output_dir: Path, + schedule_sha256: str, + predictor_config: dict[str, Any], +) -> None: + destination = output_dir / f"checkpoint_step_{step:04d}" + destination.mkdir(parents=True, exist_ok=True) + weights = { + key: value.detach().cpu().contiguous() + for key, value in model.state_dict().items() + } + weights_tmp = destination / "predictor.safetensors.tmp" + save_file( + weights, + weights_tmp, + metadata={ + "predictor_config": json.dumps( + predictor_config, sort_keys=True, separators=(",", ":") + ) + }, + ) + os.replace(weights_tmp, destination / "predictor.safetensors") + state_tmp = destination / "training_state.pt.tmp" + torch.save( + { + "step": step, + "model": weights, + "optimizer": optimizer.state_dict(), + "schedule_sha256": schedule_sha256, + "predictor_config": predictor_config, + }, + state_tmp, + ) + os.replace(state_tmp, destination / "training_state.pt") + atomic_json( + output_dir / "latest_checkpoint.json", + { + "step": step, + "path": str(destination), + "schedule_sha256": schedule_sha256, + }, + ) + + +def main() -> None: + args = parse_args() + dist.init_process_group("nccl") + rank = dist.get_rank() + world_size = dist.get_world_size() + local_rank = int(os.environ["LOCAL_RANK"]) + torch.cuda.set_device(local_rank) + device = torch.device("cuda", local_rank) + is_main = rank == 0 + + args.dataset_root = resolve(args.dataset_root) + args.output_dir = resolve(args.output_dir) + args.checkpoint_path = resolve(args.checkpoint_path) + args.config_path = resolve(args.config_path) + if is_main: + args.output_dir.mkdir(parents=True, exist_ok=True) + dist.barrier() + set_seed(args.seed) + torch.set_num_threads(4) + torch.set_num_interop_threads(1) + torch.backends.cuda.matmul.allow_tf32 = True + torch.set_float32_matmul_precision("high") + + predictor_config = { + "source_layer": 17, + "input_variant": args.input_variant, + "atc_previous_scope": args.atc_previous_scope, + "atc_freq_dim": args.atc_freq_dim, + "atc_mlp_hidden_dim": args.atc_mlp_hidden_dim, + "atc_gate_hidden_dim": args.atc_gate_hidden_dim, + "atc_transport_residual_scale": args.atc_transport_residual_scale, + "atc_gate_initial_probability": args.atc_gate_initial_probability, + } + + start_step = 0 + latest_metadata = args.output_dir / "latest_checkpoint.json" + if latest_metadata.exists(): + start_step = int(json.loads(latest_metadata.read_text())["step"]) + batches, schedule_sha256 = build_rank_schedule( + rank=rank, + world_size=world_size, + batch_size=args.per_device_batch_size, + accumulation=args.gradient_accumulation_steps, + steps=args.max_steps, + seed=args.seed, + start_step=start_step, + ) + if is_main: + atomic_json( + args.output_dir / "config.json", + { + **{ + key: str(value) if isinstance(value, Path) else value + for key, value in vars(args).items() + }, + "world_size": world_size, + "effective_global_batch_size": ( + world_size + * args.per_device_batch_size + * args.gradient_accumulation_steps + ), + "source_layer": 17, + "initialization": "fresh Self-Forcing Teacher Layer 17", + "predictor_config": predictor_config, + "input_variant": args.input_variant, + "input_features": ( + "current_tokens + same_chunk_previous_timestep_hidden + " + "current_timestep; no previous_chunk_hidden channel" + if args.input_variant == "disca" + else ( + "ATC(current_tokens, same_chunk_previous_timestep_hidden, " + "previous_chunk_same_timestep_hidden, target_time_condition, " + "anchor_distance)" + if args.input_variant == "atc" + else "current_tokens + same_chunk_previous_timestep_hidden + " + "previous_chunk_same_timestep_hidden + current_timestep" + ) + ), + "training_prompt_ids": "0..999", + "validation": None, + "selection": f"checkpoint_step_{args.max_steps:04d}", + "schedule_sha256": schedule_sha256, + }, + ) + + print(f"[rank {rank}] loading frozen Self-Forcing Teacher", flush=True) + teacher = load_teacher(args.checkpoint_path, args.config_path, device) + shared = build_shared_nonblock_state(args.seed) + model = make_model( + teacher, + 17, + "teacher_full", + shared, + args.seed, + True, + device, + input_variant=args.input_variant, + atc_previous_scope=args.atc_previous_scope, + atc_freq_dim=args.atc_freq_dim, + atc_mlp_hidden_dim=args.atc_mlp_hidden_dim, + atc_gate_hidden_dim=args.atc_gate_hidden_dim, + atc_transport_residual_scale=args.atc_transport_residual_scale, + atc_gate_initial_probability=args.atc_gate_initial_probability, + ) + model.set_block_trainable(True) + fusion_parameters = model.stage1_input_parameters() + block_parameters = model.block_parameters() + optimizer = AdamW( + [ + {"params": fusion_parameters, "lr": args.fusion_lr}, + {"params": block_parameters, "lr": args.block_lr}, + ], + betas=(0.9, 0.95), + weight_decay=args.weight_decay, + ) + ddp = DDP( + model, + device_ids=[local_rank], + output_device=local_rank, + broadcast_buffers=False, + find_unused_parameters=True, + ) + if start_step: + state_path = ( + args.output_dir + / f"checkpoint_step_{start_step:04d}" + / "training_state.pt" + ) + state = torch.load(state_path, map_location="cpu", weights_only=False) + if state["schedule_sha256"] != schedule_sha256: + raise RuntimeError("Checkpoint schedule differs from requested run") + if state.get("predictor_config", predictor_config) != predictor_config: + raise RuntimeError("Checkpoint Predictor architecture differs") + ddp.module.load_state_dict(state["model"], strict=True) + optimizer.load_state_dict(state["optimizer"]) + print(f"[rank {rank}] resumed at step {start_step}", flush=True) + + dataset = LazyLayer17Dataset( + args.dataset_root, + layer_id=17, + include_previous_hidden=args.input_variant != "disca", + ) + loader = DataLoader( + dataset, + batch_sampler=ScheduledBatchSampler(batches), + num_workers=args.num_workers, + collate_fn=collate_lazy_samples, + pin_memory=False, + persistent_workers=args.num_workers > 0, + prefetch_factor=1 if args.num_workers > 0 else None, + ) + iterator = iter(loader) + optimizer.zero_grad(set_to_none=True) + started = time.perf_counter() + log_path = args.output_dir / "train_log.jsonl" + + for step in range(start_step, args.max_steps): + fusion_lr, block_lr = lr_values( + step, + args.max_steps, + args.fusion_lr, + args.block_lr, + args.fusion_warmup_steps, + args.block_freeze_steps, + args.block_warmup_steps, + ) + optimizer.param_groups[0]["lr"] = fusion_lr + optimizer.param_groups[1]["lr"] = block_lr + step_started = time.perf_counter() + totals = torch.zeros(3, device=device, dtype=torch.float64) + atc_diagnostic_totals: dict[str, torch.Tensor] = {} + last_chunk = last_target = -1 + for micro in range(args.gradient_accumulation_steps): + cpu_batch = next(iterator) + batch = move_training_batch(cpu_batch, teacher, device) + last_chunk, last_target = batch["chunk"], batch["target_step"] + sync = ( + ddp.no_sync() + if micro + 1 < args.gradient_accumulation_steps + else nullcontext() + ) + with sync: + pred_hidden, pred_flow = forward_predictor( + ddp, batch, teacher, device + ) + if args.input_variant == "atc": + for key, value in ddp.module.last_atc_diagnostics.items(): + contribution = ( + value.detach().double() + / args.gradient_accumulation_steps + ) + if key in atc_diagnostic_totals: + atc_diagnostic_totals[key] += contribution + else: + atc_diagnostic_totals[key] = contribution.clone() + hidden_loss = F.mse_loss( + pred_hidden.float(), batch["target_hidden"].float() + ) + flow_loss = F.mse_loss( + pred_flow.float(), batch["target_flow"].float() + ) + loss = ( + args.hidden_weight * hidden_loss + + args.flow_weight * flow_loss + ) + (loss / args.gradient_accumulation_steps).backward() + totals += torch.stack( + [loss.detach(), hidden_loss.detach(), flow_loss.detach()] + ).double() / args.gradient_accumulation_steps + del cpu_batch, batch, pred_hidden, pred_flow + del hidden_loss, flow_loss, loss + + dist.all_reduce(totals) + totals /= world_size + for value in atc_diagnostic_totals.values(): + dist.all_reduce(value) + value /= world_size + if step < args.block_freeze_steps: + for parameter in block_parameters: + parameter.grad = None + fusion_grad = gradient_norm(fusion_parameters) + block_grad = gradient_norm(block_parameters) + total_grad = torch.nn.utils.clip_grad_norm_( + ddp.module.parameters(), args.grad_clip + ) + optimizer.step() + optimizer.zero_grad(set_to_none=True) + completed = step + 1 + + if completed == 1 or completed % args.log_every == 0: + torch.cuda.synchronize() + if is_main: + record = { + "step": completed, + "train_total_loss": float(totals[0]), + "train_hidden_mse": float(totals[1]), + "train_flow_mse": float(totals[2]), + "fusion_grad_norm": fusion_grad, + "block_grad_norm": block_grad, + "total_grad_norm_before_clip": float(total_grad), + "fusion_lr": fusion_lr, + "block_lr": block_lr, + "chunk": last_chunk, + "target_step": last_target, + "step_time_s": time.perf_counter() - step_started, + "peak_gpu_gib": torch.cuda.max_memory_allocated() / 2**30, + } + record.update( + { + f"atc_{key}": float(value) + for key, value in atc_diagnostic_totals.items() + } + ) + append_jsonl(log_path, record) + print( + f"[train] {completed}/{args.max_steps} " + f"flow={record['train_flow_mse']:.8f} " + f"time={record['step_time_s']:.2f}s " + f"mem={record['peak_gpu_gib']:.1f}G", + flush=True, + ) + + if completed % args.save_every == 0: + dist.barrier() + if is_main: + save_checkpoint( + model=ddp.module, + optimizer=optimizer, + step=completed, + output_dir=args.output_dir, + schedule_sha256=schedule_sha256, + predictor_config=predictor_config, + ) + print(f"[checkpoint] step={completed}", flush=True) + dist.barrier() + del totals, total_grad + + if is_main: + atomic_json( + args.output_dir / "selected_checkpoint.json", + { + "selection_rule": "final requested optimizer step", + "step": args.max_steps, + "path": str( + args.output_dir + / f"checkpoint_step_{args.max_steps:04d}" + ), + "training_time_s": time.perf_counter() - started, + }, + ) + print(f"[complete] selected step={args.max_steps}", flush=True) + dist.barrier() + dist.destroy_process_group() + + +if __name__ == "__main__": + main() diff --git a/scripts/train_long_layer17_ddp.py b/scripts/train_long_layer17_ddp.py new file mode 100644 index 0000000000000000000000000000000000000000..d97b908525d48f64fe2e532d009e23ceec4802ad --- /dev/null +++ b/scripts/train_long_layer17_ddp.py @@ -0,0 +1,423 @@ +#!/usr/bin/env python3 +"""Train the Layer-17 predictor on 2x/4x offline trajectories with 2-GPU DDP.""" + +from __future__ import annotations + +import argparse +import hashlib +import json +import os +import random +import sys +import time +from contextlib import nullcontext +from pathlib import Path + +import swanlab +import torch +import torch.distributed as dist +import torch.nn.functional as F +from safetensors.torch import save_file +from torch.nn.parallel import DistributedDataParallel as DDP +from torch.optim import AdamW + +ROOT = Path(__file__).resolve().parents[1] +if str(ROOT) not in sys.path: + sys.path.insert(0, str(ROOT)) + +from predictor_training.offline_data import OfflinePredictorStore +from scripts.run_single_block_init_sweep import ( + append_jsonl, + atomic_json, + build_shared_nonblock_state, + evaluate, + forward_predictor, + gradient_norm, + load_teacher, + lr_values, + make_model, + move_batch, + normalized_auc, + parameter_norm, + resolve, +) +from utils.misc import set_seed + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser() + parser.add_argument("--train_root", type=Path, required=True) + parser.add_argument( + "--val_root", type=Path, + default=Path("outputs/predictor_offline_100_all_blocks"), + ) + parser.add_argument("--output_dir", type=Path, required=True) + parser.add_argument("--num_chunks", type=int, choices=(14, 28), required=True) + parser.add_argument("--max_steps", type=int, default=1000) + parser.add_argument("--per_device_batch_size", type=int, default=32) + parser.add_argument("--gradient_accumulation_steps", type=int, default=2) + parser.add_argument("--eval_every", type=int, default=100) + parser.add_argument("--log_every", type=int, default=20) + parser.add_argument("--save_every", type=int, default=100) + parser.add_argument("--fusion_lr", type=float, default=1e-4) + parser.add_argument("--block_lr", type=float, default=1e-5) + parser.add_argument("--weight_decay", type=float, default=0.01) + parser.add_argument("--hidden_weight", type=float, default=0.1) + parser.add_argument("--flow_weight", type=float, default=1.0) + parser.add_argument( + "--gate_mode", choices=("baseline", "learned", "constant"), + default="baseline", + ) + parser.add_argument("--gate_hidden_dim", type=int, default=128) + parser.add_argument("--gate_initial_bias", type=float, default=4.6) + parser.add_argument("--gate_floor", type=float, default=0.0) + parser.add_argument("--constant_gate", type=float, default=1.0) + parser.add_argument("--grad_clip", type=float, default=1.0) + parser.add_argument("--fusion_warmup_steps", type=int, default=100) + parser.add_argument("--block_freeze_steps", type=int, default=100) + parser.add_argument("--block_warmup_steps", type=int, default=100) + parser.add_argument("--seed", type=int, default=0) + parser.add_argument( + "--checkpoint_path", type=Path, + default=Path("checkpoints/self_forcing_dmd.pt"), + ) + parser.add_argument( + "--config_path", type=Path, + default=Path("configs/self_forcing_sid.yaml"), + ) + parser.add_argument("--swanlab_project", default="Self-Forcing") + parser.add_argument("--swanlab_name", required=True) + args = parser.parse_args() + if not 0.0 <= args.constant_gate <= 1.0: + parser.error("--constant_gate must be in [0, 1]") + return args + + +class DistributedBatchSchedule: + """One denoising group plus an independent batch per rank/micro-step.""" + + def __init__( + self, prompt_ids: list[int], batch_size: int, accumulation_steps: int, + world_size: int, num_chunks: int, steps: int, seed: int, + ) -> None: + self.entries: list[tuple[int, int, list[list[int]]]] = [] + rng = random.Random(seed) + groups = [ + (chunk, target_step) + for chunk in range(1, num_chunks) + for target_step in range(1, 4) + ] + while len(self.entries) < steps: + epoch_groups = groups.copy() + rng.shuffle(epoch_groups) + for chunk, target_step in epoch_groups: + batches = [ + rng.sample(prompt_ids, batch_size) + for _ in range(world_size * accumulation_steps) + ] + self.entries.append((chunk, target_step, batches)) + if len(self.entries) == steps: + break + + def fingerprint(self) -> str: + payload = json.dumps(self.entries, separators=(",", ":")).encode() + return hashlib.sha256(payload).hexdigest() + + +def save_weights(model: torch.nn.Module, path: Path) -> None: + temporary = path.with_suffix(path.suffix + ".tmp") + save_file( + {k: v.detach().cpu().contiguous() for k, v in model.state_dict().items()}, + temporary, + ) + os.replace(temporary, path) + + +def main() -> None: + args = parse_args() + dist.init_process_group("nccl") + rank = dist.get_rank() + world_size = dist.get_world_size() + local_rank = int(os.environ["LOCAL_RANK"]) + torch.cuda.set_device(local_rank) + device = torch.device("cuda", local_rank) + is_main = rank == 0 + + args.train_root = resolve(args.train_root) + args.val_root = resolve(args.val_root) + args.output_dir = resolve(args.output_dir) + if is_main: + args.output_dir.mkdir(parents=True, exist_ok=True) + dist.barrier() + + set_seed(args.seed) + torch.backends.cuda.matmul.allow_tf32 = True + torch.backends.cudnn.allow_tf32 = True + torch.set_float32_matmul_precision("high") + + train_ids = list(range(80)) + val_ids = list(range(80, 100)) + schedule = DistributedBatchSchedule( + train_ids, args.per_device_batch_size, + args.gradient_accumulation_steps, world_size, + args.num_chunks, args.max_steps, args.seed, + ) + micro_batch_size = args.per_device_batch_size + + if is_main: + config = { + **{k: str(v) if isinstance(v, Path) else v for k, v in vars(args).items()}, + "world_size": world_size, + "micro_batch_size_per_gpu": micro_batch_size, + "effective_global_batch_size": ( + micro_batch_size * world_size * args.gradient_accumulation_steps + ), + "source_layer": 17, + "max_history_chunks": 7, + "train_prompt_ids": train_ids, + "val_prompt_ids": val_ids, + "validation_num_chunks": 7, + "batch_schedule_sha256": schedule.fingerprint(), + } + atomic_json(args.output_dir / "config.json", config) + swanlab.init( + project=args.swanlab_project, + experiment_name=args.swanlab_name, + config=config, + logdir=str(args.output_dir / "swanlog"), + ) + + print(f"[rank {rank}] loading Teacher", flush=True) + teacher = load_teacher(args.checkpoint_path, args.config_path, device) + shared_state = build_shared_nonblock_state(args.seed) + model = make_model( + teacher, 17, "teacher_full", shared_state, args.seed, True, device, + args.gate_mode, args.gate_hidden_dim, args.gate_initial_bias, + args.gate_floor, args.constant_gate, + ) + fusion_parameters = model.fusion_parameters() + block_parameters = model.block_parameters() + optimizer = AdamW( + [ + {"params": fusion_parameters, "lr": args.fusion_lr}, + {"params": block_parameters, "lr": args.block_lr}, + ], + betas=(0.9, 0.95), + weight_decay=args.weight_decay, + ) + initial_block_norm = parameter_norm(block_parameters) + + print(f"[rank {rank}] loading {args.num_chunks}-chunk training data", flush=True) + train_store = OfflinePredictorStore( + args.train_root, train_ids, + num_chunks=args.num_chunks, max_history_chunks=7, + ) + train_store.load_layer_cache(17, teacher, device) + + val_store = None + if is_main: + print("[rank 0] loading 7-chunk validation data", flush=True) + val_store = OfflinePredictorStore( + args.val_root, val_ids, num_chunks=7, max_history_chunks=7, + ) + val_store.load_layer_cache(17, teacher, device) + + ddp = DDP( + model, device_ids=[local_rank], output_device=local_rank, + broadcast_buffers=False, find_unused_parameters=False, + ) + dist.barrier() + + evaluations: list[dict] = [] + latest_path = args.output_dir / "training_latest.pt" + start_step = 0 + if latest_path.exists(): + state = torch.load(latest_path, map_location="cpu", weights_only=False) + ddp.module.load_state_dict(state["model"], strict=True) + optimizer.load_state_dict(state["optimizer"]) + evaluations = state["evaluations"] + start_step = int(state["step"]) + print(f"[rank {rank}] resumed at step {start_step}", flush=True) + + if is_main and not evaluations: + initial = evaluate( + ddp.module, val_store, val_ids, 10, teacher, device, + args.hidden_weight, args.flow_weight, + ) + evaluations.append({"step": 0, **initial}) + swanlab.log( + { + f"val/{k}": v for k, v in initial.items() + if isinstance(v, (int, float)) + }, + step=0, + ) + print(f"[eval] step=0 flow={initial['flow_mse']:.8f}", flush=True) + dist.barrier() + + ddp.train() + optimizer.zero_grad(set_to_none=True) + run_started = time.perf_counter() + + for step in range(start_step, args.max_steps): + block_enabled = step >= args.block_freeze_steps + + fusion_lr, block_lr = lr_values( + step, args.max_steps, args.fusion_lr, args.block_lr, + args.fusion_warmup_steps, args.block_freeze_steps, + args.block_warmup_steps, + ) + optimizer.param_groups[0]["lr"] = fusion_lr + optimizer.param_groups[1]["lr"] = block_lr + + chunk, target_step, distributed_prompts = schedule.entries[step] + step_started = time.perf_counter() + totals = torch.zeros(3, device=device, dtype=torch.float64) + for accumulation_index in range(args.gradient_accumulation_steps): + prompt_ids = distributed_prompts[ + rank * args.gradient_accumulation_steps + accumulation_index + ] + batch = move_batch( + train_store.batch(prompt_ids, chunk, target_step), device, + ) + sync_context = ( + ddp.no_sync() + if accumulation_index < args.gradient_accumulation_steps - 1 + else nullcontext() + ) + with sync_context: + pred_hidden, pred_flow = forward_predictor( + ddp, batch, teacher, device, + ) + hidden_loss = F.mse_loss( + pred_hidden.float(), batch["target_hidden"].float(), + ) + flow_loss = F.mse_loss( + pred_flow.float(), batch["target_flow"].float(), + ) + loss = args.hidden_weight * hidden_loss + args.flow_weight * flow_loss + (loss / args.gradient_accumulation_steps).backward() + totals += torch.stack( + [loss.detach(), hidden_loss.detach(), flow_loss.detach()] + ).double() / args.gradient_accumulation_steps + del batch, pred_hidden, pred_flow, hidden_loss, flow_loss, loss + + dist.all_reduce(totals, op=dist.ReduceOp.SUM) + totals /= world_size + # DDP parameters remain registered throughout training. Dropping the + # block gradients exactly preserves the original 100-step freeze and + # prevents AdamW state from advancing for the frozen block. + if not block_enabled: + for parameter in block_parameters: + parameter.grad = None + fusion_grad_norm = gradient_norm(fusion_parameters) + block_grad_norm = gradient_norm(block_parameters) + total_grad_norm = torch.nn.utils.clip_grad_norm_( + ddp.module.parameters(), args.grad_clip, + ) + optimizer.step() + optimizer.zero_grad(set_to_none=True) + completed_step = step + 1 + + if is_main and (completed_step == 1 or completed_step % args.log_every == 0): + torch.cuda.synchronize() + record = { + "step": completed_step, + "train_total_loss": float(totals[0]), + "train_hidden_mse": float(totals[1]), + "train_flow_mse": float(totals[2]), + "fusion_grad_norm": fusion_grad_norm, + "block_grad_norm": block_grad_norm, + "total_grad_norm_before_clip": float(total_grad_norm), + "fusion_lr": fusion_lr, + "block_lr": block_lr, + "chunk": chunk, + "target_step": target_step, + "step_time_s": time.perf_counter() - step_started, + "peak_gpu_gib": torch.cuda.max_memory_allocated() / 2**30, + } + append_jsonl(args.output_dir / "train_log.jsonl", record) + swanlab.log({f"train/{k}": v for k, v in record.items() if k != "step"}, step=completed_step) + print( + f"[train] {completed_step}/{args.max_steps} " + f"flow={record['train_flow_mse']:.8f} " + f"time={record['step_time_s']:.2f}s", + flush=True, + ) + + if completed_step % args.eval_every == 0 or completed_step == args.max_steps: + dist.barrier() + if is_main: + validation = evaluate( + ddp.module, val_store, val_ids, 10, teacher, device, + args.hidden_weight, args.flow_weight, + ) + evaluations.append({"step": completed_step, **validation}) + swanlab.log( + { + f"val/{k}": v for k, v in validation.items() + if isinstance(v, (int, float)) + }, + step=completed_step, + ) + print( + f"[eval] step={completed_step} " + f"flow={validation['flow_mse']:.8f}", flush=True, + ) + dist.barrier() + + if completed_step % args.save_every == 0 or completed_step == args.max_steps: + if is_main: + temporary = latest_path.with_suffix(".pt.tmp") + torch.save( + { + "model": {k: v.detach().cpu() for k, v in ddp.module.state_dict().items()}, + "optimizer": optimizer.state_dict(), + "evaluations": evaluations, + "step": completed_step, + }, + temporary, + ) + os.replace(temporary, latest_path) + dist.barrier() + + if is_main: + final = evaluations[-1] + result = { + "status": "complete", + "name": args.swanlab_name, + "num_chunks": args.num_chunks, + "latent_length": args.num_chunks * 3, + "source_layer": 17, + "max_steps": args.max_steps, + "per_device_batch_size": args.per_device_batch_size, + "effective_global_batch_size": ( + args.per_device_batch_size * world_size + * args.gradient_accumulation_steps + ), + "gradient_accumulation_steps": args.gradient_accumulation_steps, + "world_size": world_size, + "micro_batch_size_per_gpu": micro_batch_size, + "final_val_hidden_mse": final["hidden_mse"], + "final_val_flow_mse": final["flow_mse"], + "final_val_total_loss": final["total_loss"], + "val_hidden_mse_auc": normalized_auc(evaluations, "hidden_mse", args.max_steps), + "val_flow_mse_auc": normalized_auc(evaluations, "flow_mse", args.max_steps), + "evaluations": evaluations, + "training_time_s": time.perf_counter() - run_started, + "initial_block_parameter_norm": initial_block_norm, + "final_block_parameter_norm": parameter_norm(block_parameters), + } + save_weights(ddp.module, args.output_dir / "predictor_final.safetensors") + atomic_json(args.output_dir / "metrics.json", result) + if latest_path.exists(): + latest_path.unlink() + swanlab.log({"final/val_flow_mse": final["flow_mse"]}, step=args.max_steps) + swanlab.finish() + + dist.barrier() + dist.destroy_process_group() + + +if __name__ == "__main__": + main() diff --git a/scripts/vbench8_protocol.py b/scripts/vbench8_protocol.py new file mode 100644 index 0000000000000000000000000000000000000000..364752b510d2d7f47535eacc96a9d2d04fcb2422 --- /dev/null +++ b/scripts/vbench8_protocol.py @@ -0,0 +1,97 @@ +"""Shared constants and pure functions for the VBench-8 protocol.""" + +from __future__ import annotations + +from typing import Mapping + + +PROTOCOL_NAME = "Self-Forcing Extended-251 Full Evaluation" + +DIMENSIONS = ( + "subject_consistency", + "background_consistency", + "motion_smoothness", + "dynamic_degree", + "aesthetic_quality", + "imaging_quality", + "scene", + "overall_consistency", +) + +SUITE_COUNTS = { + "subject_consistency": 72, + "overall_consistency": 93, + "scene": 86, +} + +SUITE_DIMENSIONS = { + "subject_consistency": ( + "subject_consistency", + "dynamic_degree", + "motion_smoothness", + ), + "overall_consistency": ( + "overall_consistency", + "aesthetic_quality", + "imaging_quality", + ), + "scene": ("scene", "background_consistency"), +} + +# These are the empirical VBench Selected Score ranges specified by the +# Self-Forcing Extended-251 Full Evaluation protocol. The installed vbench==0.1.5 +# wheel does not ship the Selected Score constants. +NORMALIZE_RANGE = { + "subject_consistency": (0.1462, 1.0), + "background_consistency": (0.2615, 1.0), + "motion_smoothness": (0.7060, 0.9975), + "dynamic_degree": (0.0, 1.0), + "aesthetic_quality": (0.0, 1.0), + "imaging_quality": (0.0, 1.0), + "scene": (0.0, 0.8222), + "overall_consistency": (0.0, 0.3640), +} + +QUALITY_WEIGHTS = { + "subject_consistency": 1.0, + "background_consistency": 1.0, + "motion_smoothness": 1.0, + "dynamic_degree": 0.5, + "aesthetic_quality": 1.0, + "imaging_quality": 1.0, +} +SEMANTIC_WEIGHTS = {"scene": 1.0, "overall_consistency": 1.0} + + +def normalize_scores(raw_scores: Mapping[str, float]) -> dict[str, float]: + missing = [dimension for dimension in DIMENSIONS if dimension not in raw_scores] + if missing: + raise KeyError(f"Missing raw VBench scores: {missing}") + normalized: dict[str, float] = {} + for dimension in DIMENSIONS: + value = float(raw_scores[dimension]) + if not 0.0 <= value <= 1.0: + raise ValueError(f"Raw score for {dimension} is outside [0, 1]: {value}") + minimum, maximum = NORMALIZE_RANGE[dimension] + normalized[dimension] = (value - minimum) / (maximum - minimum) + return normalized + + +def aggregate_selected_score(raw_scores: Mapping[str, float]) -> dict[str, float | dict[str, float]]: + normalized = normalize_scores(raw_scores) + quality_score = sum( + normalized[dimension] * weight + for dimension, weight in QUALITY_WEIGHTS.items() + ) / sum(QUALITY_WEIGHTS.values()) + semantic_score = sum( + normalized[dimension] * weight + for dimension, weight in SEMANTIC_WEIGHTS.items() + ) / sum(SEMANTIC_WEIGHTS.values()) + selected_score = (4.0 * quality_score + semantic_score) / 5.0 + return { + "normalized_scores": normalized, + "quality_score": quality_score, + "semantic_score": semantic_score, + "selected_vbench_score": selected_score, + "selected_vbench_percent": selected_score * 100.0, + }