Downloads · 30 days
0
BurnyCoder/grokking-modular-addition-transformer
grokking-modular-addition-transformer is a other model from BurnyCoder. Use it for the other task on the model card, and read the license before you ship it in a product. It is set up for transformer-lens. The card lists the license as mit.
A 1-layer transformer trained on modular addition (a + b) mod 113 that exhibits grokking -- the phenomenon where the model first memorizes the training data, then suddenly generalizes to the test set after continued t…
Downloads · 30 days
0
Access
Public
Updated Mar 1, 2026
Repo size
230 MB
Likes
0
Public
Click a slice to open those files.
.pth230 MB · 100%
From the Hugging Face model README
A 1-layer transformer trained on modular addition (a + b) mod 113 that exhibits grokking -- the phenomenon where the model first memorizes the training data, then suddenly generalizes to the test set after continued training.
This model is a reproduction of the setup from Progress Measures for Grokking via Mechanistic Interpretability (Nanda et al., 2023), built with TransformerLens.
The model learns a Fourier-based algorithm to perform modular addition:
a and b into Fourier components (sin/cos at key frequencies)= position to a and b, computing sin(ka), cos(ka), sin(kb), cos(kb)cos(k(a+b)) and sin(k(a+b)) via trigonometric identitiescos(k(a+b-c)) for each candidate output c| Parameter | Value |
|---|---|
| Layers | 1 |
| Attention heads | 4 |
| d_model | 128 |
| d_head | 32 |
| d_mlp | 512 |
| Activation | ReLU |
| Normalization | None |
| Vocabulary (input) | 114 (0-112 for numbers, 113 for =) |
| Vocabulary (output) | 113 |
| Context length | 3 tokens: [a, b, =] |
| Parameters | ~2.5M |
Design choices (no LayerNorm, ReLU, no biases) were made to simplify mechanistic interpretability analysis.
import torch
from transformer_lens import HookedTransformer
# Download and load
cached_data = torch.load("grokking_demo.pth", weights_only=False)
model = HookedTransformer(cached_data["config"])
model.load_state_dict(cached_data["model"])
# Training history is also included
model_checkpoints = cached_data["checkpoints"] # 250 intermediate checkpoints
checkpoint_epochs = cached_data["checkpoint_epochs"] # Every 100 epochs
train_losses = cached_data["train_losses"]
test_losses = cached_data["test_losses"]
train_indices = cached_data["train_indices"]
test_indices = cached_data["test_indices"]
import torch
p = 113
a, b = 37, 58
input_tokens = torch.tensor([[a, b, p]]) # [a, b, =]
logits = model(input_tokens)
prediction = logits[0, -1].argmax().item()
print(f"{a} + {b} mod {p} = {prediction}") # Should print 95
pip install torch transformer-lens
| Setting | Value |
|---|---|
| Task | (a + b) mod 113 |
| Total data | 113^2 = 12,769 pairs |
| Train split | 30% (3,830 examples) |
| Test split | 70% (8,939 examples) |
| Optimizer | AdamW |
| Learning rate | 1e-3 |
| Weight decay | 1.0 |
| Betas | (0.9, 0.98) |
| Epochs | 25,000 |
| Batch size | Full batch |
| Checkpoints | Every 100 epochs (250 total) |
| Seed | 999 (model), 598 (data split) |
| Training time | ~2 minutes on GPU |
The high weight decay (1.0) is critical for grokking -- it gradually erodes memorization weights in favor of the compact generalizing Fourier circuit.
The training exhibits three distinct phases:
Analysis of the trained model reveals:
cos(freq * 2pi/p * (a + b - c)) for key frequenciesFull analysis notebook and training code: GitHub repository