What changing batch size actually changes
Batch size sets how jagged your curve looks, how many steps an epoch contains and how far one epoch travels — so two runs with different batch sizes cannot be compared on an epoch axis at all.
- 12 min read
- 3 reading levels
- Published
Read these first
On this page 5
One lesson, three depths. Pick the one that fits you today — you can switch any time.
Beginner — No maths. Plain English.
Batch size is how many examples the model sees before it adjusts itself once — and changing it reshapes your curve on its own.
Rolling chapatis for dinner works either way. You can roll one, put it on the tawa, taste it, adjust the flour, and repeat. Or you can roll twenty, cook them together, taste, and adjust once.
Tasting one chapati tells you something noisy — maybe that one was thin. Tasting twenty gives a steadier verdict. But you also get far fewer chances to adjust the dough.
Batch size is the number of chapatis per round. It trades how reliable each correction is against how many corrections you get.
Why this had to be invented
Early on, people computed the correction using every training example at once. Reliable, and painfully slow — a single adjustment required reading the entire dataset.
The alternative was one example at a time. Fast adjustments, but each one pulled in a slightly wrong direction, because one example is not the world.
Batches are the compromise, and almost all training uses them. Once you pick a batch size, three things move together, and confusing them is where the trouble starts.
How it works
SMALL BATCH (say 8) LARGE BATCH (say 512)
many corrections per pass few corrections per pass
each correction is noisy each correction is steady
curve looks hairy and jagged curve looks smooth and calm
one pass through the data = one pass through the data =
MANY adjustments FEW adjustmentsHere is the trap. Most people plot loss against epochs — one epoch being one full pass through the training data.
But an epoch at batch size 8 contains 64 times more adjustments than an epoch at batch size 512. The same "epoch 10" on two charts means two completely different amounts of learning.
Comparing two batch sizes on an epoch axis is like comparing two journeys by counting fuel stops instead of kilometres.
A real example you have seen
Think of checking exam papers. Mark one paper, then immediately change how strictly you mark. Your standard jumps around, because one weak paper makes you think the whole batch is weak.
Mark fifty, then adjust once. Your standard is steadier and fairer. You also correct your marking far less often across the whole pile. Neither approach is wrong; they produce different-looking behaviour from the same examiner.
Remember this
- Batch size sets how jagged the curve looks. Small batch, hairy curve. Large batch, smooth curve.
- Batch size sets how many adjustments an epoch contains.
- Never compare two batch sizes on an epoch axis. Compare on steps, on examples seen, or on wall-clock time.
What to learn next
- Finding a learning rate — the knob that must move whenever batch size does.
- Gradient accumulation — keeping a large effective batch on a small GPU.
- How smoothing and log scales mislead you — the other way a curve's shape lies about the model.
Developer — Code and libraries.
Setup
pip install torch numpyVerified with torch 2.5.1 (CPU), numpy 1.26.4, Python 3.10. Runs in about a minute on a laptop CPU. Data is generated in memory; nothing is downloaded.
Three batch sizes, one recipe
Everything is held fixed except the batch size — same seed, same learning rate, same network, same 25 epochs.
import numpy as np, torch, torch.nn as nn
def make_data(n, seed):
g = torch.Generator().manual_seed(seed)
X = torch.randn(n, 8, generator=g)
y = (X[:, 0]*X[:, 1] + 0.4*X[:, 2] + 0.3*torch.randn(n, generator=g) > 0).float()
return X, y[:, None]
Xtr, ytr = make_data(2048, 0)
Xva, yva = make_data(4000, 999)
def train(bs, lr=0.05, epochs=25, seed=0):
torch.manual_seed(seed)
net = nn.Sequential(nn.Linear(8, 32), nn.ReLU(), nn.Linear(32, 1))
opt, loss_fn = torch.optim.SGD(net.parameters(), lr=lr, momentum=0.9), nn.BCEWithLogitsLoss()
step_losses, steps, at200 = [], 0, None
for _ in range(epochs):
perm = torch.randperm(len(Xtr))
for i in range(0, len(Xtr), bs):
idx = perm[i:i+bs]
opt.zero_grad(); l = loss_fn(net(Xtr[idx]), ytr[idx]); l.backward(); opt.step()
step_losses.append(l.item()); steps += 1
if steps == 200: # a fair, batch-independent checkpoint
with torch.no_grad(): at200 = loss_fn(net(Xva), yva).item()
with torch.no_grad():
at_end = loss_fn(net(Xva), yva).item()
return steps, np.array(step_losses), at200, at_end
print("batch steps/epoch total steps jitter 1/sqrt(B) rule val @200 steps val @25 epochs")
base = None
for bs in [8, 32, 128]:
steps, sl, a200, aend = train(bs)
j = sl[len(sl)*3//4:].std() # jitter measured on the last quarter only
if base is None: base = (bs, j)
pred = base[1] * (base[0] / bs) ** 0.5
print(f"{bs:5d} {steps//25:11d} {steps:11d} {j:.4f} {pred:13.4f} {a200:14.4f} {aend:14.4f}")batch steps/epoch total steps jitter 1/sqrt(B) rule val @200 steps val @25 epochs
8 256 6400 0.2274 0.2274 0.3574 0.3579
32 64 1600 0.0880 0.1137 0.3034 0.3102
128 16 400 0.0385 0.0568 0.2949 0.3010The walkthrough
jitter is the standard deviation of the per-step training loss, measured over the last quarter of training so the early descent does not pollute it. It is the number behind "the curve looks hairy". It falls from 0.2274 to 0.0385 as the batch grows sixteenfold — and the model is identical in all three rows.
The 1/sqrt(B) column is theory, and the data beats it. Averaging $B$ independent gradient estimates should shrink their spread by $\sqrt{B}$, giving 0.2274 → 0.1137 → 0.0568. Measured jitter falls faster than that. The reason is that the rule assumes the model stands still; at batch size 8 the weights also move 256 times per epoch, and that movement adds extra variation to the plotted loss. Treat the rule as an order-of-magnitude guide, not a prediction.
Read the last two columns side by side. They are the whole lesson. Judged at 200 steps, the three runs are 0.3574 / 0.3034 / 0.2949 — but the batch-128 run has seen 16 times more data to get there. Judged at 25 epochs, they are 0.3579 / 0.3102 / 0.3010 — but the batch-8 run took 6,400 steps against 400. Neither axis is neutral. Each one hands the win to a different side.
Wall-clock time is the third axis, and it is the one that pays your bill. It depends entirely on hardware: on a GPU a batch of 128 often costs barely more than a batch of 8 per step, so large batches finish an epoch far faster; on this CPU example the gap is much smaller. Measure it on your machine rather than trusting any published number — timing GPU code correctly explains why naive timings are usually wrong.
Batch size and learning rate are not independent. Raise batch size by a factor $k$, and the usual starting rule raises the learning rate by $k$ (linear scaling) or by $\sqrt{k}$, plus a warmup. Change one without the other and you will misread a learning-rate problem as a batch-size result — see finding a learning rate.
What to do about it
- Plot against steps or examples seen, not epochs, whenever batch size is one of the things you vary. Keep the epoch axis for runs that share a batch size.
- State the batch size next to every curve. A curve without it cannot be interpreted, and cannot be reproduced.
- When memory forces a smaller batch, use gradient accumulation to keep the effective batch size fixed. The curve then stays comparable with your earlier runs, which is worth more than it sounds.
- Change batch size and learning rate as a pair, and record both in the run config.
Common mistakes
"Batch size 512 trains better, look at the epoch chart." It had 16 times more data per epoch. Fix the axis before the conclusion.
Raising batch size to fix a jagged curve, then celebrating. The jagged curve becomes smooth because the measurement is smoother, not because learning improved. Compare validation loss at matched steps before believing anything.
Forgetting the last, smaller batch. With 2,048 rows and batch 128 the split is exact. With 2,050 rows the final batch holds 2 examples, whose loss is wildly noisy. drop_last=True in a DataLoader removes the artefact.
Batch-norm layers changing behaviour with batch size. Batch normalisation computes statistics within a batch, so tiny batches give unstable statistics and genuinely worse models — not a plotting artefact but a real effect. BatchNorm in PyTorch covers the group-norm alternatives.
Assuming a bigger batch is always faster overall. Past the point where the hardware is saturated, doubling the batch roughly doubles the time per step, and you gain nothing but a smoother chart. Is the GPU waiting for data? shows how to check where the time is actually going.
Try it yourself
Rerun with lr scaled linearly against batch size — train(8, lr=0.05), train(32, lr=0.2), train(128, lr=0.8) — and compare the val @200 steps column against the table above. Then predict what happens to the jitter column, and check whether you were right. Most people guess wrong, because the learning rate scales the steps, not the gradient noise.
What to learn next
- Finding a learning rate — the knob that must move whenever batch size does.
- Gradient accumulation — keeping a large effective batch on a small GPU.
- How smoothing and log scales mislead you — the other way a curve's shape lies about the model.
Researcher — Mathematics and papers.
The gradient as an estimator
A mini-batch gradient is an unbiased estimator of the full-batch gradient. For a batch $\mathcal{B}$ of size $B$ drawn uniformly with replacement from a training set of size $N$,
$$ g_{\mathcal{B}} = \frac{1}{B}\sum_{i \in \mathcal{B}} \nabla \ell_i(\theta), \qquad \mathbb{E}[g_{\mathcal{B}}] = \nabla \hat{R}(\theta), \qquad \operatorname{Cov}(g_{\mathcal{B}}) = \frac{\Sigma(\theta)}{B} $$
where $\ell_i$ is the loss on example $i$, $\hat{R}$ the empirical risk, and $\Sigma(\theta)$ the per-example gradient covariance. The $1/B$ in the covariance is the whole story: gradient noise falls as $1/B$ in variance and $1/\sqrt{B}$ in standard deviation, while the signal is unchanged. Sampling without replacement (as real epoch shuffling does) attaches a finite-population factor $(N-B)/(N-1)$, negligible when $B \ll N$.
Temperature: what the pair $(\eta, B)$ actually sets
The continuous-time view treats SGD as a stochastic differential equation whose stationary distribution depends on $\eta$ and $B$ only through the ratio $\eta/B$, often called the noise scale or temperature (Jastrzębski et al., 2017, Three Factors Influencing Minima in SGD; Smith and Le, 2018, A Bayesian Perspective on Generalization and SGD, ICLR). This is the theoretical basis of the linear scaling rule: multiply $B$ by $k$, multiply $\eta$ by $k$, and the trajectory statistics are approximately preserved.
Goyal et al. (2017), Accurate, Large Minibatch SGD: Training ImageNet in 1 Hour, is the practical demonstration — batch 8,192 with linear scaling plus a gradual warmup matches small-batch accuracy, where the warmup exists precisely because the linear rule breaks in the first epochs when curvature is changing fast. You, Gitman and Ginsburg (2017), LARS, and You et al. (2020), LAMB, push further by making the effective rate layer-wise, which is how batch sizes in the tens of thousands became trainable at all. The $\sqrt{k}$ alternative arises naturally for adaptive optimisers, where the preconditioner already absorbs part of the scaling.
Diminishing returns, and where they bite
Shallue et al. (2019), Measuring the Effects of Data Parallelism on Neural Network Training (JMLR), is the definitive empirical study: across models and datasets, steps-to-target-accuracy falls roughly linearly with batch size up to a workload-dependent knee, after which extra batch buys almost nothing. McCandlish et al. (2018), An Empirical Model of Large-Batch Training, give a predictor for that knee — the gradient noise scale
$$ B_{\text{noise}} = \frac{\operatorname{tr}(\Sigma)}{|\nabla \hat{R}|^2} $$
the ratio of gradient variance to squared gradient norm. Below $B_{\text{noise}}$, doubling the batch roughly halves the steps needed; far above it, steps barely improve and you are paying compute for a smoother chart. The quantity is estimable during training from the variance of gradients across workers or micro-batches, and it grows over the course of training — which is the principled argument for batch-size schedules rather than one fixed value.
The generalisation-gap debate
Keskar et al. (2017), On Large-Batch Training for Deep Learning (ICLR), reported that large-batch training converges to sharper minima and generalises worse, proposing sharpness as the mechanism. The claim was substantially qualified afterwards: Hoffer, Hubara and Soudry (2017), Train Longer, Generalize Better (NeurIPS), showed much of the gap disappears when large-batch runs are given a matched number of updates rather than epochs — the same axis error this lesson's last two columns demonstrate at toy scale. Dinh et al. (2017) further showed sharpness as usually defined is not reparameterisation-invariant, so a sharp minimum can be made flat without changing the function. The current consensus is narrow but useful: with the learning rate rescaled, an adequate warmup and a matched update budget, large-batch training usually matches small-batch quality up to the noise-scale knee, and the residual differences are regularisation effects rather than a law about minima.
For the curve-reading practitioner, three invariants survive all of the above: plotted jitter scales as $1/\sqrt{B}$; the epoch axis is not comparable across batch sizes; and $\eta$ and $B$ must move together or the experiment measures neither.
What to learn next
- Finding a learning rate — the knob that must move whenever batch size does.
- Gradient accumulation — keeping a large effective batch on a small GPU.
- How smoothing and log scales mislead you — the other way a curve's shape lies about the model.