File size: 11,969 Bytes
38a51ff
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
#!/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()