A reproduction of Power, Burda, Edwards, Babuschkin & Misra, "Grokking: Generalization Beyond Overfitting on Small Algorithmic Datasets" (arXiv:2201.02177, 2022) -- the paper that showed a network can sit at chance-level validation accuracy for tens of thousands of steps after perfectly fitting its training set, and then suddenly generalise.
This repository implements the setup from scratch in PyTorch -- the modular-arithmetic datasets, the ~400k-parameter transformer, full-batch AdamW training, and the mechanistic progress measures from the Nanda et al. (2023) follow-up -- and runs the whole experiment suite (26 runs) on a single consumer GPU in 57 minutes.
It is written as a guided, documented reproduction: docs/ walks through the
ideas in the order they are implemented, and ships with a live dashboard
that streams the run as it trains, because the phenomenon is fundamentally
about time and a static plot understates it.
| Doc | Covers |
|---|---|
docs/01-background.md |
What grokking is, what the paper claims, and the mechanism established afterwards |
docs/02-task-and-data.md |
Modular-arithmetic tables, tokenisation, and why the train/test split makes this a clean probe |
docs/03-model-and-optimisation.md |
The transformer, full-batch AdamW, why weight decay is the experiment, log-spaced logging |
docs/04-progress-measures.md |
Fourier concentration, restricted/excluded loss, and the three phases underneath the jump |
docs/05-dashboard.md |
The live UI and the append-only-file streaming contract behind it |
docs/06-results.md |
Full results, compared point-by-point against the paper's claims |
(a + b) mod 97, 30% of the table used for training, 418,816-parameter
transformer, full-batch AdamW with weight_decay = 1.0, single RTX 5070:
| Event | Step |
|---|---|
| Train accuracy >= 99% (memorised) | 98 |
| Test accuracy >= 95% (grokked) | 7,656 |
| Grokking gap | 78x |
Final test accuracy 100.0%, in 101 seconds of wall clock. The identical run
with weight_decay = 0 never groks: 100% train accuracy, 3.1% test
accuracy after 25,000 steps.
The two models are qualitatively different objects, not two points on a quality scale. With key frequencies fixed from the final model:
wd = 1.0 (grokked) |
wd = 0.0 (memorised) |
|
|---|---|---|
| Power in top 5 of 48 Fourier frequencies | 94.2% | 15.2% (uniform ≈ 10.4%) |
| Logit variance explained by the answer alone | 97.5% | 3.7% |
And the jump is not sudden underneath: restricted loss overtakes excluded loss at step ~4,600, while test accuracy still reads 20%.
See docs/06-results.md for the training-fraction sweep
(28x change in grokking step from a 2.8x change in data), the weight-decay
sweep (grokking step scales roughly as 1/wd), all seven binary operations, and
a point-by-point comparison against the paper's claims.
On error bars. Grokking step is noisy. Across three initialisation seeds it ranges 7,401–13,375 (±25%), and three argument-for-argument identical runs in the suite gave 6,575 / 7,155 / 7,401 (±6%, from GPU floating-point non-determinism). Read every step count in this repository as accurate to about ±25%; the orders of magnitude are what reproduce, not the digits.
paper/ the source paper (PDF)
docs/ concept notes, paired with the implementation, read in order
src/grokking/ library code
data.py modular-arithmetic tables, tokenisation, train/test split
model.py the ~400k-param decoder-only transformer
metrics.py Fourier progress measures, restricted/excluded loss
train.py full-batch AdamW training loop with log-spaced logging
analysis.py post-hoc mechanistic analysis of a finished run
dashboard/ single-file live dashboard (no build step, no dependencies)
scripts/ CLI entry points
download_paper.py fetch the source PDF from arXiv
serve_dashboard.py the live dashboard server
run_experiments.py the full 26-run experiment suite
analyze.py post-hoc mechanistic analysis of one run
phase_table.py print the three-phase decomposition as a table
tests/ unit tests for every module above
results/ results.json / results.csv from the suite (runs/ is gitignored)
Requires Python >= 3.10 and (recommended) a CUDA-capable GPU -- everything runs on CPU too, roughly 20x slower.
pip install -e .
python scripts/download_paper.py
There is no dataset to download. The task is a 97x97 arithmetic table generated in memory.
python -m grokking.train # canonical run
python -m grokking.train --weight_decay 0 --name no_wd # the ablation
Key flags: --op (which binary operation), --p (modulus), --train_frac
(fraction of table cells used for training -- the knob that controls how long
the plateau lasts), --weight_decay, --steps, --seed, --evals
(log-spaced evaluation points), --layernorm (use the paper's original
architecture). Defaults match the configuration that groks reliably: 2-layer
transformer, d_model 128, 4 heads, AdamW at lr = 1e-3 with
weight_decay = 1.0 and betas = (0.9, 0.98), full batch, 25,000 steps.
python scripts/serve_dashboard.py # http://localhost:8000
Run this in a second terminal, then start a training run. The page streams metrics live: accuracy on a log x axis with the grokking gap shaded, loss, the Fourier progress measures, the embedding frequency spectrum (scrub with the mouse to watch it go from noise to five sharp spikes), and the parameter norm. Tick two runs to overlay them.
Details in docs/05-dashboard.md.
python scripts/analyze.py results/runs/expA_canonical
Recomputes the restricted/excluded progress measures over the run's saved
checkpoints with the key frequencies held fixed at the final model's -- the
faithful version of Nanda et al.'s measures, which train.py can only
approximate online. Also reports how much of the model's logit variance is
explained by a function of the correct answer alone, a mechanism-agnostic test
of "did it learn the rule or memorise the table". Writes analysis.json, which
the dashboard picks up.
python scripts/run_experiments.py
Runs all five experiments described in docs/06-results.md
-- the headline run and weight-decay ablation, the training-fraction sweep, the
weight-decay sweep, all seven binary operations, and a seed-variance check --
and writes results/results.json / results/results.csv. 26 runs, 57 minutes
on a single RTX 5070. Already-completed configurations are reused unless you
pass --force; --only A,B restricts which experiments run.
To print the three-phase decomposition quoted in docs/06:
python scripts/phase_table.py results/runs/expA_canonical
pytest tests/
30 tests covering table construction and the train/test split (disjointness,
coverage, that div really inverts mul), the model (parameter count,
causal masking, gradient flow to every parameter), the progress measures
(basis orthonormality, 2D Fourier round-trip, and a synthetic logit grid built
to be the trig algorithm at a known frequency, where restricted and excluded
loss have exactly known values), the analysis measures, and end-to-end training
artifacts.
- The permutation-group datasets. The paper also studies composition in the
symmetric group S5, which is non-abelian and has no Fourier basis over
Z_p. Everything here is modular arithmetic onZ_p, so the progress measures indocs/04would not transfer without being rebuilt on group representations. - The full regulariser comparison. The paper compares several regularisers and reports weight decay as the most effective. This repository sweeps weight decay thoroughly (Experiment C) but does not implement the alternatives it is being compared against.
- The paper's exact step counts. Absolute grokking steps depend on the
operation, modulus, training fraction, initialisation scale and optimiser
settings, several of which the paper does not fully specify. What reproduces
is the shape -- the order-of-magnitude gap, its growth as data shrinks, and
its dependence on weight decay -- not the specific numbers.
docs/06says where ours differ and why.
MIT. See LICENSE.