Overfit a single batch first
Before training on everything, prove your pipeline can memorise eight samples — a healthy setup drives that loss to zero, and a setup that cannot has a bug, not a data problem.
- 7 min read
- 3 reading levels
- Published
Read these first
One lesson, three depths. Pick the one that fits you today — you can switch any time.
Beginner — No maths. Plain English.
The single most useful sanity check in deep learning: take one tiny batch of data and train on it, alone, until the model knows it perfectly.
Before cooking for a two-hundred-guest wedding, a sensible cook makes one plate. If one plate comes out wrong, the problem is the recipe — and no amount of bigger cooking will fix a recipe.
One batch is your one plate. A healthy model, optimiser and loss can memorise eight examples with ease. If yours cannot, something in the machinery is broken, and you have found this out in thirty seconds instead of three hours.
Why it exists
Training on full data mixes two questions that deserve separate answers. Is my machinery wired correctly? And is my task learnable from this data? When full-scale training disappoints, you cannot tell which question failed.
The one-batch test isolates the first question completely. Memorising eight examples needs no generalisation, no data quality, no clever features — nothing but working machinery. Failure here is a pure machinery verdict.
How it works
pick 8 samples, freeze them
train on those 8, over and over
│
├── loss dives toward zero → machinery works; scale up
└── loss stalls → bug in model/loss/optimiser wiring
(data volume is not the issue)Note the twist that confuses newcomers: everywhere else, memorisation — overfitting — is the disease. Here it is deliberately induced, as proof of health.
Where you have seen this
Sound engineers run a test tone before the concert. Pilots check control surfaces before taxiing. Every serious profession has a cheap pre-flight ritual that catches wiring faults before the expensive event. This is deep learning's.
Remember this
- One tiny batch, memorised perfectly = machinery certified.
- Failure to memorise 8 samples is a bug, never a data-quantity problem.
- Run this ritual before every long training, not after it disappoints.
What to learn next
- Tracking down a NaN loss — when the one-batch test explodes instead of stalling.
- Finding label and target bugs — the bugs this test is blind to.
- When the loss will not go down — the full checklist this ritual belongs to.
Developer — Code and libraries.
Setup
pip install torchCaptured with torch 2.5.1 on CPU; seeded and deterministic.
The ritual
import torch
import torch.nn as nn
torch.manual_seed(0)
xb = torch.randn(8, 10) # one small batch, frozen
yb = torch.randint(0, 3, (8,)) # 8 samples, 3 classes
model = nn.Sequential(nn.Linear(10, 64), nn.ReLU(), nn.Linear(64, 3))
optimizer = torch.optim.Adam(model.parameters(), lr=1e-3)
loss_fn = nn.CrossEntropyLoss()
for step in range(401):
optimizer.zero_grad()
loss = loss_fn(model(xb), yb)
loss.backward()
optimizer.step()
if step % 100 == 0:
acc = (model(xb).argmax(dim=1) == yb).float().mean()
print(f"step {step:3d} loss {loss.item():.5f} batch accuracy {acc:.0%}")step 0 loss 1.04707 batch accuracy 62% step 100 loss 0.12591 batch accuracy 100% step 200 loss 0.02483 batch accuracy 100% step 300 loss 0.01023 batch accuracy 100% step 400 loss 0.00558 batch accuracy 100%
That is what health looks like: loss in freefall toward zero, accuracy pinned at 100% within a hundred steps. The labels here are random — there is no pattern to learn — and the model memorises them anyway, which is the whole point. Only the machinery is being examined.
Reading the verdict
Pass: loss heads for zero (below ~0.01 for a batch this size), accuracy hits 100%. Machinery certified. Scale up.
Fail — loss stuck at the guessing floor (about 1.10 for 3 classes, from the previous lesson): the model's output ignores its input, or the optimiser is disconnected. Run the weights-moved test.
Fail — loss falls, then flattens well above zero: gradients flow but something limits expressiveness — a stray softmax before CrossEntropyLoss, a bottleneck layer of width 1, targets mangled by a bad cast.
Fail — loss explodes: learning rate, or a NaN factory in a custom layer.
Each failure mode now reproduces in seconds, on 8 samples, with no DataLoader, no GPU queueing, no hour-long feedback loop. That speed is the entire value.
Rules for a fair test
Freeze the batch. Sample it once, outside the loop. Drawing fresh random data each step turns the test into real (impossible) learning — the classic way this ritual gets accidentally invalidated.
Disable regularisation. Dropout and augmentation exist to prevent memorisation, and this test demands memorisation. Turn them off, or run the model in .eval()-style configuration for the test. Weight decay at default strength is usually harmless here; heavy decay is not.
Keep the batch genuinely small. Eight to thirty-two samples. Memorising 4,096 samples is legitimately hard and muddies the verdict.
Expect near-perfect, not "pretty good". 87% accuracy on 8 memorisable samples is a fail. The bar is deliberately absolute; softness here defeats the diagnostic.
Common mistakes
Judging by accuracy alone. Accuracy saturates early (62% before any training here, by luck). The loss falling toward zero is the real certificate — it proves confidence, not lucky argmaxes.
Skipping the ritual because the model "worked last time". The test certifies the current wiring: this refactor, this new loss, this new data format. Machinery bugs arrive precisely with changes.
Running it with the full validation apparatus attached. Early stopping watching a validation set will halt your memorisation run. Strip the test to its bones: one batch, one loop, one printout.
Concluding data is fine because the test passed. The pass certifies machinery only. Label alignment, leakage and class balance live in their own lesson — the one-batch test cannot see them.
Try it yourself
Sabotage the run three ways and watch the three failure signatures: add nn.Softmax(dim=1) as a final layer; change the hidden width from 64 to 1; set lr=1.0. Reading those three curves is the skill this lesson exists to build.
What to learn next
- Tracking down a NaN loss — when the one-batch test explodes instead of stalling.
- Finding label and target bugs — the bugs this test is blind to.
- When the loss will not go down — the full checklist this ritual belongs to.
Researcher — Mathematics and papers.
Why memorisation is the right null test
Zhang et al. (2017), Understanding Deep Learning Requires Rethinking Generalization, showed standard architectures fit random labels to zero training error — memorisation demands no structure in the data whatsoever. This makes one-batch overfitting a clean necessary condition: any correctly wired (model, loss, optimiser) triple of adequate capacity must pass, independent of task learnability. Failure therefore implicates the wiring with high specificity. Arpit et al. (2017), A Closer Look at Memorization in Deep Networks, refine the picture — networks learn patterns before noise — which is why real-data one-batch runs typically converge even faster than the random-label bound.
Capacity arithmetic
Interpolating $n$ points with $C$ classes needs parameter count comfortably above $n$; the demo's 771-parameter MLP against 8 points is overparameterised by two orders of magnitude, deliberately. In the overparameterised regime, gradient descent on such tiny problems enjoys strong convergence guarantees (Du et al., 2019, Gradient Descent Provably Optimizes Over-parameterized Neural Networks; NTK analysis, Jacot et al., 2018) — theory agreeing with practice that a stall cannot be blamed on optimisation difficulty at this scale.
Provenance and tooling
The ritual is folklore made canon by Karpathy (2019), A Recipe for Training Neural Networks — step 2 of the recipe, "overfit one batch", sandwiched between data audit and regularisation. Framework support acknowledges its status: Lightning ships Trainer(overfit_batches=1) and fast_dev_run, both encountered in the Lightning lesson.
What the test cannot certify, formally
The test evaluates the composite $\text{loss} \circ \text{model} \circ \text{batch}$ on a fixed batch; anything outside that composite escapes it: dataset-wide label permutations (the batch's own labels are consistent within the test), train/val leakage (see train-test split for the discipline), distribution shift, and DataLoader-level bugs (shuffling, collation, worker seeding) that only manifest across batches. A disciplined escalation: one batch → two batches → one epoch on 1% of data → full data, each step adding one new subsystem to the certified set.
What to learn next
- Tracking down a NaN loss — when the one-batch test explodes instead of stalling.
- Finding label and target bugs — the bugs this test is blind to.
- When the loss will not go down — the full checklist this ritual belongs to.