Grokking on modular addition¶

Reproduce grokking — generalization long after overfitting — on (a + b) mod p, and ablate weight decay to test whether it drives the effect. We compare a 1-layer Transformer and a 2-layer MLP, each trained with weight decay on (wd=1.0) and off (wd=0.0).

In [1]:
import json, os, time, math
from pathlib import Path
import numpy as np
import torch
import torch.nn as nn
import torch.nn.functional as F
%matplotlib inline
import matplotlib
# matplotlib.use("Agg")   # ← 빼기 (headless 실행 때만)
import matplotlib.pyplot as plt

# ----- config -----
P            = 97          # modulus (prime); dataset is all p*p pairs
TRAIN_FRAC   = 0.5         # fraction of pairs used for training
STEPS        = 30000       # full-batch optimizer steps per run
EVAL_EVERY   = 100
LR           = 1e-3
BETAS        = (0.9, 0.98)
WDS          = [1.0, 0.0]            # weight-decay ablation
MODELS       = ["transformer", "mlp"]
SEED         = 0
D_MODEL      = 128

FIG_DIR = Path("figures"); FIG_DIR.mkdir(exist_ok=True)
RES_DIR = Path("results"); RES_DIR.mkdir(exist_ok=True)

device = "cuda" if torch.cuda.is_available() else "cpu"
print("device:", device, "| torch:", torch.__version__,
      "| gpu:", torch.cuda.get_device_name(0) if torch.cuda.is_available() else "-")
device: cuda | torch: 2.4.1+cu124 | gpu: NVIDIA RTX A6000
In [2]:
# ----- data: all (a,b) pairs, target (a+b) % P -----
def make_data(p, train_frac, seed):
    a = torch.arange(p).repeat_interleave(p)      # (p*p,)
    b = torch.arange(p).repeat(p)                  # (p*p,)
    y = (a + b) % p
    eq = torch.full_like(a, p)                     # "=" token id = p
    x_tf = torch.stack([a, b, eq], dim=1)          # transformer input (N,3)
    x_mlp = torch.stack([a, b], dim=1)             # mlp input (N,2)
    g = torch.Generator().manual_seed(seed)
    perm = torch.randperm(p * p, generator=g)
    n_tr = int(train_frac * p * p)
    tr, va = perm[:n_tr], perm[n_tr:]
    return {
        "tf":  (x_tf[tr].to(device),  x_tf[va].to(device)),
        "mlp": (x_mlp[tr].to(device), x_mlp[va].to(device)),
        "y":   (y[tr].to(device),     y[va].to(device)),
    }

DATA = make_data(P, TRAIN_FRAC, SEED)
print("train:", DATA["y"][0].shape[0], "| val:", DATA["y"][1].shape[0])
train: 4704 | val: 4705
In [3]:
# ----- models -----
class Block(nn.Module):
    def __init__(self, d_model, n_heads, d_mlp):
        super().__init__()
        self.ln1 = nn.LayerNorm(d_model)
        self.attn = nn.MultiheadAttention(d_model, n_heads, batch_first=True)
        self.ln2 = nn.LayerNorm(d_model)
        self.mlp = nn.Sequential(nn.Linear(d_model, d_mlp), nn.GELU(),
                                 nn.Linear(d_mlp, d_model))

    def forward(self, x):
        h = self.ln1(x)
        a, _ = self.attn(h, h, h, need_weights=False)
        x = x + a
        x = x + self.mlp(self.ln2(x))
        return x

class Transformer(nn.Module):
    def __init__(self, p, d_model=128, n_heads=4, d_mlp=512, n_ctx=3):
        super().__init__()
        self.tok_emb = nn.Embedding(p + 1, d_model)   # +1 for the "=" token
        self.pos_emb = nn.Embedding(n_ctx, d_model)
        self.block = Block(d_model, n_heads, d_mlp)
        self.ln = nn.LayerNorm(d_model)
        self.unembed = nn.Linear(d_model, p, bias=False)

    def forward(self, x):                              # x: (B, n_ctx)
        pos = torch.arange(x.shape[1], device=x.device)
        h = self.tok_emb(x) + self.pos_emb(pos)[None]
        h = self.ln(self.block(h))                     # layer norm
        return self.unembed(h[:, -1])                  # logits at "=" position (B, p)

