"""Plot train/val curves from a metrics.csv. Usage: python -m scripts.plot [--runs-dir DIR]""" import argparse import csv import matplotlib matplotlib.use("Agg") import matplotlib.pyplot as plt def main() -> None: ap = argparse.ArgumentParser() ap.add_argument("model") ap.add_argument("seed") ap.add_argument("--runs-dir", default="runs") a = ap.parse_args() model, seed = a.model, a.seed path = f"{a.runs_dir}/{model}/seed{seed}/metrics.csv" out = f"{a.runs_dir}/{model}/seed{seed}/curves.png" with open(path) as fh: rows = list(csv.DictReader(fh)) steps = [int(r["step"]) for r in rows] train_loss = [float(r["train_loss"]) for r in rows] train_em = [float(r["train_em"]) for r in rows] val_em = [float(r["val_em"]) for r in rows] train_tok = [float(r["train_token_acc"]) for r in rows] val_tok = [float(r["val_token_acc"]) for r in rows] halt = [float(r["mean_halt_steps"]) for r in rows] fig, axes = plt.subplots(2, 2, figsize=(12, 8)) axes[0, 0].plot(steps, train_loss, label="train loss", color="tab:blue") axes[0, 0].set_title("Train loss (token CE + halt penalty)") axes[0, 0].set_xlabel("step") axes[0, 1].plot(steps, train_em, label="train EM", color="tab:orange") axes[0, 1].plot(steps, val_em, label="val EM", color="tab:green") axes[0, 1].axhline(0.9, ls="--", c="gray", lw=0.7) axes[0, 1].set_title(f"Exact-match (model={model}, seed={seed})") axes[0, 1].set_ylim(-0.05, 1.05) axes[0, 1].legend() axes[1, 0].plot(steps, train_tok, label="train tok acc", color="tab:red") axes[1, 0].plot(steps, val_tok, label="val tok acc", color="tab:purple") axes[1, 0].set_title("Token accuracy") axes[1, 0].set_ylim(-0.05, 1.05) axes[1, 0].legend() axes[1, 1].plot(steps, halt, label="mean halt steps", color="tab:brown") axes[1, 1].set_title("RNN halting (mean steps used)") axes[1, 1].set_xlabel("step") fig.suptitle(f"prime-grokking — {model} seed {seed}") plt.tight_layout() plt.savefig(out, dpi=110) print(f"saved {out}") if __name__ == "__main__": main()