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