class MLP(nn.Module):
    def __init__(self, p, d_model=128, hidden=256):
        super().__init__()
        self.emb = nn.Embedding(p, d_model)
        self.net = nn.Sequential(nn.Linear(2 * d_model, hidden), nn.ReLU(),
                                 nn.Linear(hidden, p))

    def forward(self, ab):                             # ab: (B, 2) indices
        e = self.emb(ab).reshape(ab.shape[0], -1)      # (B, 2*d_model)
        return self.net(e)

def build(kind):
    return (Transformer(P, D_MODEL) if kind == "transformer"
            else MLP(P, D_MODEL)).to(device)
In [4]:
# ----- full-batch training with periodic eval -----
def evaluate(model, X, y):
    model.eval()
    with torch.no_grad():
        logits = model(X)
        loss = F.cross_entropy(logits, y).item()
        acc = (logits.argmax(-1) == y).float().mean().item()
    return loss, acc

def train_run(kind, wd, seed=SEED):
    torch.manual_seed(seed)
    model = build(kind)
    opt = torch.optim.AdamW(model.parameters(), lr=LR, betas=BETAS, weight_decay=wd)
    Xtr, Xva = DATA["tf"] if kind == "transformer" else DATA["mlp"]
    ytr, yva = DATA["y"]

    hist = {"step": [], "train_loss": [], "train_acc": [], "val_loss": [], "val_acc": []}
    t0 = time.time()
    for step in range(STEPS + 1):
        model.train()
        loss = F.cross_entropy(model(Xtr), ytr)
        opt.zero_grad(); loss.backward(); opt.step()
        if step % EVAL_EVERY == 0:
            tr_loss, tr_acc = evaluate(model, Xtr, ytr)
            va_loss, va_acc = evaluate(model, Xva, yva)
            hist["step"].append(step)
            hist["train_loss"].append(tr_loss); hist["train_acc"].append(tr_acc)
            hist["val_loss"].append(va_loss);   hist["val_acc"].append(va_acc)
            if step % (EVAL_EVERY * 10) == 0:
                print(f"[{kind} wd={wd}] step {step:6d}  "
                      f"train_acc {tr_acc:.3f}  val_acc {va_acc:.3f}  "
                      f"val_loss {va_loss:.3f}  ({time.time()-t0:.0f}s)", flush=True)
    return hist
In [5]:
# ----- run the 2x2 grid (model x weight decay) -----
results = {}
for kind in MODELS:
    for wd in WDS:
        print(f"=== {kind}  weight_decay={wd} ===", flush=True)
        results[f"{kind}|wd{wd}"] = train_run(kind, wd)

with open(RES_DIR / "metrics.json", "w") as f:
    json.dump({"config": {"P": P, "TRAIN_FRAC": TRAIN_FRAC, "STEPS": STEPS,
                          "LR": LR, "WDS": WDS, "MODELS": MODELS, "SEED": SEED},
               "results": results}, f)
