diff options
| author | yurenh <blackhao0426@gmail.com> | 2026-08-31 18:34:12 -0500 |
|---|---|---|
| committer | yurenh <blackhao0426@gmail.com> | 2026-08-31 18:34:12 -0500 |
| commit | b270eb58e22deb9f6a1c5342db41d531232ded0d (patch) | |
| tree | 9829b9b32e7c76b0ce6f8863c25eb58f581f66a0 /scripts | |
| parent | f68f6b00f047d9412113d7c2627bd7f6e96f2d62 (diff) | |
results flow: collect.py (node-side bundle) + plot_ladder.py (gap-vs-scale analysis)
Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_01GkgLsACEF6CCP7EUfA5fZe
Diffstat (limited to 'scripts')
| -rw-r--r-- | scripts/collect.py | 27 | ||||
| -rw-r--r-- | scripts/plot_ladder.py | 42 |
2 files changed, 69 insertions, 0 deletions
diff --git a/scripts/collect.py b/scripts/collect.py new file mode 100644 index 0000000..c943385 --- /dev/null +++ b/scripts/collect.py @@ -0,0 +1,27 @@ +"""Collect run logs into a small committable bundle. +On the training node after (some) runs finish: + python scripts/collect.py --runs runs --out results/h200node1 && git add results && git commit -m "results" && git push +Checkpoints stay on the node (runs/ is gitignored); only JSONL summaries+curves are collected.""" +import os, json, glob, argparse, shutil + +p = argparse.ArgumentParser() +p.add_argument("--runs", default="runs") +p.add_argument("--out", default="results/local") +a = p.parse_args() +os.makedirs(a.out, exist_ok=True) +summary = {} +for d in sorted(glob.glob(os.path.join(a.runs, "*"))): + name = os.path.basename(d) + lp = os.path.join(d, "log_rank0.jsonl") + if not os.path.exists(lp): + continue + rows = [json.loads(l) for l in open(lp)] + meta = next((r for r in rows if r.get("kind") == "meta"), {}) + evals = [r for r in rows if r.get("kind") == "eval"] + shutil.copy(lp, os.path.join(a.out, f"{name}.jsonl")) + summary[name] = {"params": meta.get("params"), "world": meta.get("world"), "mode": (meta.get("cfg") or {}).get("mode"), + "n_probes": (meta.get("cfg") or {}).get("n_probes"), "done": os.path.exists(os.path.join(d, "DONE")), + "steps": evals[-1]["step"] if evals else 0, "val_loss": evals[-1]["val_loss"] if evals else None, + "tok_per_s": evals[-1].get("tok_per_s") if evals else None} +json.dump(summary, open(os.path.join(a.out, "summary.json"), "w"), indent=1) +print(json.dumps(summary, indent=1)) diff --git a/scripts/plot_ladder.py b/scripts/plot_ladder.py new file mode 100644 index 0000000..4473e02 --- /dev/null +++ b/scripts/plot_ladder.py @@ -0,0 +1,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) |
