Neural NetworksintermediateGPU optional~25 minColab

Grokking: the plateau your stopping rule would have believed

الاستيعاب المتأخّر: هضبة تخدع قاعدة التوقّف

The delay is a number you can measure, and early stopping fires inside it

You are training a small model on a task it can certainly learn. Training climbs to perfect within a few hundred steps and then stops moving. Held-out accuracy is not at chance — it is BELOW chance, because a model that has memorised is confidently wrong on what it has not seen rather than randomly wrong. Every instinct says the same thing: nothing is improving, kill the run and change something.

is what happens if you do not. Thousands of steps later — long past the point where a careful practitioner would have stopped — test accuracy leaves chance and climbs to near-perfect. That it happens is the paper's finding, and reading the paper is enough to learn it. What reading cannot give you is the size of the gap on your own machine, or the fact that an ordinary early-stopping rule, applied to your own logged curve, fires inside it and throws the result away.

The goal

Train a tiny transformer on modular addition until it memorises, keep training well past that point, and measure three things: the step at which it memorised, the step at which it generalised, and the distance between them. Then apply an ordinary early-stopping rule to your own logged curve and count the steps it would have discarded.

Colab opens a read-only copy. Save a copy to Drive to keep your edits.

The notebook needs a keyboard — best opened on a desktop.

The papers behind this

Almost nothing here is exotic. The task is addition modulo a prime, the model is one block, and the whole workshop is a single training run watched for longer than a training run is usually watched. The one value worth staring at is the , which is larger than the number most write-ups quote — at the usual value this same run memorises and then does nothing for forty thousand steps.

Read the curve in three parts. Training accuracy rises and saturates, and that is memorisation. Held-out accuracy drops to the floor and stays there, and that is the part that reads as failure. Then it moves. What you should look at while it is flat is the training , which is still falling the entire time — the run is not finished, it just stopped reporting anything on the axis you were watching.

Preflight. No accelerator is requested — this one is small enough that asking for a GPU would be theatre.

import azimuth_nb as azimuth

env = azimuth.setup(SLUG, lang=LANG, profile=PROFILE)
Workshop code
Grokking: the plateau your stopping rule would have believed
no GPU · 12.7 GB RAM · PyTorch 2.11.0+cpu
profile: free
ready · modulus=47, trainFrac=0.6, dModel=64, nHeads=4, dFF=256, steps=14000, evalEvery=100, logEvery=1000, learningRate=0.001, weightDecay=2.0, memoriseAt=0.99, generaliseAt=0.9, seed=17
code · 9fc593a2610292a0

Every pair the task admits is enumerated, then split once. Note the chance level printed here and keep it: it is what a model that knows nothing would score, and the held-out curve spends the plateau UNDER it. Reading that as a broken run is the first mistake this workshop exists to prevent.

import numpy as np

p = env.cfg["modulus"]

# Every pair the task admits, enumerated. This is what makes an algorithmic
# dataset useful here: "unseen" is exact rather than approximate, so the
# held-out number means what it says.
a_all, b_all = np.meshgrid(np.arange(p), np.arange(p), indexing="ij")
a_all, b_all = a_all.ravel(), b_all.ravel()

# Token p is the "=" that separates the operands from the answer position.
# Vocabulary is p operand tokens plus that one; the output layer predicts over
# the p possible answers only .
EQUALS = p
inputs = np.stack([a_all, b_all, np.full_like(a_all, EQUALS)], axis=1)
targets = (a_all + b_all) % p

rng = np.random.default_rng(env.cfg["seed"])

# THE SPLIT IS OVER UNORDERED PAIRS, NOT OVER ROWS.
#
# Addition commutes, so (a, b) and (b, a) have the same answer. Splitting the
# 3481 rows at random puts one ordering in training and the other in the
# held-out set for a large fraction of them, and a model that has merely
# MEMORISED the training row can answer its transpose. The first run of this
# workshop did exactly that: the held-out curve rose immediately to a plateau
# and sat there for thousands of steps, which read as partial generalization
# and was not.
#
# Grouping by the unordered pair and assigning whole groups keeps both
# orderings on the same side of the split. The `leakage` cell measures that it
# worked rather than assuming it.
lo, hi = np.minimum(a_all, b_all), np.maximum(a_all, b_all)
group = lo * p + hi

