summaryrefslogtreecommitdiff
path: root/scripts/plot_ladder.py
blob: 4473e02c9aaa48d205654be9022b3294e8b6a766 (plain)
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
"""Paper-side analysis from collected results: gap-vs-scale table + curves.
    python scripts/plot_ladder.py --results results/h200node1 --fig results/h200node1/fig_ladder.png"""
import os, json, glob, argparse
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt

p = argparse.ArgumentParser()
p.add_argument("--results", required=True)
p.add_argument("--fig", default=None)
a = p.parse_args()
summary = json.load(open(os.path.join(a.results, "summary.json")))
sizes = sorted({k.rsplit("_", 2)[0] if "_zbp_" in k else k.split("_")[0] for k in summary})
print(f"{'run':22s} {'params':>10s} {'steps':>7s} {'val_loss':>9s} {'tok/s':>9s}")
for k, v in sorted(summary.items()):
    print(f"{k:22s} {v['params'] or 0:>10d} {v['steps']:>7d} {v['val_loss'] or float('nan'):>9.3f} {v['tok_per_s'] or 0:>9.0f}")
fig, axes = plt.subplots(1, 2, figsize=(10, 3.8))
for k in sorted(summary):
    rows = [json.loads(l) for l in open(os.path.join(a.results, f"{k}.jsonl"))]
    ev = [(r["step"], r["val_loss"]) for r in rows if r.get("kind") == "eval" and r["step"] > 0]
    if ev:
        axes[0].plot(*zip(*ev), label=k, lw=1.4)
axes[0].set_xlabel("step"); axes[0].set_ylabel("val loss"); axes[0].legend(fontsize=6); axes[0].grid(alpha=0.3)
# gap vs scale: for each size, zbp - bp final
gaps = {}
for k, v in summary.items():
    if v["val_loss"] is None: continue
    size = k.split("_")[0]
    gaps.setdefault(size, {})[k[len(size) + 1:]] = (v["params"], v["val_loss"])
xs, arms = [], {}
for size, d in sorted(gaps.items(), key=lambda kv: (kv[1].get("bp") or (0, 0))[0]):
    if "bp" not in d: continue
    for arm, (params, loss) in d.items():
        if arm == "bp": continue
        arms.setdefault(arm, []).append((params, loss - d["bp"][1]))
for arm, pts in arms.items():
    pts.sort()
    axes[1].plot(*zip(*pts), "o-", label=f"{arm} − bp")
axes[1].set_xscale("log"); axes[1].axhline(0, color="k", lw=0.8)
axes[1].set_xlabel("params"); axes[1].set_ylabel("val-loss gap vs BP"); axes[1].legend(fontsize=7); axes[1].grid(alpha=0.3)
fig.tight_layout(); out = a.fig or os.path.join(a.results, "fig_ladder.png")
fig.savefig(out, dpi=150); print("saved", out)