Download build_points.py from Cccccz/comparison: direct link, hf CLI and curl.
- Browser
- Download file 4.35 kB
-
https://huggingface.co/Cccccz/comparison/resolve/main/build_points.py
- Command line
-
hf download hf://Cccccz/comparison/build_points.py
-
curl -L -o build_points.py https://huggingface.co/Cccccz/comparison/resolve/main/build_points.py
4.35 kB
| #!/usr/bin/env python | |
| """Turn the scans into the strategy files the evaluation reads. | |
| Two slots are added on top of the existing ~1.3x points: `FxxF` aimed at | |
| 1.70-1.80x and `Fxxx` aimed at 2.70-3.00x, both on *measured* wall-clock speedup. | |
| Where a method cannot reach its slot -- TeaCache because its reachable set is | |
| quantised, MotionCache because its per-token decision goes vacuous first -- the | |
| nearest setting that still behaves like the method is used and the achieved | |
| speedup is recorded, so the table can report what actually happened. | |
| """ | |
| import glob, json, os | |
| ROOT = os.path.dirname(os.path.abspath(__file__)) | |
| BANDS = {"FxxF": (1.70, 1.80, 1.75), "Fxxx": (2.70, 3.00, 2.85)} | |
| ACT_MIN, SEL_MIN = 0.25, 0.05 | |
| TOKENWISE = {"motioncache", "flowcache"} | |
| LONG = {"sf": "self_forcing", "cf": "causal_forcing", "hy": "hy_worldplay"} | |
| scans = {} | |
| for p in glob.glob(os.path.join(ROOT, "results/scan_*.json")): | |
| if p.endswith("_dirty.json"): | |
| continue | |
| d = json.load(open(p)) | |
| b = {"self_forcing": "sf", "causal_forcing": "cf", "hy_worldplay": "hy"}[d["base"]] | |
| scans.setdefault((b, d["method"], d["schedule"]), []).extend(d["rows"]) | |
| def choose(base, method, sched): | |
| lo, hi, _ = BANDS[sched] | |
| cand = scans.get((base, method, sched), []) | |
| if not cand: | |
| return None, "no scan" | |
| alive = ((lambda r: r["active_ratio"] >= ACT_MIN and r["sel_mean"] >= SEL_MIN) | |
| if method in TOKENWISE else (lambda r: True)) | |
| band = [r for r in cand if lo <= r["measured_speedup"] <= hi and alive(r)] | |
| if band: | |
| return max(band, key=lambda r: r["measured_speedup"]), "" | |
| live = [r for r in cand if alive(r) and r["measured_speedup"] <= hi] | |
| if live: | |
| r = max(live, key=lambda r: r["measured_speedup"]) | |
| why = ("token-wise decision goes vacuous above this point" | |
| if method in TOKENWISE else "reachable speedups are quantised") | |
| return r, f"below band ({why})" | |
| r = min(cand, key=lambda r: abs(r["measured_speedup"] - hi)) | |
| return r, "above band (reachable speedups are quantised)" | |
| rows = {"sf": [], "cf": [], "hy": []} | |
| print(f"{'base':>4} {'method':12s}{'sched':>6}{'param':>10}{'measured':>10}{'compute':>9}" | |
| f"{'active%':>9}{'sel%':>7} note") | |
| for base in ("sf", "cf", "hy"): | |
| for method in ("teacache", "taylorseer", "motioncache"): | |
| for sched in ("FxxF", "Fxxx"): | |
| # CF TeaCache already has a fully evaluated Fxxx point at 2.891x. | |
| if (base, method, sched) == ("cf", "teacache", "Fxxx"): | |
| print(f"{base:>4} {method:12s}{sched:>6}{1.70117:>10.5g}{2.891:>10.3f}" | |
| f"{'—':>9}{'—':>9}{'—':>7} keep existing cf_teacache_Fxxx_x3") | |
| continue | |
| r, note = choose(base, method, sched) | |
| if r is None: | |
| print(f"{base:>4} {method:12s}{sched:>6} {note}"); continue | |
| rows[base].append({ | |
| "base": LONG[base], "method": method, "param": r["param"], | |
| "value": r["value"], "schedule": sched, | |
| "target_speedup": BANDS[sched][2], | |
| "scan_measured_speedup": r["measured_speedup"], | |
| "scan_compute_speedup": r["compute_speedup"], | |
| "scan_compute_equivalent": r["compute_equivalent"], | |
| "scan_active_step_ratio": r["active_ratio"], | |
| "scan_mean_selected_fraction": r["sel_mean"], | |
| "band_note": note, | |
| "note": "operating point targeted on measured wall-clock speedup", | |
| }) | |
| print(f"{base:>4} {method:12s}{sched:>6}{r['value']:>10.5g}" | |
| f"{r['measured_speedup']:>10.3f}{r['compute_speedup']:>9.3f}" | |
| f"{100*r['active_ratio']:>9.1f}{100*r['sel_mean']:>7.1f} {note}") | |
| for b, long in (("sf", "self_forcing"), ("cf", "causal_forcing")): | |
| p = os.path.join(ROOT, f"results/final_{long}_v2.json") | |
| json.dump({"base": long, "source": "scan (measured-speedup targeting)", | |
| "rows": rows[b]}, open(p, "w"), indent=2) | |
| print("wrote", p, len(rows[b]), "rows") | |
| p = os.path.join(ROOT, "results/hy_point_v2.json") | |
| json.dump({"base": "hy_worldplay", "method": "multi", "param": "-", | |
| "schedule": "-", "operating_points": [], "rows": rows["hy"]}, | |
| open(p, "w"), indent=2) | |
| print("wrote", p, len(rows["hy"]), "rows") | |