Gala-598M-MLX / plot_loss.py
junafinity's picture
Gala-598M: weights, code, logs, report
76bbe95 verified
Raw History Blame Contribute Delete
2.48 kB
"""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()