Download eval/generate_predictor_eval.py from Cccccz/comparison: direct link, hf CLI and curl.
- Browser
- Download file 28.6 kB
-
https://huggingface.co/Cccccz/comparison/resolve/main/eval/generate_predictor_eval.py
- Command line
-
hf download hf://Cccccz/comparison/eval/generate_predictor_eval.py
-
curl -L -o generate_predictor_eval.py https://huggingface.co/Cccccz/comparison/resolve/main/eval/generate_predictor_eval.py
28.6 kB
| #!/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_<model>_fpff chunks 1..6: F P F F (22 full + 6 Predictor) | |
| cfa_<model>_fppf chunks 1..6: F P P F (16 full + 12 Predictor) | |
| cfa_<model>_fppp chunks 1..6: F P P P (10 full + 18 Predictor) | |
| where <model> 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. | |
| ``<model>@<step>`` selects another checkpoint of that run and names the strategy | |
| cfa_<model>_s<step>..., 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_<N>/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() | |