| |
| """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, |
| 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 |
| |
| |
| |
| 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) |
| |
| |
| |
| 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() |
|
|