Skip to content

Latest commit

 

History

1 Commit

Folders and files

NameName
Last commit message
Last commit date
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 

Repository files navigation

grokking: a from-scratch reproduction

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.

Reading order

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

Results at a glance

(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.

Repository layout

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)

Setup

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.

Training a model

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.

Watching it happen

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.

Analysing a finished run

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.

Reproducing the full experiment suite

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

Tests

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.

What this does not reproduce

  • 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 on Z_p, so the progress measures in docs/04 would 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/06 says where ours differ and why.

License

MIT. See LICENSE.

About

A documented from-scratch reproduction of grokking (Power et al. 2022), with mechanistic progress measures and a live training dashboard.

Resources

Stars

0 stars

Watchers

0 watching

Forks

Releases

Packages

Contributors

Languages