Download eval/retime_eval.py from Cccccz/comparison: direct link, hf CLI and curl.
- Browser
- Download file 8.27 kB
-
https://huggingface.co/Cccccz/comparison/resolve/main/eval/retime_eval.py
- Command line
-
hf download hf://Cccccz/comparison/eval/retime_eval.py
-
curl -L -o retime_eval.py https://huggingface.co/Cccccz/comparison/resolve/main/eval/retime_eval.py
8.27 kB
| #!/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() | |