groups = rng.permutation(np.unique(group))
n_train_groups = round(env.cfg["trainFrac"] * len(groups))
train_groups = set(groups[:n_train_groups].tolist())
is_train = np.array([g in train_groups for g in group.tolist()])

train_in, test_in = inputs[is_train], inputs[~is_train]
train_out, test_out = targets[is_train], targets[~is_train]

n_train, n_test = len(train_in), len(test_in)

# Chance is one over the modulus, NOT zero. A held-out curve resting just
# above the axis is a model guessing uniformly, and mistaking that floor for
# zero makes the eventual jump look larger than it is.
chance = 1.0 / p

if env.lang == "ar":
    print(f"القياس {p} · {len(inputs)} زوجاً ممكناً")
    print(f"تدريب {n_train} · اختبار {n_test}")
    print(f"مستوى الصدفة {chance:.4f}")
else:
    print(f"modulus {p} · {len(inputs)} possible pairs")
    print(f"train {n_train} · test {n_test}")
    print(f"chance level {chance:.4f}")
Workshop code
modulus 47 · 2209 possible pairs
train 1323 · test 886
chance level 0.0213

Set arithmetic, no training. The first number must be zero. The second is what a random split over rows would have given away instead — and that fraction, not chance, is the height of the floor the held-out curve would have rested on while looking like slow progress.

# What the split gave away. Set arithmetic only — nothing trains here.
#
# `leaked_pairs` is how many held-out rows could be answered by transposing a
# row the model was trained on. It must be zero. `naive_leaked` is what a
# random split over rows would have handed over instead, and it is the height
# of the floor the held-out curve would have rested on: not chance, and not a
# result.
train_rows = set(zip(train_in[:, 0].tolist(), train_in[:, 1].tolist()))
test_rows = list(zip(test_in[:, 0].tolist(), test_in[:, 1].tolist()))
leaked_pairs = sum(1 for x, y in test_rows if (y, x) in train_rows)

naive_order = np.random.default_rng(env.cfg["seed"]).permutation(len(inputs))
naive_cut = int(env.cfg["trainFrac"] * len(inputs))
naive_train = set(
    zip(
        inputs[naive_order[:naive_cut], 0].tolist(),
        inputs[naive_order[:naive_cut], 1].tolist(),
    )
)
naive_test = list(
    zip(
        inputs[naive_order[naive_cut:], 0].tolist(),
        inputs[naive_order[naive_cut:], 1].tolist(),
    )
)
naive_leaked = sum(1 for x, y in naive_test if (y, x) in naive_train) / len(naive_test)

env.explain("leakage")
if env.lang == "ar":
    print(f"أزواج محجوزة يمكن الإجابة عنها بالتبديل: {leaked_pairs}")
    print(f"القسمة الساذجة على الصفوف كانت لتسلّم {naive_leaked:.3f} منها")
else:
    print(f"held-out pairs answerable by transposition: {leaked_pairs}")
    print(f"a naive split over rows would have handed over {naive_leaked:.3f} of them")

split_ok = env.check("split-is-clean", leaked_pairs)
Workshop code
leakage — Information reaching the held-out set from the training set. Here it arrives through a symmetry of the task rather than through a duplicated row, which is why the split looked random and was not.
held-out pairs answerable by transposition: 0
a naive split over rows would have handed over 0.578 of them
✓ Held-out pairs answerable by transposing a training pair: 0 (needs ≤ 0)

One block, and small enough to read in full. Compare the parameter count against the number of training pairs above — memorising them is well within reach, which is the whole reason the first plateau happens.

import torch
import torch.nn as nn

torch.manual_seed(env.cfg["seed"])
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")

SEQ_LEN = 3