print("saved results/metrics.json", flush=True)
=== transformer  weight_decay=1.0 ===
[transformer wd=1.0] step      0  train_acc 0.012  val_acc 0.010  val_loss 4.655  (1s)
[transformer wd=1.0] step   1000  train_acc 1.000  val_acc 0.367  val_loss 2.967  (8s)
[transformer wd=1.0] step   2000  train_acc 1.000  val_acc 1.000  val_loss 0.038  (15s)
[transformer wd=1.0] step   3000  train_acc 0.999  val_acc 0.997  val_loss 0.174  (22s)
[transformer wd=1.0] step   4000  train_acc 1.000  val_acc 1.000  val_loss 0.015  (29s)
[transformer wd=1.0] step   5000  train_acc 1.000  val_acc 1.000  val_loss 0.016  (36s)
[transformer wd=1.0] step   6000  train_acc 0.997  val_acc 0.997  val_loss 0.116  (43s)
[transformer wd=1.0] step   7000  train_acc 1.000  val_acc 1.000  val_loss 0.017  (50s)
[transformer wd=1.0] step   8000  train_acc 1.000  val_acc 1.000  val_loss 0.015  (57s)
[transformer wd=1.0] step   9000  train_acc 1.000  val_acc 1.000  val_loss 0.042  (64s)
[transformer wd=1.0] step  10000  train_acc 1.000  val_acc 1.000  val_loss 0.017  (72s)
[transformer wd=1.0] step  11000  train_acc 1.000  val_acc 1.000  val_loss 0.014  (79s)
[transformer wd=1.0] step  12000  train_acc 1.000  val_acc 1.000  val_loss 0.016  (86s)
[transformer wd=1.0] step  13000  train_acc 1.000  val_acc 1.000  val_loss 0.028  (93s)
[transformer wd=1.0] step  14000  train_acc 1.000  val_acc 1.000  val_loss 0.028  (101s)
[transformer wd=1.0] step  15000  train_acc 1.000  val_acc 1.000  val_loss 0.012  (108s)
[transformer wd=1.0] step  16000  train_acc 1.000  val_acc 1.000  val_loss 0.011  (115s)
[transformer wd=1.0] step  17000  train_acc 1.000  val_acc 1.000  val_loss 0.015  (122s)
[transformer wd=1.0] step  18000  train_acc 1.000  val_acc 1.000  val_loss 0.013  (129s)
[transformer wd=1.0] step  19000  train_acc 1.000  val_acc 1.000  val_loss 0.033  (137s)
[transformer wd=1.0] step  20000  train_acc 1.000  val_acc 1.000  val_loss 0.017  (144s)
[transformer wd=1.0] step  21000  train_acc 1.000  val_acc 1.000  val_loss 0.013  (151s)
[transformer wd=1.0] step  22000  train_acc 1.000  val_acc 1.000  val_loss 0.014  (158s)
[transformer wd=1.0] step  23000  train_acc 1.000  val_acc 1.000  val_loss 0.011  (166s)
[transformer wd=1.0] step  24000  train_acc 1.000  val_acc 1.000  val_loss 0.017  (173s)
[transformer wd=1.0] step  25000  train_acc 1.000  val_acc 1.000  val_loss 0.011  (180s)
[transformer wd=1.0] step  26000  train_acc 1.000  val_acc 1.000  val_loss 0.011  (187s)
[transformer wd=1.0] step  27000  train_acc 1.000  val_acc 1.000  val_loss 0.006  (195s)
[transformer wd=1.0] step  28000  train_acc 1.000  val_acc 1.000  val_loss 0.016  (202s)
[transformer wd=1.0] step  29000  train_acc 1.000  val_acc 1.000  val_loss 0.011  (209s)
[transformer wd=1.0] step  30000  train_acc 1.000  val_acc 1.000  val_loss 0.011  (217s)
=== transformer  weight_decay=0.0 ===
[transformer wd=0.0] step      0  train_acc 0.012  val_acc 0.010  val_loss 4.655  (0s)
[transformer wd=0.0] step   1000  train_acc 1.000  val_acc 0.108  val_loss 9.840  (7s)
[transformer wd=0.0] step   2000  train_acc 1.000  val_acc 0.123  val_loss 11.230  (15s)
[transformer wd=0.0] step   3000  train_acc 1.000  val_acc 0.116  val_loss 11.820  (22s)
[transformer wd=0.0] step   4000  train_acc 1.000  val_acc 0.139  val_loss 10.457  (29s)
[transformer wd=0.0] step   5000  train_acc 1.000  val_acc 0.159  val_loss 9.811  (36s)
[transformer wd=0.0] step   6000  train_acc 1.000  val_acc 0.155  val_loss 11.476  (44s)
[transformer wd=0.0] step   7000  train_acc 1.000  val_acc 0.182  val_loss 9.233  (51s)
[transformer wd=0.0] step   8000  train_acc 1.000  val_acc 0.195  val_loss 8.899  (58s)
[transformer wd=0.0] step   9000  train_acc 1.000  val_acc 0.210  val_loss 8.564  (66s)
[transformer wd=0.0] step  10000  train_acc 1.000  val_acc 0.225  val_loss 8.274  (73s)
[transformer wd=0.0] step  11000  train_acc 1.000  val_acc 0.236  val_loss 8.054  (80s)
[transformer wd=0.0] step  12000  train_acc 1.000  val_acc 0.246  val_loss 7.920  (88s)
[transformer wd=0.0] step  13000  train_acc 1.000  val_acc 0.256  val_loss 7.703  (95s)
[transformer wd=0.0] step  14000  train_acc 1.000  val_acc 0.270  val_loss 7.550  (102s)
[transformer wd=0.0] step  15000  train_acc 1.000  val_acc 0.279  val_loss 7.420  (109s)
[transformer wd=0.0] step  16000  train_acc 1.000  val_acc 0.284  val_loss 7.296  (117s)
[transformer wd=0.0] step  17000  train_acc 1.000  val_acc 0.291  val_loss 7.172  (124s)
[transformer wd=0.0] step  18000  train_acc 1.000  val_acc 0.298  val_loss 7.106  (131s)
[transformer wd=0.0] step  19000  train_acc 1.000  val_acc 0.305  val_loss 6.999  (138s)
[transformer wd=0.0] step  20000  train_acc 1.000  val_acc 0.312  val_loss 6.953  (146s)
[transformer wd=0.0] step  21000  train_acc 1.000  val_acc 0.314  val_loss 6.842  (153s)
[transformer wd=0.0] step  22000  train_acc 1.000  val_acc 0.323  val_loss 6.782  (160s)
[transformer wd=0.0] step  23000  train_acc 1.000  val_acc 0.329  val_loss 6.731  (168s)
[transformer wd=0.0] step  24000  train_acc 1.000  val_acc 0.332  val_loss 6.711  (175s)
[transformer wd=0.0] step  25000  train_acc 1.000  val_acc 0.338  val_loss 6.632  (182s)
[transformer wd=0.0] step  26000  train_acc 1.000  val_acc 0.342  val_loss 6.578  (189s)
[transformer wd=0.0] step  27000  train_acc 1.000  val_acc 0.346  val_loss 6.562  (197s)
[transformer wd=0.0] step  28000  train_acc 1.000  val_acc 0.351  val_loss 6.521  (204s)
[transformer wd=0.0] step  29000  train_acc 1.000  val_acc 0.351  val_loss 6.503  (211s)
[transformer wd=0.0] step  30000  train_acc 1.000  val_acc 0.355  val_loss 6.485  (218s)
=== mlp  weight_decay=1.0 ===
[mlp wd=1.0] step      0  train_acc 0.011  val_acc 0.008  val_loss 4.608  (0s)
[mlp wd=1.0] step   1000  train_acc 1.000  val_acc 0.097  val_loss 6.397  (2s)
[mlp wd=1.0] step   2000  train_acc 1.000  val_acc 0.880  val_loss 0.445  (3s)
[mlp wd=1.0] step   3000  train_acc 1.000  val_acc 0.998  val_loss 0.013  (5s)
[mlp wd=1.0] step   4000  train_acc 1.000  val_acc 1.000  val_loss 0.000  (8s)
[mlp wd=1.0] step   5000  train_acc 1.000  val_acc 1.000  val_loss 0.000  (10s)
[mlp wd=1.0] step   6000  train_acc 1.000  val_acc 1.000  val_loss 0.000  (13s)
[mlp wd=1.0] step   7000  train_acc 1.000  val_acc 1.000  val_loss 0.000  (16s)
[mlp wd=1.0] step   8000  train_acc 1.000  val_acc 1.000  val_loss 0.000  (19s)
[mlp wd=1.0] step   9000  train_acc 1.000  val_acc 1.000  val_loss 0.000  (21s)
[mlp wd=1.0] step  10000  train_acc 1.000  val_acc 1.000  val_loss 0.000  (24s)
[mlp wd=1.0] step  11000  train_acc 1.000  val_acc 1.000  val_loss 0.000  (27s)
[mlp wd=1.0] step  12000  train_acc 1.000  val_acc 1.000  val_loss 0.000  (30s)
[mlp wd=1.0] step  13000  train_acc 1.000  val_acc 1.000  val_loss 0.000  (32s)
[mlp wd=1.0] step  14000  train_acc 1.000  val_acc 1.000  val_loss 0.000  (35s)
[mlp wd=1.0] step  15000  train_acc 1.000  val_acc 1.000  val_loss 0.000  (38s)
[mlp wd=1.0] step  16000  train_acc 1.000  val_acc 1.000  val_loss 0.000  (40s)
[mlp wd=1.0] step  17000  train_acc 1.000  val_acc 1.000  val_loss 0.000  (43s)
[mlp wd=1.0] step  18000  train_acc 1.000  val_acc 1.000  val_loss 0.000  (46s)
[mlp wd=1.0] step  19000  train_acc 1.000  val_acc 1.000  val_loss 0.000  (49s)
[mlp wd=1.0] step  20000  train_acc 1.000  val_acc 1.000  val_loss 0.000  (51s)
[mlp wd=1.0] step  21000  train_acc 1.000  val_acc 1.000  val_loss 0.000  (54s)
[mlp wd=1.0] step  22000  train_acc 1.000  val_acc 1.000  val_loss 0.000  (57s)
[mlp wd=1.0] step  23000  train_acc 1.000  val_acc 1.000  val_loss 0.000  (60s)
[mlp wd=1.0] step  24000  train_acc 1.000  val_acc 1.000  val_loss 0.000  (62s)
[mlp wd=1.0] step  25000  train_acc 1.000  val_acc 1.000  val_loss 0.000  (65s)
[mlp wd=1.0] step  26000  train_acc 1.000  val_acc 1.000  val_loss 0.000  (68s)
[mlp wd=1.0] step  27000  train_acc 1.000  val_acc 1.000  val_loss 0.000  (71s)
[mlp wd=1.0] step  28000  train_acc 1.000  val_acc 1.000  val_loss 0.000  (74s)
[mlp wd=1.0] step  29000  train_acc 1.000  val_acc 1.000  val_loss 0.000  (76s)
[mlp wd=1.0] step  30000  train_acc 1.000  val_acc 1.000  val_loss 0.000  (79s)
=== mlp  weight_decay=0.0 ===
[mlp wd=0.0] step      0  train_acc 0.011  val_acc 0.008  val_loss 4.608  (0s)
[mlp wd=0.0] step   1000  train_acc 1.000  val_acc 0.001  val_loss 23.260  (3s)
[mlp wd=0.0] step   2000  train_acc 1.000  val_acc 0.003  val_loss 25.002  (5s)
[mlp wd=0.0] step   3000  train_acc 1.000  val_acc 0.004  val_loss 28.883  (8s)
[mlp wd=0.0] step   4000  train_acc 1.000  val_acc 0.005  val_loss 35.339  (11s)
[mlp wd=0.0] step   5000  train_acc 1.000  val_acc 0.008  val_loss 40.488  (13s)
[mlp wd=0.0] step   6000  train_acc 1.000  val_acc 0.016  val_loss 43.192  (16s)
[mlp wd=0.0] step   7000  train_acc 1.000  val_acc 0.012  val_loss 42.784  (19s)
[mlp wd=0.0] step   8000  train_acc 1.000  val_acc 0.016  val_loss 42.765  (22s)
[mlp wd=0.0] step   9000  train_acc 1.000  val_acc 0.018  val_loss 43.063  (25s)
[mlp wd=0.0] step  10000  train_acc 1.000  val_acc 0.017  val_loss 43.349  (27s)
[mlp wd=0.0] step  11000  train_acc 1.000  val_acc 0.021  val_loss 43.769  (30s)
[mlp wd=0.0] step  12000  train_acc 1.000  val_acc 0.019  val_loss 44.437  (33s)
[mlp wd=0.0] step  13000  train_acc 1.000  val_acc 0.021  val_loss 43.787  (35s)
[mlp wd=0.0] step  14000  train_acc 1.000  val_acc 0.021  val_loss 43.919  (38s)
[mlp wd=0.0] step  15000  train_acc 1.000  val_acc 0.021  val_loss 44.257  (41s)
[mlp wd=0.0] step  16000  train_acc 1.000  val_acc 0.023  val_loss 45.076  (43s)
[mlp wd=0.0] step  17000  train_acc 1.000  val_acc 0.022  val_loss 44.577  (46s)
[mlp wd=0.0] step  18000  train_acc 1.000  val_acc 0.024  val_loss 44.562  (49s)
[mlp wd=0.0] step  19000  train_acc 1.000  val_acc 0.025  val_loss 45.031  (51s)
[mlp wd=0.0] step  20000  train_acc 1.000  val_acc 0.025  val_loss 46.175  (54s)
[mlp wd=0.0] step  21000  train_acc 1.000  val_acc 0.027  val_loss 44.962  (57s)
[mlp wd=0.0] step  22000  train_acc 1.000  val_acc 0.026  val_loss 45.412  (59s)
[mlp wd=0.0] step  23000  train_acc 1.000  val_acc 0.026  val_loss 46.006  (62s)
[mlp wd=0.0] step  24000  train_acc 1.000  val_acc 0.027  val_loss 45.799  (65s)
[mlp wd=0.0] step  25000  train_acc 1.000  val_acc 0.028  val_loss 45.889  (68s)
[mlp wd=0.0] step  26000  train_acc 1.000  val_acc 0.029  val_loss 46.450  (70s)
[mlp wd=0.0] step  27000  train_acc 0.530  val_acc 0.028  val_loss 70.992  (73s)
[mlp wd=0.0] step  28000  train_acc 1.000  val_acc 0.032  val_loss 46.191  (76s)
[mlp wd=0.0] step  29000  train_acc 1.000  val_acc 0.034  val_loss 47.065  (78s)
[mlp wd=0.0] step  30000  train_acc 1.000  val_acc 0.031  val_loss 47.854  (81s)
saved results/metrics.json
In [6]:
# ----- figures -----
def xs(h):  # avoid step 0 on log axis
    s = np.array(h["step"]); s[0] = 1; return s

