#!/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__.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()