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)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 · 9fc593a2610292a0Every 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}")modulus 47 · 2209 possible pairs
train 1323 · test 886
chance level 0.0213Set 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)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")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 pairsThe 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}")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 9200The 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")
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 generalisingExercise
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}")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)✓ 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 complete.
Completion code: AZ-██████████
Paste it on the workshop's page on Azimuth to record it.Terms in this workshop
- Grokkingالاستيعاب المتأخّر
- Memorizationالحفظ
- Weight Decayاضمحلال الأوزان
- Early Stoppingالإيقاف المبكر للتدريب
- Modular Arithmeticالحساب المعياري