# 1) grokking curves (wd=1.0): train vs val accuracy, log-x
fig, axes = plt.subplots(1, 2, figsize=(11, 4.2), sharey=True)
for ax, kind in zip(axes, MODELS):
    h = results[f"{kind}|wd1.0"]
    ax.plot(xs(h), h["train_acc"], label="train", color="tab:blue")
    ax.plot(xs(h), h["val_acc"], label="val", color="tab:red")
    ax.set_xscale("log"); ax.set_xlabel("optimizer step"); ax.set_title(kind)
    ax.set_ylim(-0.02, 1.02); ax.grid(alpha=0.3)
axes[0].set_ylabel("accuracy"); axes[0].legend()
fig.suptitle("Grokking on (a+b) mod %d  (weight decay = 1.0)" % P)
fig.tight_layout(); fig.savefig(FIG_DIR / "grokking_curve.png", dpi=130); plt.close(fig)

# 2) weight-decay ablation: val accuracy, wd on vs off
fig, axes = plt.subplots(1, 2, figsize=(11, 4.2), sharey=True)
for ax, kind in zip(axes, MODELS):
    for wd, c in zip(WDS, ["tab:green", "tab:gray"]):
        h = results[f"{kind}|wd{wd}"]
        ax.plot(xs(h), h["val_acc"], label=f"wd={wd}", color=c)
    ax.set_xscale("log"); ax.set_xlabel("optimizer step"); ax.set_title(kind)
    ax.set_ylim(-0.02, 1.02); ax.grid(alpha=0.3)
