#!/usr/bin/env python """Extended-251 videos for the Causal-Forcing-a single-block Predictors. Same protocol as generate_eval.py: work is sharded by prompt, FFFF is generated first in the same process and kept as the pixel-metric reference, every strategy sees the identical noise draw, and the per-prompt record has the layout aggregate.py / run_vbench.py expect. Strategies (prefix ``cfa`` = Causal-Forcing-a stack): cfa_ffff reference, 28 full DiT forwards cfa__fpff chunks 1..6: F P F F (22 full + 6 Predictor) cfa__fppf chunks 1..6: F P P F (16 full + 12 Predictor) cfa__fppp chunks 1..6: F P P P (10 full + 18 Predictor) where is one of atc_chunk, atc_last_frame, disca (Stage-1 Layer-17 Predictors) or atc_chunk_s2 (Stage-2 DMD EMA). Chunk 0 is always fully denoised. ``@`` selects another checkpoint of that run and names the strategy cfa__s..., e.g. ``disca@500`` -> cfa_disca_s500_fppf; a bare model uses --checkpoint-step (default 2000) and the unsuffixed name. cfa_naive2step chunks 1..: the 2-step subset of the grid (t = 1000 / 250, the same subset as cf_naive2step_c0FFFF), no Predictor CUDA_VISIBLE_DEVICES=3 python eval/generate_predictor_eval.py --shard 0 --num-shards 4 """ import argparse import json import os import sys import time ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) REPO = os.environ.get("CFA_REPO", "/local/zoubin/cz/projects/Causal-Forcing-a") MAPPING = os.path.join(ROOT, "assets/vbench8_extended_subset_mapping.json") BASE = { "repo": REPO, "config": "configs/causal_forcing_dmd_chunkwise.yaml", "checkpoint": "checkpoints/chunkwise/causal_forcing.pt", } MODELS = { "atc_chunk": "output/layer17_stage1_atc_chunk_1000p_21f_seed0_4gpu_b16_acc1_2000steps", "atc_last_frame": "output/layer17_stage1_atc_last_frame_1000p_21f_seed0_4gpu_b16_acc1_2000steps", "disca": "output/layer17_stage1_disca_1000p_21f_seed0_4gpu_b16_acc1_2000steps", # Stage-2 random-exit DMD (LoRA critic), EMA weights (a (run dir, file name) pair; # a plain directory gets checkpoint_step_/predictor.safetensors). "atc_chunk_s2": ("output/layer17_stage2_dmd_atc_chunk_lora128_4gpu_b2_2000steps", "predictor_ema.safetensors"), } NAIVE = {"naive2step": [0, 3]} # pseudo-models: reduced step grid, no Predictor def checkpoint_weights(model, step): """Weights file of ``model`` at training step ``step`` (Stage-1 runs zero-pad the checkpoint directory, Stage-2 does not).""" spec = MODELS[model] run_dir, fname = spec if isinstance(spec, tuple) else (spec, "predictor.safetensors") for d in (f"checkpoint_step_{step:04d}", f"checkpoint_step_{step}"): w = os.path.join(REPO, run_dir, d, fname) if os.path.exists(w): return w raise FileNotFoundError(f"{model} step {step}: no checkpoint under {os.path.join(REPO, run_dir)}") PATTERNS = {"FPFF": {1}, "FPPF": {1, 2}, "FPPP": {1, 2, 3}} # Reduced schedules: which step indices of the 4-step grid a chunk (after chunk 0) # actually visits. FPF keeps steps 0 and 3 (t = 1000 / 250) as Full and runs one # Predictor step at step 1 (t = 750, anchored on step 0 as in training), skipping # step 2 entirely: 3 denoising positions per chunk, 2 Full + 1 Predictor. RUN_STEPS = {"FPF": [0, 1, 3]} PATTERNS["FPF"] = {1} NUM_CHUNKS, FRAMES_PER_CHUNK, NUM_STEPS = 7, 3, 4 # A Predictor step runs 1 of the 30 Teacher blocks (plus a small fusion MLP); # its compute-equivalent cost is booked as 1/30 of a full forward. PREDICTOR_COMPUTE_EQUIV = 1.0 / 30.0 def parse_args(): p = argparse.ArgumentParser() p.add_argument("--shard", type=int, default=0) p.add_argument("--num-shards", type=int, default=1) p.add_argument("--out-root", default="eval_out") p.add_argument("--seed", type=int, default=0) p.add_argument("--fps", type=int, default=16) p.add_argument("--models", default=",".join(MODELS)) p.add_argument("--patterns", default="FPPF,FPPP") p.add_argument("--checkpoint-step", type=int, default=2000) p.add_argument("--num-chunks", type=int, default=NUM_CHUNKS, help="Video length in 3-latent-frame chunks (7 = the 81-frame protocol; longer " "videos use a rolling 21-frame KV window like Self-Forcing's long-video mode)") p.add_argument("--ref-tag", default="", help="Inserted into the reference name, e.g. _14c -> cfa_14c_ffff (a separate FFFF " "reference for a different video length)") p.add_argument("--tag", default="", help="Inserted into strategy names after the model, e.g. _s1000 -> " "cfa_disca_s1000_fppf, so another checkpoint does not overwrite the default rows") p.add_argument("--regenerate-ffff", action="store_true", help="Regenerate the FFFF reference instead of reading cfa_ffff's MP4") p.add_argument("--limit", type=int, default=None, help="Smoke-test: first N prompts") p.add_argument("--overwrite", action="store_true") p.add_argument("--retime-only", action="store_true", help="Do not regenerate videos/metrics; re-time FFFF and each strategy in the same " "process per prompt and update policy_latency_ms / matched_ffff_policy_latency_ms.") p.add_argument("--retime-max-ffff-ms", type=float, default=3500.0, help="With --retime-only, records whose paired FFFF exceeds this are re-timed again " "(up to --retime-passes passes); clean records are left alone.") p.add_argument("--retime-passes", type=int, default=3) p.add_argument("--retime-strategy-only", action="store_true", help="Re-time only the strategies (no FFFF run): rerun each selected record's " "rollout on an idle GPU and replace policy_latency_ms; the generation-run " "value is kept under generation_run_latency.") p.add_argument("--written-before", type=float, default=None, help="With --retime-strategy-only: only records whose file mtime is before this " "epoch second (the window in which another job shared the GPU)") return p.parse_args() def build_strategies(args): out = [{"name": f"cfa{args.ref_tag}_ffff", "model": None, "pattern": "FFFF", "is_reference": True}] for token in args.models.split(","): if token in NAIVE: out.append({"name": f"cfa_{token}{args.tag}", "model": None, "pattern": token.upper(), "naive": token, "is_reference": False}) continue model, _, step = token.partition("@") step = int(step) if step else args.checkpoint_step key = f"{model}@{step}" # one loaded Predictor per (model, step) label = f"{model}_s{step}" if "@" in token else model for pattern in args.patterns.split(","): out.append({"name": f"cfa_{label}{args.tag}_{pattern.lower()}", "model": key, "base_model": model, "checkpoint_step": step, "pattern": pattern.upper(), "is_reference": False}) return out class FinalHiddenCapture: """Grab the Teacher's final hidden (input of its output head) on a Full step.""" def __init__(self, teacher): self.enabled, self.value = False, None self.handle = teacher.head.register_forward_pre_hook(self._hook) def _hook(self, _module, inputs): if self.enabled: self.value = inputs[0].detach() def start(self): self.value, self.enabled = None, True def finish(self): 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 main(): args = parse_args() out_root = args.out_root if os.path.isabs(args.out_root) else os.path.join(ROOT, args.out_root) with open(MAPPING) as f: rows = json.load(f)["rows"] if args.limit: rows = rows[:args.limit] shard_rows = rows[args.shard::args.num_shards] strategies = build_strategies(args) ref_name = strategies[0]["name"] print(f"cfa shard {args.shard}/{args.num_shards}: {len(shard_rows)} prompts x " f"{len(strategies)} strategies", flush=True) os.chdir(REPO) sys.path.insert(0, ROOT) sys.path.insert(0, REPO) # pipeline/, utils/, predictor_training/ from Causal-Forcing-a import torch from einops import rearrange from safetensors.torch import load_file from torchvision.io import read_video, write_video from eval.pixel_metrics import PixelMetrics from harness import load_pipeline from predictor_training.metadata import build_single_block_predictor, read_predictor_config from predictor_training.online import online_predictor_step as _online_predictor_step from predictor_training.offline_data import TOKENS_PER_CHUNK def online_predictor_step(*, history_cache, chunk, **kw): """Stage-1 contract: the Predictor sees the clean K/V of the *earlier* chunks. With the rolling 21-frame window (long videos) the cache holds the last six clean chunks plus the current noisy one, so the history is everything before the current chunk's tokens rather than ``chunk * TOKENS_PER_CHUNK`` (which the stock helper asserts and which exceeds the window from chunk 7 on). For chunk <= 6 the two are identical.""" if chunk * TOKENS_PER_CHUNK <= int(history_cache["local_end_index"].item()): return _online_predictor_step(history_cache=history_cache, chunk=chunk, **kw) end = int(history_cache["local_end_index"].item()) - TOKENS_PER_CHUNK window = {"k": history_cache["k"][:, :end], "v": history_cache["v"][:, :end], "local_end_index": torch.tensor([end], device=history_cache["k"].device)} # left-pad to the absolute length the stock helper slices; the attention # window (max_attention_size = 21 frames) never reaches the padding pad = chunk * TOKENS_PER_CHUNK - end if pad > 0: z = history_cache["k"].new_zeros(history_cache["k"].shape[0], pad, *history_cache["k"].shape[2:]) window = {"k": torch.cat([z, window["k"]], 1), "v": torch.cat([z, window["v"]], 1), "local_end_index": torch.tensor([chunk * TOKENS_PER_CHUNK], device=z.device)} return _online_predictor_step(history_cache=window, chunk=chunk, **kw) from utils.misc import set_seed torch.set_grad_enabled(False) device = torch.device("cuda") pipeline = load_pipeline(BASE) n_chunks = args.num_chunks if n_chunks * FRAMES_PER_CHUNK > 21: # Longer than the 21-latent-frame training context: roll the KV cache over a # 21-frame local window (the model's long-video mechanism; RoPE is relative # so positions inside the window stay in the trained range). model = pipeline.generator.model for m in [model] + [b for b in model.blocks] + [b.self_attn for b in model.blocks]: m.local_attn_size = 21 if hasattr(m, "max_attention_size"): m.max_attention_size = 21 * pipeline.frame_seq_length pipeline.local_attn_size = 21 pipeline.kv_cache1 = None print(f"long video: {n_chunks} chunks = {n_chunks * FRAMES_PER_CHUNK} latent frames, " f"rolling KV window 21 frames", flush=True) teacher = pipeline.generator.model capture = FinalHiddenCapture(teacher) metrics = PixelMetrics(device) timesteps = pipeline.denoising_step_list.to(device) assert len(timesteps) == NUM_STEPS and pipeline.num_frame_per_block == FRAMES_PER_CHUNK predictors, predictor_configs = {}, {} for model in {s["model"] for s in strategies if s["model"]}: base, step = model.split("@") weights = checkpoint_weights(base, int(step)) config = read_predictor_config(weights) predictor = build_single_block_predictor(teacher, config, gradient_checkpointing=False) predictor.load_state_dict(load_file(weights, device="cpu"), strict=True) predictors[model] = predictor.to(device).eval().requires_grad_(False) predictor_configs[model] = {**config, "weights": weights} print(f"loaded {model}: {config}", flush=True) def reset_caches(): 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"] = torch.tensor([0], dtype=torch.long, device=device) cache["local_end_index"] = torch.tensor([0], dtype=torch.long, device=device) for cache in pipeline.crossattn_cache: cache["is_init"] = False def rollout(conditional, noise, predictor=None, predictor_steps=frozenset(), run_steps=None): """Mirror of cachelib.runner.cached_inference with Predictor steps spliced in.""" reset_caches() output = torch.zeros_like(noise) denoise_events, context_events = [], [] full_calls = predictor_calls = 0 previous_hidden = None for chunk in range(n_chunks): start_frame = chunk * FRAMES_PER_CHUNK token_start = start_frame * pipeline.frame_seq_length noisy_input = noise[:, start_frame:start_frame + FRAMES_PER_CHUNK] current_hidden = [None] * NUM_STEPS # chunk 0 always runs the full grid; later chunks may visit a subset steps = list(range(NUM_STEPS)) if (run_steps is None or chunk == 0) else list(run_steps) for si, step in enumerate(steps): timestep = torch.ones([1, FRAMES_PER_CHUNK], device=device, dtype=torch.int64) * timesteps[step] start_ev = torch.cuda.Event(enable_timing=True) end_ev = torch.cuda.Event(enable_timing=True) start_ev.record() if predictor is not None and chunk > 0 and step in predictor_steps: hidden, flow = online_predictor_step( predictor=predictor, teacher=teacher, noisy_input=noisy_input, timestep=timestep, anchor_timestep=torch.ones_like(timestep) * timesteps[step - 1], anchor_hidden=current_hidden[step - 1], previous_hidden=previous_hidden[step], history_cache=pipeline.kv_cache1[predictor.source_layer], cross_cache=pipeline.crossattn_cache[predictor.source_layer], chunk=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, timestep=timestep, kv_cache=pipeline.kv_cache1, crossattn_cache=pipeline.crossattn_cache, current_start=token_start) current_hidden[step] = capture.finish() full_calls += 1 end_ev.record() denoise_events.append((start_ev, end_ev)) if si < len(steps) - 1: # Same RNG draw order as the protocol runner, so every strategy # re-noises with identical samples. noisy_input = pipeline.scheduler.add_noise( denoised_pred.flatten(0, 1), torch.randn_like(denoised_pred.flatten(0, 1)), timesteps[steps[si + 1]] * torch.ones([FRAMES_PER_CHUNK], device=device, dtype=torch.long), ).unflatten(0, denoised_pred.shape[:2]) output[:, start_frame:start_frame + FRAMES_PER_CHUNK] = denoised_pred # Clean-context KV refresh: full DiT, timed separately, never added in. ctx_start = torch.cuda.Event(enable_timing=True) ctx_end = torch.cuda.Event(enable_timing=True) ctx_start.record() pipeline.generator( noisy_image_or_video=denoised_pred, 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) ctx_end.record() context_events.append((ctx_start, ctx_end)) previous_hidden = current_hidden torch.cuda.synchronize() timing = { "denoise_dit_ms": sum(s.elapsed_time(e) for s, e in denoise_events), "context_kv_dit_ms": sum(s.elapsed_time(e) for s, e in context_events), "num_denoise_forwards": len(denoise_events), } video = pipeline.vae.decode_to_pixel(output, use_cache=False) video = (video * 0.5 + 0.5).clamp(0, 1) pipeline.vae.model.clear_cache() return video, timing, {"full_forwards": full_calls, "predictor_forwards": predictor_calls} def run_one(strategy, row): set_seed(args.seed) noise = torch.randn([1, n_chunks * FRAMES_PER_CHUNK, 16, 60, 104], device=device, dtype=torch.bfloat16) conditional = pipeline.text_encoder(text_prompts=[row["extended_prompt"]]) wall0 = time.time() if strategy["is_reference"]: video, timing, counts = rollout(conditional, noise) elif strategy.get("naive"): video, timing, counts = rollout(conditional, noise, None, frozenset(), NAIVE[strategy["naive"]]) else: video, timing, counts = rollout(conditional, noise, predictors[strategy["model"]], PATTERNS[strategy["pattern"]], RUN_STEPS.get(strategy["pattern"])) return video, timing, counts, time.time() - wall0 def record_path(strategy, row): d = os.path.join(out_root, "per_prompt", strategy) os.makedirs(d, exist_ok=True) return os.path.join(d, f"{row['prompt_suite']}_{row['suite_index']:03d}.json") def video_path(strategy, row): d = os.path.join(out_root, "generated_videos", strategy, row["prompt_suite"]) os.makedirs(d, exist_ok=True) return os.path.join(d, f"{row['suite_index']:03d}.mp4") def is_done(strategy, row): p = record_path(strategy, row) if args.overwrite or not os.path.exists(p): return False try: with open(p) as f: rec = json.load(f) return rec.get("status") == "complete" and os.path.exists(rec.get("video", "")) except Exception: return False def write_json(path, payload): tmp = path + ".tmp" with open(tmp, "w") as f: json.dump(payload, f, indent=2) os.replace(tmp, path) if shard_rows: print("warm-up pass (not recorded)", flush=True) for strategy in strategies: run_one(strategy, shard_rows[0]) torch.cuda.empty_cache() if args.retime_strategy_only: targets = [s for s in strategies if not s["is_reference"]] n_done = 0 for n, row in enumerate(shard_rows): for strategy in targets: rp = record_path(strategy["name"], row) if not os.path.exists(rp): continue if args.written_before and os.path.getmtime(rp) >= args.written_before: continue rec = json.load(open(rp)) if rec.get("status") != "complete" or "generation_run_latency" in rec: continue _, timing, counts, _ = run_one(strategy, row) assert counts["full_forwards"] == rec["cache_diagnostics"]["full_forwards"], rp rec["generation_run_latency"] = { "policy_latency_ms": rec["policy_latency_ms"], "excluded_context_kv_latency_ms": rec["excluded_context_kv_latency_ms"]} rec["policy_latency_ms"] = timing["denoise_dit_ms"] rec["excluded_context_kv_latency_ms"] = timing["context_kv_dit_ms"] rec["latency_retimed"] = "strategy-only rerun on an idle GPU (FFFF not re-run)" write_json(rp, rec) n_done += 1 if (n + 1) % 20 == 0: print(f"[strategy retime shard {args.shard}] {n + 1}/{len(shard_rows)}", flush=True) print(f"shard {args.shard}: {n_done} records re-timed", flush=True) return if args.retime_only: # Protocol "matched FFFF": pair every strategy with an FFFF run of the same # prompt in the same process, seconds apart, so contention cancels in the # ratio. Only records whose pairing is missing or dirty are touched. targets = [s for s in strategies if not s["is_reference"]] for pass_index in range(args.retime_passes): dirty = 0 for n, row in enumerate(shard_rows): todo = [] for strategy in targets: rp = record_path(strategy["name"], row) if not os.path.exists(rp): continue rec = json.load(open(rp)) m = rec.get("matched_ffff_policy_latency_ms") if m is None or m > args.retime_max_ffff_ms: todo.append((strategy, rp, rec)) if not todo: continue dirty += len(todo) _, ref_timing, _, _ = run_one(strategies[0], row) for strategy, rp, rec in todo: _, timing, _, _ = run_one(strategy, row) rec["policy_latency_ms"] = timing["denoise_dit_ms"] rec["excluded_context_kv_latency_ms"] = timing["context_kv_dit_ms"] rec["matched_ffff_policy_latency_ms"] = ref_timing["denoise_dit_ms"] rec["matched_ffff_context_kv_latency_ms"] = ref_timing["context_kv_dit_ms"] rec["latency_source"] = "paired retime (same process as FFFF)" write_json(rp, rec) if (n + 1) % 10 == 0: print(f"[retime pass {pass_index} shard {args.shard}] {n + 1}/{len(shard_rows)} " f"ffff={ref_timing['denoise_dit_ms']:.0f}ms", flush=True) print(f"[retime pass {pass_index} shard {args.shard}] re-timed {dirty} records", flush=True) if dirty == 0: break return t_start = time.time() for n, row in enumerate(shard_rows): todo = [s for s in strategies if not is_done(s["name"], row)] if not todo: continue need_ref = any(not s["is_reference"] for s in todo) ref_frames = None ref_timing = None ref_from_mp4 = False if need_ref and os.path.exists(video_path(ref_name, row)) and not args.regenerate_ffff: # Reference from the cfa_ffff run's MP4 (H.264, 40+ dB above the strategy # PSNRs); FFFF is not regenerated, so there is no same-process FFFF timing # and the speedup is taken against cfa_ffff's own records (aggregate.py). vid, _, _ = read_video(video_path(ref_name, row), pts_unit="sec", output_format="TCHW") ref_frames = vid.to(device, torch.float16) / 255.0 ref_from_mp4 = True need_ref = False for strategy in strategies: if strategy["is_reference"]: if strategy not in todo and not need_ref: continue elif strategy not in todo: continue video, timing, counts, wall = run_one(strategy, row) frames = video[0] if strategy["is_reference"]: ref_frames = frames.detach().to(torch.float16) ref_timing = timing pix = PixelMetrics.identity(frames.shape[0]) else: pix = metrics.compute(frames, ref_frames) if strategy in todo: vp = video_path(strategy["name"], row) arr = (255.0 * rearrange(frames, "t c h w -> t h w c")).clamp(0, 255) write_video(vp, arr.to(torch.uint8).cpu(), fps=args.fps) compute_equiv = counts["full_forwards"] + PREDICTOR_COMPUTE_EQUIV * counts["predictor_forwards"] model = strategy["model"] naive = strategy.get("naive") schedule = "F" * len(NAIVE[naive]) if naive else strategy["pattern"] write_json(record_path(strategy["name"], row), { "status": "complete", "protocol": "Self-Forcing Extended-251 Full Evaluation", "strategy": strategy["name"], "base_model": "causal_forcing_a", "method": "none" if model is None else f"predictor_{strategy['base_model']}", "param": None if model is None else "pattern", "param_value": None if model is None else strategy["pattern"], "target_speedup": (n_chunks * NUM_STEPS) / compute_equiv, "schedule": schedule, "num_inference_steps": len(NAIVE[naive]) if naive else NUM_STEPS, "first_chunk_schedule": "FFFF" if (naive or model) else None, "predictor_config": None if model is None else predictor_configs[model], "global_index": row["global_index"], "prompt_suite": row["prompt_suite"], "suite_index": row["suite_index"], "prompt": row["extended_prompt"], "original_prompt": row["original_prompt"], "seed": args.seed, "video": vp, "num_frames": int(frames.shape[0]), "height": int(frames.shape[2]), "width": int(frames.shape[3]), "fps": args.fps, "policy_latency_ms": timing["denoise_dit_ms"], "excluded_context_kv_latency_ms": timing["context_kv_dit_ms"], # Same-process FFFF timing of this prompt (protocol "matched FFFF"). "matched_ffff_policy_latency_ms": ref_timing["denoise_dit_ms"] if ref_timing else None, "matched_ffff_context_kv_latency_ms": ref_timing["context_kv_dit_ms"] if ref_timing else None, "checkpoint_step": strategy.get("checkpoint_step"), "num_chunks": n_chunks, "wall_generation_s": wall, "reference_strategy": ref_name, "pixel_metrics_vs_ffff": pix, "reference_source": "ffff_mp4" if ref_from_mp4 else "ffff_frames", "cache_diagnostics": { "denoise_forwards": timing["num_denoise_forwards"], "full_forwards": counts["full_forwards"], "predictor_forwards": counts["predictor_forwards"], "predictor_compute_equivalent": PREDICTOR_COMPUTE_EQUIV, "compute_equivalent_forwards": compute_equiv, "middle_steps": n_chunks * 2, "middle_compute_equivalent": compute_equiv - n_chunks * 2.0, }, }) del frames, video ref_frames = None done = n + 1 rate = (time.time() - t_start) / done print(f"[cfa shard {args.shard}] {done}/{len(shard_rows)} prompts " f"({row['prompt_suite']}/{row['suite_index']:03d}) {rate:.1f}s/prompt " f"eta {(len(shard_rows) - done) * rate / 60:.0f} min", flush=True) print(f"shard {args.shard} done in {(time.time() - t_start) / 60:.1f} min", flush=True) if __name__ == "__main__": main()