Download finalize.py from Cccccz/comparison: direct link, hf CLI and curl.
- Browser
- Download file 12 kB
-
https://huggingface.co/Cccccz/comparison/resolve/main/finalize.py
- Command line
-
hf download hf://Cccccz/comparison/finalize.py
-
curl -L -o finalize.py https://huggingface.co/Cccccz/comparison/resolve/main/finalize.py
12 kB
| #!/usr/bin/env python | |
| """Authoritative pass: re-target and cleanly re-measure every operating point. | |
| Two things the parallel sweeps could not do well: | |
| 1. **Targeting.** The sweeps aimed at ``target / measured_overhead``, but that | |
| overhead is itself a timing measurement on a contended node -- one probe came | |
| back at 1.110, which is impossible, and it dragged that method's operating | |
| points well below their targets. Here the search aims at the **compute budget** | |
| directly, which is deterministic, and also puts all four methods on identical | |
| compute so a later quality comparison is like-for-like. | |
| 2. **Timing.** The sweeps ran four at a time across the node, so their wall-clock | |
| column includes contention from each other. Here every point is measured | |
| serially in one session, paired against the baseline on the same prompt. | |
| Each point is also re-checked on a held-out prompt slice: a threshold that only | |
| reached its target by sitting inside a narrow band of the indicator distribution | |
| shows up here as compute-fraction drift. | |
| python finalize.py --base self_forcing --out results/final_self_forcing.json | |
| """ | |
| import argparse | |
| import json | |
| import os | |
| import sys | |
| ROOT = os.path.dirname(os.path.abspath(__file__)) | |
| sys.path.insert(0, ROOT) | |
| from harness import (BASES, enter_base, evaluate, load_pipeline, # noqa: E402 | |
| load_prompts, paired_evaluate) | |
| METHOD_ORDER = ["teacache", "flowcache", "taylorseer", "motioncache"] | |
| def main(): | |
| ap = argparse.ArgumentParser() | |
| ap.add_argument("--base", choices=sorted(BASES), required=True) | |
| ap.add_argument("--pair-prompts", type=int, default=10) | |
| ap.add_argument("--search-prompts", type=int, default=4) | |
| ap.add_argument("--tolerance", type=float, default=0.02, | |
| help="Re-bisect when |flops - target| exceeds this fraction") | |
| ap.add_argument("--bisect-iters", type=int, default=18) | |
| ap.add_argument("--correction-iters", type=int, default=9, | |
| help="Bisection steps for the correction pass on the timed prompts") | |
| ap.add_argument("--heldout-offset", type=int, default=64) | |
| ap.add_argument("--heldout-prompts", type=int, default=8) | |
| ap.add_argument("--num-output-frames", type=int, default=21) | |
| ap.add_argument("--seed", type=int, default=0) | |
| ap.add_argument("--save-video-dir", default=None) | |
| ap.add_argument("--methods", default=None, | |
| help="Comma-separated subset of methods to (re)measure") | |
| ap.add_argument("--suffix", default="", | |
| help="Look for results/sweep_<base>_<method><suffix>.json") | |
| ap.add_argument("--targets", default=None, | |
| help="Comma-separated subset of target speedups to (re)measure") | |
| ap.add_argument("--resume", action="store_true", | |
| help="Keep rows already in --out and skip those (method, target) pairs") | |
| ap.add_argument("--out", required=True) | |
| args = ap.parse_args() | |
| out_path = args.out if os.path.isabs(args.out) else os.path.join(ROOT, args.out) | |
| done = {} | |
| rows = [] | |
| if args.resume and os.path.exists(out_path): | |
| with open(out_path) as f: | |
| rows = json.load(f).get("rows", []) | |
| done = {(r["method"], r["target_speedup"]) for r in rows} | |
| print(f"resuming: {len(rows)} rows already done", flush=True) | |
| else: | |
| done = set() | |
| wanted = set(args.methods.split(",")) if args.methods else None | |
| sweeps = [] | |
| for m in METHOD_ORDER: | |
| if wanted and m not in wanted: | |
| continue | |
| sp = os.path.join(ROOT, f"results/sweep_{args.base}_{m}{args.suffix}.json") | |
| if os.path.exists(sp): | |
| with open(sp) as f: | |
| sweeps.append(json.load(f)) | |
| if not sweeps: | |
| print(f"no sweep files for {args.base}") | |
| return | |
| print(f"methods: {[s['method'] for s in sweeps]}", flush=True) | |
| base = enter_base(args.base) | |
| import torch | |
| from cachelib import CacheController, build_method, install, parse_schedule | |
| torch.set_grad_enabled(False) | |
| pipeline = load_pipeline(base) | |
| num_steps = len(pipeline.denoising_step_list) | |
| tune_prompts = load_prompts(args.pair_prompts + 1) | |
| held_prompts = load_prompts(args.heldout_prompts, offset=args.heldout_offset) | |
| def make_ctrl(sweep, method_name, value): | |
| kw = dict(indicator=sweep.get("indicator", "modulated_input"), | |
| coefficients=sweep.get("coefficients")) | |
| if method_name != "none": | |
| kw[sweep["param"]] = value | |
| ctrl = CacheController( | |
| build_method(method_name, **kw), num_steps=num_steps, | |
| forced_steps=parse_schedule(sweep.get("schedule", "FxxF"), num_steps)) | |
| install(pipeline.generator.model, ctrl) | |
| return ctrl | |
| def flops_at(sweep, value, prompts=None): | |
| r = evaluate(pipeline, make_ctrl(sweep, sweep["method"], value), | |
| prompts if prompts is not None else tune_prompts[:args.search_prompts], | |
| num_output_frames=args.num_output_frames, seed=args.seed, | |
| warmup=0) | |
| return r["flops_speedup_estimate"], r["mean_compute_equivalent_forwards"] | |
| def retarget(sweep, target, prompts=None, iters=None): | |
| """Bisect the stored curve's bracket for the requested compute budget.""" | |
| curve = sweep["curve"] | |
| max_flops = max(c["flops_speedup"] for c in curve) | |
| goal = min(target, max_flops) | |
| lo = min(c["value"] for c in curve) | |
| hi = max(c["value"] for c in curve) | |
| a, b = lo, hi | |
| for c in curve: | |
| if c["flops_speedup"] < goal: | |
| a = max(a, c["value"]) | |
| for c in reversed(curve): | |
| if c["flops_speedup"] >= goal: | |
| b = min(b, c["value"]) | |
| best = None | |
| for _ in range(iters or args.bisect_iters): | |
| mid = 0.5 * (a + b) | |
| f, ce = flops_at(sweep, mid, prompts) | |
| if best is None or abs(f - goal) < abs(best[1] - goal): | |
| best = (mid, f, ce) | |
| if f < goal: | |
| a = mid | |
| else: | |
| b = mid | |
| # Bisection only probes interior points; when the curve steps up exactly at | |
| # the upper bracket it converges from below and never samples the value that | |
| # actually reaches the target. Score the closing ends too. | |
| for endpoint in (a, b): | |
| f, ce = flops_at(sweep, endpoint, prompts) | |
| if abs(f - goal) < abs(best[1] - goal): | |
| best = (endpoint, f, ce) | |
| return best | |
| print("=== warm-up ===", flush=True) | |
| evaluate(pipeline, make_ctrl(sweeps[0], "none", None), tune_prompts[:3], | |
| num_output_frames=args.num_output_frames, seed=args.seed, warmup=0) | |
| for sweep in sweeps: | |
| pname = sweep["param"] | |
| for p in sweep["operating_points"]: | |
| target = p["target_speedup"] | |
| if (sweep["method"], target) in done: | |
| continue | |
| if args.targets and target not in [float(t) for t in args.targets.split(",")]: | |
| continue | |
| value = p[pname] | |
| f, ce = flops_at(sweep, value) | |
| retargeted = False | |
| if abs(f - target) > args.tolerance * target: | |
| new = retarget(sweep, target) | |
| if abs(new[1] - target) < abs(f - target): | |
| value, f, ce = new | |
| retargeted = True | |
| print(f" retargeted {sweep['method']} {target}x -> " | |
| f"{pname}={value:.6g} flops={f:.3f}", flush=True) | |
| vid_dir = None | |
| if args.save_video_dir: | |
| sched = sweep.get("schedule", "FxxF") | |
| tag = "" if sched.upper() == "FXXF" else f"_{sched}" | |
| vid_dir = os.path.join(ROOT, args.save_video_dir, | |
| f"{args.base}_{sweep['method']}{tag}_x{target:g}") | |
| def measure(v, save=None): | |
| return paired_evaluate( | |
| pipeline, | |
| lambda s=sweep: make_ctrl(s, "none", None), | |
| lambda s=sweep, vv=v: make_ctrl(s, s["method"], vv), | |
| tune_prompts, num_output_frames=args.num_output_frames, | |
| seed=args.seed, warmup=1, save_video_dir=save) | |
| paired = measure(value, save=vid_dir) | |
| # The compute fraction of the timed runs is the one that counts. Near a | |
| # decision boundary it can differ from the search subset's, so if it | |
| # misses, re-tune on the timed prompts themselves and measure again. | |
| if abs(paired["paired_flops_speedup"] - target) > args.tolerance * target: | |
| new_pt = retarget(sweep, target, prompts=tune_prompts[1:], | |
| iters=args.correction_iters) | |
| if abs(new_pt[1] - target) < abs(paired["paired_flops_speedup"] - target): | |
| value, _, _ = new_pt | |
| retargeted = True | |
| print(f" corrected {sweep['method']} {target}x on timed prompts " | |
| f"-> {pname}={value:.6g} flops={new_pt[1]:.3f}", flush=True) | |
| paired = measure(value, save=vid_dir) | |
| f = paired["paired_flops_speedup"] | |
| ce = paired["paired_compute_equivalent"] | |
| held = evaluate(pipeline, make_ctrl(sweep, sweep["method"], value), | |
| held_prompts, num_output_frames=args.num_output_frames, | |
| seed=args.seed, warmup=0) | |
| row = { | |
| "base": args.base, | |
| "method": sweep["method"], | |
| "param": pname, | |
| "value": value, | |
| "schedule": sweep.get("schedule", "FxxF"), | |
| "retargeted": retargeted, | |
| "target_speedup": target, | |
| "flops_speedup": f, | |
| "compute_equivalent_forwards": ce, | |
| "denoise_forwards": 28.0, | |
| "measured_speedup": paired["speedup_from_minima"], | |
| "measured_speedup_paired_median": paired["paired_speedup_median"], | |
| "measured_speedup_paired_mean": paired["paired_speedup_mean"], | |
| "measured_speedup_paired_stdev": paired["paired_speedup_stdev"], | |
| "measured_speedup_paired_range": [paired["paired_speedup_min"], | |
| paired["paired_speedup_max"]], | |
| "min_baseline_ms": paired["min_baseline_ms"], | |
| "min_method_ms": paired["min_method_ms"], | |
| "median_baseline_ms": paired["median_baseline_ms"], | |
| "median_method_ms": paired["median_method_ms"], | |
| "baseline_ms": paired["baseline_ms"], | |
| "method_ms": paired["method_ms"], | |
| "num_pairs": paired["num_pairs"], | |
| "heldout_flops_speedup": held["flops_speedup_estimate"], | |
| "heldout_compute_equivalent": held["mean_compute_equivalent_forwards"], | |
| "heldout_flops_drift": held["flops_speedup_estimate"] - f, | |
| "video_dir": vid_dir, | |
| } | |
| rows.append(row) | |
| print(f"{sweep['method']:12s} target {target:g}x " | |
| f"{pname}={value:<10.6g} measured={row['measured_speedup']:.3f}x " | |
| f"(paired {row['measured_speedup_paired_median']:.3f}" | |
| f"±{row['measured_speedup_paired_stdev']:.3f}) flops={f:.3f}x " | |
| f"heldout_flops={row['heldout_flops_speedup']:.3f}x " | |
| f"({row['heldout_flops_drift']:+.3f})", flush=True) | |
| with open(out_path, "w") as fh: | |
| json.dump({"base": args.base, "pair_prompts": args.pair_prompts, | |
| "heldout_offset": args.heldout_offset, "rows": rows}, | |
| fh, indent=2) | |
| print(f"wrote {out_path}", flush=True) | |
| if __name__ == "__main__": | |
| main() | |