File size: 2,484 Bytes
76bbe95
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
55
56
57
58
59
60
61
62
63
64
65
66
67
68
"""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()