"""Plot train/val loss curves from one or more runs' log.txt. python plot_loss.py runs/hoard_small_* runs/transformer_small_* --out compare.png """ import argparse, json, os, re import matplotlib matplotlib.use("Agg") import matplotlib.pyplot as plt STEP_RE = re.compile(r"step\s+(\d+) \| loss ([\d.]+) ce ([\d.]+).*?\| ([\d,]+) tok/s") VAL_RE = re.compile(r"val loss ([\d.]+) @ step (\d+) \(([\d.]+)M tokens\)") def parse(run_dir): steps, ce, val_steps, val_loss, val_tokens = [], [], [], [], [] tok_s = [] with open(os.path.join(run_dir, "log.txt")) as f: for line in f: m = STEP_RE.search(line) if m: steps.append(int(m.group(1))) ce.append(float(m.group(3))) tok_s.append(float(m.group(4).replace(",", ""))) m = VAL_RE.search(line) if m: val_loss.append(float(m.group(1))) val_steps.append(int(m.group(2))) val_tokens.append(float(m.group(3))) cfg = json.load(open(os.path.join(run_dir, "config.json"))) bt = cfg["args"].get("batch_tokens", 65536) return {"steps": steps, "ce": ce, "tok_s": tok_s, "val_steps": val_steps, "val_loss": val_loss, "val_tokens": val_tokens, "batch_tokens": bt, "name": cfg["args"]["config"]} def main(): ap = argparse.ArgumentParser() ap.add_argument("runs", nargs="+") ap.add_argument("--out", default="compare.png") ap.add_argument("--title", default="val loss vs tokens") a = ap.parse_args() fig, ax = plt.subplots(figsize=(8, 5), dpi=150) colors = {"hoard_small": "#1E7A64", "gdn_hybrid_small": "#4C6EF5", "transformer_small": "#A8862B", "hoard_m3": "#C2452D"} for rd in a.runs: r = parse(rd) col = colors.get(r["name"]) # train ce vs tokens (light), val loss vs tokens (solid) train_tokens = [s * r["batch_tokens"] / 1e6 for s in r["steps"]] ax.plot(train_tokens, r["ce"], alpha=0.25, lw=1, color=col) if r["val_tokens"]: ax.plot(r["val_tokens"], r["val_loss"], marker="o", ms=3, lw=1.8, label=f"{r['name']} (val {r['val_loss'][-1]:.3f})", color=col) ax.set_xlabel("tokens seen (M)") ax.set_ylabel("cross-entropy (nats)") ax.set_title(a.title) ax.grid(alpha=0.3) ax.legend() fig.tight_layout() fig.savefig(a.out) print("wrote", a.out) if __name__ == "__main__": main()