axes[0].set_ylabel("val accuracy"); axes[0].legend()
fig.suptitle("Weight-decay ablation: generalization vs no generalization")
fig.tight_layout(); fig.savefig(FIG_DIR / "wd_ablation.png", dpi=130); plt.close(fig)

# 3) loss curves (wd=1.0): train vs val loss, log-x
fig, axes = plt.subplots(1, 2, figsize=(11, 4.2))
for ax, kind in zip(axes, MODELS):
    h = results[f"{kind}|wd1.0"]
    ax.plot(xs(h), h["train_loss"], label="train", color="tab:blue")
    ax.plot(xs(h), h["val_loss"], label="val", color="tab:red")
    ax.set_xscale("log"); ax.set_yscale("log")
    ax.set_xlabel("optimizer step"); ax.set_title(kind); ax.grid(alpha=0.3)
axes[0].set_ylabel("loss"); axes[0].legend()
fig.suptitle("Loss: validation overfits, then drops at grokking  (weight decay = 1.0)")
fig.tight_layout(); fig.savefig(FIG_DIR / "loss_curves.png", dpi=130); plt.close(fig)

print("saved figures:", [p.name for p in FIG_DIR.glob("*.png")])
saved figures: ['grokking_curve.png', 'wd_ablation.png', 'loss_curves.png']