comparison / eval /retime_eval.py
Cccccz's picture
Add files using upload-large-folder tool
f70ac4f verified
Raw History Blame Contribute Delete
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()