summaryrefslogtreecommitdiff
path: root/scripts/plot.py
blob: 14536e514cdae5f91444b64d341337a3c101830e (plain)
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
"""Plot train/val curves from a metrics.csv. Usage: python -m scripts.plot <model> <seed>"""
import csv
import sys

import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt


def main() -> None:
    model, seed = sys.argv[1], sys.argv[2]
    path = f"runs/{model}/seed{seed}/metrics.csv"
    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()
    out = f"runs/{model}/seed{seed}/curves.png"
    plt.savefig(out, dpi=110)
    print(f"saved {out}")


if __name__ == "__main__":
    main()