class GrokFormer(nn.Module):
    """One pre-norm transformer block, reading the answer off the last position.

    Deliberately the smallest thing that groks rather than a scaled-down copy
    of anything. Depth is not the variable here — the training run is.
    """

    def __init__(self, vocab: int, d_model: int, n_heads: int, d_ff: int, n_answers: int):
        super().__init__()
        self.tok = nn.Embedding(vocab, d_model)
        self.pos = nn.Parameter(torch.randn(SEQ_LEN, d_model) * 0.02)
        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_ff),
            nn.GELU(),
            nn.Linear(d_ff, d_model),
        )
        self.unembed = nn.Linear(d_model, n_answers, bias=False)

    def forward(self, x):
        h = self.tok(x) + self.pos
        normed = self.ln1(h)
        h = h + self.attn(normed, normed, normed, need_weights=False)[0]
        h = h + self.mlp(self.ln2(h))
        return self.unembed(h[:, -1])


model = GrokFormer(
    vocab=p + 1,
    d_model=env.cfg["dModel"],
    n_heads=env.cfg["nHeads"],
    d_ff=env.cfg["dFF"],
    n_answers=p,
).to(device)

param_count = sum(w.numel() for w in model.parameters())

env.explain("memorization")
if env.lang == "ar":
    print(f"{param_count:,} معامل مقابل {n_train} زوج تدريب")
else:
    print(f"{param_count:,} parameters against {n_train} training pairs")
Workshop code
memorization — Fitting the training examples without learning a rule that covers unseen ones. Perfect on what it has seen, chance on everything else.
56,256 parameters against 1323 training pairs

The long part. Held-out accuracy will sit AT OR BELOW chance while this runs — a memorising model is confidently wrong on what it has not seen, not randomly wrong — so treat a figure under the chance line as the expected reading rather than a bug. After the jump, expect the occasional spike where training accuracy falls out of memorisation for an evaluation or two and the held-out figure drops with it; that is the decay and the loss trading places, and it is why the check below reads a peak rather than a last value.

import time

train_x = torch.from_numpy(train_in).long().to(device)
train_y = torch.from_numpy(train_out).long().to(device)
test_x = torch.from_numpy(test_in).long().to(device)
test_y = torch.from_numpy(test_out).long().to(device)

# Full batch. The dataset fits in memory many times over, and a fixed batch
# removes sampling noise from an axis where a long flat line is the evidence.
optimizer = torch.optim.AdamW(
    model.parameters(),
    lr=env.cfg["learningRate"],
    weight_decay=env.cfg["weightDecay"],
    betas=(0.9, 0.98),
)
criterion = nn.CrossEntropyLoss()


def accuracy(x, y) -> float:
    model.eval()
    with torch.no_grad():
        return float((model(x).argmax(dim=-1) == y).float().mean())


steps_log: list[int] = []
train_acc_log: list[float] = []
test_acc_log: list[float] = []
train_loss_log: list[float] = []

# The two crossings, recorded as they happen rather than reconstructed later:
# the first step at which the model has memorised, and the first at which it
# has generalised. Both stay None if the crossing never occurs, which is a
# distinct state from "crossed at step 0" and must not be flattened into it.
memorise_step = None
generalise_step = None

started = time.time()
model.train()
for step in range(1, env.cfg["steps"] + 1):
    optimizer.zero_grad()
    loss = criterion(model(train_x), train_y)
    loss.backward()
    optimizer.step()
    # .detach() before the scalar conversion: float() on a tensor that still
    # tracks gradients emits a UserWarning, and warnings printed inside a
    # training loop end up on the published page next to the numbers.
    loss_value = loss.detach().item()

    if step == 1 or step % env.cfg["evalEvery"] == 0:
        train_acc = accuracy(train_x, train_y)
        test_acc = accuracy(test_x, test_y)
        model.train()

        steps_log.append(step)
        train_acc_log.append(train_acc)
        test_acc_log.append(test_acc)
        train_loss_log.append(loss_value)

        if memorise_step is None and train_acc >= env.cfg["memoriseAt"]:
            memorise_step = step
        if generalise_step is None and test_acc >= env.cfg["generaliseAt"]:
            generalise_step = step

        if step == 1 or step % env.cfg["logEvery"] == 0:
            elapsed = time.time() - started
            print(
                f"step {step:6d}  train {train_acc:.3f}  test {test_acc:.3f}  "
                f"loss {loss_value:.5f}  {elapsed:5.0f}s"
            )

