Reading Training Curves

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.

On this page 5
  1. Why this had to be invented
  2. How it works
  3. A real example you have seen
  4. Remember this
  5. What to learn next

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 adjustments

Here 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

Developer — Code and libraries.

Setup

bash
pip install torch numpy

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

batch_size_effects.py
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}")
Output
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.3010

The 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

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