summaryrefslogtreecommitdiff
path: root/scripts/plot.py
blob: 61733eea258d68d1a7c8ec8342c269d70d04aa3d (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
50
51
52
53
54
"""Plot train/val curves from a metrics.csv. Usage: python -m scripts.plot <model> <seed> [--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()