#!/usr/bin/env python """Paired re-timing of evaluated strategies against FFFF, prompt by prompt. ``generate_eval.py`` records each strategy's denoise latency, but when it is re-run for *additional* strategies the FFFF reference is regenerated only for the pixel metrics and its latency is not recorded, so the new records would be compared against FFFF timings from a different run under different contention. The protocol's speedup is "matched FFFF" (section 9.4): same prompt, same seed, same conditions. This pass restores that: for every prompt it runs FFFF and then each requested strategy back to back (denoise only, no decode, no metrics) and writes the strategy's latency together with the FFFF latency measured seconds earlier into the strategy's per-prompt record. CUDA_VISIBLE_DEVICES=0 python eval/retime_eval.py --base self_forcing \ --strategies sf_taylorseer_FFxx_x1.3,sf_teacache_Fxxx_x3 \ --shard 0 --num-shards 4 --out-root eval_out Records already carrying ``latency_source`` are skipped, so the pass resumes. """ import argparse import json import os import sys import time ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) sys.path.insert(0, ROOT) from eval.strategies import load_strategies # noqa: E402 from harness import BASES, enter_base, load_pipeline # noqa: E402 MAPPING = os.path.join(ROOT, "assets/vbench8_extended_subset_mapping.json") SOURCE = "retime_eval: paired with FFFF in the same process, denoise only" def parse_args(): p = argparse.ArgumentParser() p.add_argument("--base", choices=sorted(BASES), required=True) p.add_argument("--strategies", required=True, help="Comma-separated strategy names, or 'all' for every " "non-reference strategy of the base") 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("--num-latent-frames", type=int, default=21) p.add_argument("--limit", type=int, default=None) p.add_argument("--force", action="store_true", help="Re-time records already re-timed") p.add_argument("--redo-ratio-bounds", default=None, help="lo,hi: also re-time records whose measured/compute speedup " "ratio falls outside this band -- a spike that hit the strategy " "half of the pair rather than the FFFF half") p.add_argument("--redo-above-ms", type=float, default=None, help="Also re-time records whose matched FFFF latency exceeds this: " "a neighbour's job hit that pair, and the pair is only fair " "when both halves ran under the same conditions") return p.parse_args() 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] all_strategies = load_strategies(args.base) reference = all_strategies[0] assert reference["is_reference"] if args.strategies == "all": wanted = [s for s in all_strategies if not s["is_reference"]] else: names = set(args.strategies.split(",")) wanted = [s for s in all_strategies if s["name"] in names] missing = names - {s["name"] for s in wanted} if missing: sys.exit(f"unknown strategies for {args.base}: {sorted(missing)}") print(f"base={args.base} shard {args.shard}/{args.num_shards}: " f"{len(shard_rows)} prompts x ({reference['name']} + {len(wanted)} strategies)", flush=True) base = enter_base(args.base) import torch from utils.misc import set_seed from cachelib import (CacheController, build_method, cached_inference, install, parse_schedule) torch.set_grad_enabled(False) device = torch.device("cuda") pipeline = load_pipeline(base) num_steps = len(pipeline.denoising_step_list) def record_path(strategy, row): return os.path.join(out_root, "per_prompt", strategy, f"{row['prompt_suite']}_{row['suite_index']:03d}.json") def time_one(strategy, row): kw = dict(indicator=strategy["indicator"], coefficients=strategy["coefficients"]) if strategy["method"] != "none": kw[strategy["param"]] = strategy["value"] first = strategy.get("first_chunk_schedule") ctrl = CacheController( build_method(strategy["method"], **kw), num_steps=num_steps, forced_steps=parse_schedule(strategy.get("schedule", "FxxF"), num_steps), first_chunk_forced_steps=(parse_schedule(first, num_steps) if first else None)) install(pipeline.generator.model, ctrl) set_seed(args.seed) noise = torch.randn([1, args.num_latent_frames, 16, 60, 104], device=device, dtype=torch.bfloat16) _, _, timing = cached_inference(pipeline, ctrl, noise, [row["extended_prompt"]], decode=False) return timing, ctrl.summary() def write_json(path, payload): tmp = path + ".tmp" with open(tmp, "w") as f: json.dump(payload, f, indent=2) os.replace(tmp, path) def needs(strategy, row): p = record_path(strategy["name"], row) if not os.path.exists(p): return False # nothing to attach the timing to; generation must run first with open(p) as f: rec = json.load(f) if rec.get("status") != "complete": return False if args.force or "latency_source" not in rec: return True if (args.redo_above_ms is not None and rec.get("matched_ffff_policy_latency_ms", 0) > args.redo_above_ms): return True if args.redo_ratio_bounds: lo, hi = (float(v) for v in args.redo_ratio_bounds.split(",")) d = rec["cache_diagnostics"] compute_x = d["denoise_forwards"] / max(d["compute_equivalent_forwards"], 1e-6) measured_x = rec["matched_ffff_policy_latency_ms"] / rec["policy_latency_ms"] if not (lo <= measured_x / compute_x <= hi): return True return False if shard_rows: print("warm-up pass (not recorded)", flush=True) for s in [reference] + wanted: time_one(s, shard_rows[0]) torch.cuda.empty_cache() t0 = time.time() for n, row in enumerate(shard_rows): todo = [s for s in wanted if needs(s, row)] if not todo: continue ref_timing, _ = time_one(reference, row) for s in todo: timing, summary = time_one(s, row) p = record_path(s["name"], row) with open(p) as f: rec = json.load(f) rec["policy_latency_ms_generation_run"] = rec.get("policy_latency_ms") 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["latency_source"] = SOURCE # The compute fraction is deterministic; recording it again guards # against a mismatch between the two runs. rec["cache_diagnostics"]["compute_equivalent_forwards_retime"] = summary.get( "compute_equivalent_forwards") write_json(p, rec) done = n + 1 rate = (time.time() - t0) / done print(f"[{args.base} retime shard {args.shard}] {done}/{len(shard_rows)} " f"({row['prompt_suite']}/{row['suite_index']:03d}) ffff={ref_timing['denoise_dit_ms']:.0f}ms " f"{rate:.1f}s/prompt eta {(len(shard_rows) - done) * rate / 60:.0f} min", flush=True) print(f"retime shard {args.shard} done in {(time.time() - t0) / 60:.1f} min", flush=True) if __name__ == "__main__": main()