|
| 1 | +#!/usr/bin/env python3 |
| 2 | +"""Paper-B figure: the size-quality frontier + the subset-overstatement |
| 3 | +pairs, from the durable verdicts (final_eval.json per run). Fetches the |
| 4 | +small JSON files from Modal volumes; no GPU. |
| 5 | +
|
| 6 | + python scripts/figures/frontier.py --out docs/paper-assets/ |
| 7 | +""" |
| 8 | +from __future__ import annotations |
| 9 | + |
| 10 | +import argparse |
| 11 | +import json |
| 12 | +import subprocess |
| 13 | +import tempfile |
| 14 | +from pathlib import Path |
| 15 | + |
| 16 | +RUNS = [ |
| 17 | + # (label, params_m, full_set_der, subset_der, volume, path) |
| 18 | + ("1.0 AdamW/r6/3ep", 300, 8.2590, 3.658, "rababa-checkpoints", |
| 19 | + "rababa_arabic_distill_small/run-002"), |
| 20 | + ("lite Muon/6ep", 190, 5.784, 2.6495, "rababa-checkpoints", |
| 21 | + "rababa_arabic_distill_small/run-009-layerdrop-6ep"), |
| 22 | + ("2.0 Muon/r7/3ep", 300, 4.8218, None, "rababa-checkpoints", |
| 23 | + "rababa_arabic_distill_small/run-006-r7-muon"), |
| 24 | + ("2.1 Muon/r7/6ep", 300, 4.5701, 2.0062, "rababa-checkpoints", |
| 25 | + "rababa_arabic_distill_small/run-007-r7-muon-6ep"), |
| 26 | + ("teacher r7", 580, 2.2864, None, "rababa-checkpoints", None), |
| 27 | + ("tiny-max (30M full levers)", 33, 73.9489, None, "rababa-checkpoints", |
| 28 | + "rababa_arabic_distill_small/run-010-tiny-max"), |
| 29 | +] |
| 30 | + |
| 31 | + |
| 32 | +def fetch_eval(volume: str, path: str | None) -> dict | None: |
| 33 | + if not path: |
| 34 | + return None |
| 35 | + with tempfile.NamedTemporaryFile(suffix=".json", delete=False) as f: |
| 36 | + out = f.name |
| 37 | + r = subprocess.run( |
| 38 | + ["modal", "volume", "get", volume, f"{path}/final_eval.json", out], |
| 39 | + capture_output=True, text=True) |
| 40 | + if r.returncode != 0: |
| 41 | + return None |
| 42 | + try: |
| 43 | + return json.loads(Path(out).read_text()) |
| 44 | + except Exception: |
| 45 | + return None |
| 46 | + |
| 47 | + |
| 48 | +def main() -> None: |
| 49 | + ap = argparse.ArgumentParser() |
| 50 | + ap.add_argument("--out", default="docs/paper-assets") |
| 51 | + args = ap.parse_args() |
| 52 | + out = Path(args.out) |
| 53 | + out.mkdir(parents=True, exist_ok=True) |
| 54 | + |
| 55 | + import matplotlib |
| 56 | + |
| 57 | + matplotlib.use("Agg") |
| 58 | + import matplotlib.pyplot as plt |
| 59 | + |
| 60 | + # Figure 1: the frontier (params vs full-set DER) |
| 61 | + fig, ax = plt.subplots(figsize=(6, 4)) |
| 62 | + pts = [(r[1], r[2]) for r in RUNS if r[2] is not None] |
| 63 | + labels = [r[0] for r in RUNS if r[2] is not None] |
| 64 | + ax.plot([p[0] for p in pts], [p[1] for p in pts], "o-") |
| 65 | + for (x, y), lab in zip(pts, labels): |
| 66 | + ax.annotate(lab, (x, y), fontsize=7, xytext=(4, 4), |
| 67 | + textcoords="offset points") |
| 68 | + ax.set_xscale("log") |
| 69 | + ax.set_yscale("log") |
| 70 | + ax.set_xlabel("parameters (M, log)") |
| 71 | + ax.set_ylabel("full-set windowed DER-CE (log)") |
| 72 | + ax.set_title("The client-tier size–quality frontier (all full-set, CI-bracketed)") |
| 73 | + fig.tight_layout() |
| 74 | + fig.savefig(out / "frontier.png", dpi=200) |
| 75 | + |
| 76 | + # Figure 2: subset overstatement pairs |
| 77 | + pairs = [(r[0], r[3], r[2]) for r in RUNS if r[3] is not None] |
| 78 | + fig2, ax2 = plt.subplots(figsize=(6, 4)) |
| 79 | + idx = range(len(pairs)) |
| 80 | + w = 0.35 |
| 81 | + ax2.bar([i - w / 2 for i in idx], [p[1] for p in pairs], w, label="first-300 subset") |
| 82 | + ax2.bar([i + w / 2 for i in idx], [p[2] for p in pairs], w, label="full 1,200") |
| 83 | + ax2.set_xticks(list(idx)) |
| 84 | + ax2.set_xticklabels([p[0] for p in pairs], fontsize=7, rotation=20) |
| 85 | + ax2.set_ylabel("DER-CE") |
| 86 | + ax2.set_title("Subset overstatement: 2–4× inflation across five instances") |
| 87 | + ax2.legend() |
| 88 | + fig2.tight_layout() |
| 89 | + fig2.savefig(out / "subset-overstatement.png", dpi=200) |
| 90 | + print(f"wrote {out}/frontier.png and {out}/subset-overstatement.png") |
| 91 | + |
| 92 | + |
| 93 | +if __name__ == "__main__": |
| 94 | + main() |
0 commit comments