summaryrefslogtreecommitdiff
path: root/scripts
diff options
context:
space:
mode:
Diffstat (limited to 'scripts')
-rw-r--r--scripts/collect.py27
-rw-r--r--scripts/plot_ladder.py42
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)