peak_train_acc = max(train_acc_log)
final_test_acc = test_acc_log[-1]
# The PEAK held-out accuracy, not the last one. At this decay the run keeps
# taking optimization spikes after it has generalised: training accuracy drops
# to around 0.92 for an evaluation or two and the held-out figure falls with
# it. Two measured seeds peaked at 0.937 and 0.917 and both happened to end
# near 0.90, so a check on the final value is really a check on where the
# oscillation was standing when the budget ran out.
peak_test_acc = max(test_acc_log)

env.explain("weight decay")
if env.lang == "ar":
    print(
        f"أعلى دقة تدريب {peak_train_acc:.3f} · أعلى دقة محجوزة {peak_test_acc:.3f} "
        f"· الدقة النهائية {final_test_acc:.3f}"
    )
    print(f"حفظ عند {memorise_step} · عمّم عند {generalise_step}")
else:
    print(
        f"peak train {peak_train_acc:.3f} · peak unseen {peak_test_acc:.3f} "
        f"· final unseen {final_test_acc:.3f}"
    )
    print(f"memorised at {memorise_step} · generalised at {generalise_step}")
Workshop code
step      1  train 0.020  test 0.018  loss 3.98689      0s
step   1000  train 1.000  test 0.000  loss 0.16864     68s
step   2000  train 1.000  test 0.002  loss 0.04802    134s
step   3000  train 1.000  test 0.000  loss 0.08134    199s
step   4000  train 1.000  test 0.025  loss 0.06125    265s
step   5000  train 1.000  test 0.335  loss 0.05982    331s
step   6000  train 1.000  test 0.416  loss 0.07515    398s
step   7000  train 1.000  test 0.710  loss 0.06052    465s
step   8000  train 1.000  test 0.841  loss 0.02337    532s
step   9000  train 1.000  test 0.896  loss 0.04213    599s
step  10000  train 1.000  test 0.921  loss 0.03648    666s
step  11000  train 0.833  test 0.676  loss 0.57979    733s
step  12000  train 1.000  test 0.921  loss 0.02582    799s
step  13000  train 1.000  test 0.930  loss 0.02328    867s
step  14000  train 1.000  test 0.923  loss 0.01406    934s
weight decay — A pull on every parameter toward zero, applied at each step. It prices complexity, and here it is what eventually makes the general solution cheaper than the memorised one.
peak train 1.000 · peak unseen 0.938 · final unseen 0.923
memorised at 300 · generalised at 9200

The accuracy axis is zero-based so that flat looks flat, and the step axis is logarithmic because the delay spans orders of magnitude. Look at the lower panel while the upper one is doing nothing.

import matplotlib.pyplot as plt

# A gap of 0 is what an incomplete run produces: either crossing missing means
# there is nothing to measure. Reporting 0 rather than skipping the number
# makes the `delay` check FAIL, which is the honest outcome — a run that never
# grokked should not certify a delay.
if memorise_step is not None and generalise_step is not None:
    grok_gap = generalise_step - memorise_step
else:
    grok_gap = 0

fig, (top, bottom) = plt.subplots(2, 1, figsize=(7, 5.4), sharex=True)

top.plot(steps_log, train_acc_log, label="train", color="#2a9d8f")
top.plot(steps_log, test_acc_log, label="held out", color="#e76f51")
top.axhline(chance, color="#8d99ae", linestyle=":", linewidth=1, label="chance")
# Zero-based: the long flat stretch is the argument, and an autoscaled axis
# would turn the noise on it into a trend.
top.set_ylim(0, 1.02)
top.set_ylabel("accuracy")
top.legend(loc="center left")
top.spines[["top", "right"]].set_visible(False)

bottom.plot(steps_log, train_loss_log, color="#264653")
bottom.set_yscale("log")
bottom.set_ylabel("train loss")
bottom.set_xlabel("optimizer step")
bottom.spines[["top", "right"]].set_visible(False)

for axis in (top, bottom):
    # Logarithmic in the step axis because the delay spans orders of
    # magnitude; on a linear axis the whole first plateau is one pixel wide.
    axis.set_xscale("log")
    if memorise_step is not None:
        axis.axvline(memorise_step, color="#2a9d8f", linewidth=1, alpha=0.5)
    if generalise_step is not None:
        axis.axvline(generalise_step, color="#e76f51", linewidth=1, alpha=0.5)

fig.tight_layout()
plt.show()

env.explain("grokking")
if env.lang == "ar":
    print(f"الفجوة {grok_gap} خطوة بين الحفظ والتعميم")
else:
    print(f"gap of {grok_gap} steps between memorising and generalising")
Workshop code
grokking — Generalization that arrives long after the training loss has stopped being interesting. A property of the training run, not of the final model.
gap of 8900 steps between memorising and generalising

Exercise

Pick a patience and see what it costs you. The rule below is the ordinary one: watch held-out accuracy, and stop when it has not improved for a given number of evaluations. Nothing retrains — it replays the curve you already have.

Choose the patience you would actually have used before you saw this plot, then read off the step it fires at and the steps between there and the jump. Then ask the harder question: what would you have had to watch instead for the rule to survive? A metric that was still moving during the plateau exists in this run, and it is on the lower panel.

# YOUR TURN.
#
# The ordinary rule: watch held-out accuracy, stop when it has not improved
# for PATIENCE consecutive evaluations. Nothing retrains — this replays the
# curve already logged above.
#
# Patience is counted in EVALUATIONS. Multiply by the eval interval to get the
# number of optimizer steps it actually buys you, which is usually smaller
# than it sounds.
PATIENCE = 20

best_acc = -1.0
stale = 0
stopped_at = steps_log[-1]
for step, acc in zip(steps_log, test_acc_log):
    if acc > best_acc:
        best_acc, stale = acc, 0
    else:
        stale += 1
        if stale >= PATIENCE:
            stopped_at = step
            break

patience_steps = PATIENCE * env.cfg["evalEvery"]
reached = generalise_step if generalise_step is not None else steps_log[-1]
discarded = max(0, reached - stopped_at)

if env.lang == "ar":
    print(f"صبر {PATIENCE} تقييماً = {patience_steps} خطوة")
    print(f"كانت القاعدة لتتوقف عند الخطوة {stopped_at} بدقة محجوزة {best_acc:.3f}")
    print(f"خطوات مهدورة قبل القفزة: {discarded}")
else:
    print(f"patience of {PATIENCE} evaluations = {patience_steps} steps")
    print(f"the rule would have stopped at step {stopped_at}, held-out {best_acc:.3f}")
    print(f"steps discarded before the jump: {discarded}")
Workshop code

A hint is available in the notebook — env.hint(4)

Memorisation is checked first and on purpose: without it there is no plateau to be surprised by, and a delay measured on a run that never learned anything would be a number about nothing.

# The control goes first. Without memorisation there is no plateau, and a
# delay measured on a run that never fit the training set would be a number
# about nothing.
memorises_ok = env.check("memorises", peak_train_acc)
generalises_ok = env.check("generalises", peak_test_acc)
delay_ok = env.check("delay", grok_gap)
Workshop code
✓ Peak accuracy on the training pairs: 1 (needs ≥ 0.99)
✓ Best accuracy reached on the held-out pairs: 0.9379 (needs ≥ 0.85)
✓ Optimizer steps between memorising and generalising: 8900 (needs ≥ 4000)
receipt = env.receipt()
Workshop code
Workshop complete.

Completion code: AZ-██████████
Paste it on the workshop's page on Azimuth to record it.
Last verified: 2026-08-29 · unknown · PyTorch unknown · Python 3.13.5 · 5a4cd21

Terms in this workshop