Building Models with nn.Module

BatchNorm and its running statistics

BatchNorm rescales each batch using the batch's own statistics during training, while quietly recording long-term averages it will replay at inference — and that double life causes most of its bugs.

On this page 5
  1. Why it exists
  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.

BatchNorm is a layer that rescales the numbers flowing through a network so no layer receives values that are wildly too big or too small.

Think of a teacher who grades on a curve. Each test, she looks at how the whole class did, then adjusts everyone's marks relative to that day's average. A brutally hard test and an easy one end up comparable.

But there is a second thing she does. All term, she keeps a running note of typical class performance. When one student shows up alone for a make-up exam, there is no class to compare against — so she grades using the term's notebook instead.

BatchNorm does both jobs. During training it uses the current group's statistics. And it keeps a notebook — the running statistics — for later, when it must judge inputs one at a time.

Why it exists

As training updates a network, the size and spread of the numbers inside it keep shifting. Layers deep in the stack receive inputs whose scale changes under their feet, which makes training slow and touchy.

BatchNorm, introduced in 2015, steadies each layer's input by rescaling it batch by batch. Networks with it train faster and tolerate bolder settings. It became standard almost overnight, especially in image models.

How it works

training:   batch → measure THIS batch's average and spread
                  → rescale the batch with those
                  → also update the notebook (a small nudge each time)

inference:  single input → no batch to measure
                        → rescale using the NOTEBOOK instead

The notebook is what makes the train/eval switch high-stakes. In train mode the notebook is being written. In eval mode it is only read.

A real example you have seen

Photo apps that recognise faces mostly run image models trained with BatchNorm. When your phone checks one photo — a batch of exactly one — it is running on the notebook, not on live statistics.

Remember this

  • Training mode: rescale with the current batch, and update the notebook.
  • Eval mode: rescale with the notebook — steady, repeatable.
  • Most BatchNorm bugs are the notebook being written or read at the wrong time.

What to learn next

Developer — Code and libraries.

Setup

bash
pip install torch

Written and tested against torch 2.5 on CPU. The inputs below are fixed numbers, so this output is exact.

Watch the notebook being written

running_stats.py
import torch
from torch import nn

bn = nn.BatchNorm1d(3)
print("fresh running_mean:", bn.running_mean.tolist())

batch = torch.tensor([[10., 200., 3.],
                      [14., 260., 5.],
                      [12., 230., 4.]])

bn.train()
for step in range(1, 4):
    bn(batch)                                   # each forward nudges the running stats
    print(f"after step {step}: {[round(v, 2) for v in bn.running_mean.tolist()]}")

bn.eval()
bn(batch)                                       # eval: stats are read, never written
print("after an eval pass:", [round(v, 2) for v in bn.running_mean.tolist()])
Output
fresh running_mean: [0.0, 0.0, 0.0]
after step 1: [1.2, 23.0, 0.4]
after step 2: [2.28, 43.7, 0.76]
after step 3: [3.25, 62.33, 1.08]
after an eval pass: [3.25, 62.33, 1.08]

The batch's true means are [12, 230, 4], and the notebook creeps toward them — ten percent of the distance per step, by default. The eval pass changes nothing. That creep is the entire mechanism behind most BatchNorm bugs.

running_mean and running_var are buffers — recorded state, saved in the checkpoint, never touched by the optimizer. They are the flagship example from buffers vs parameters. BatchNorm also owns two genuine parameters, weight and bias, which are trained.

The batch-of-one failure

Train mode needs a spread, and one sample has none:

python
bn.train()
bn(torch.tensor([[10., 200., 3.]]))     # a batch of one
Output
ValueError: Expected more than 1 value per channel when training, got input size torch.Size([1, 3])

The same call in eval mode works fine — the notebook needs no batch. If you hit this during training, your last incomplete batch has size 1; drop_last=True on the DataLoader is the usual fix.

Which BatchNorm, where

LayerInput shapeTypical use
BatchNorm1d(N, C) or (N, C, L)tabular, sequences
BatchNorm2d(N, C, H, W)images — the common one
BatchNorm3d(N, C, D, H, W)video, volumetric scans

All three normalise per channel C, across everything else.

Common mistakes

Evaluating in train mode. Every validation batch rewrites the notebook with validation statistics — your model quietly studies the exam paper. Metrics look better than they should, and inference on single inputs later disagrees. The fix is the loop discipline from the previous lesson.

Tiny batches. With 2 to 4 samples, per-batch statistics are noise, and both training and the notebook suffer. Below batch size 8 or so, prefer GroupNorm or LayerNorm, which need no batch at all.

A frozen backbone that keeps writing. Freezing parameters does not freeze the notebook — requires_grad and mode are independent switches. A "frozen" pretrained backbone left in train mode drifts its statistics toward your new data. Call .eval() on the frozen part, every epoch, after model.train(). This returns in transfer learning.

Bias before BatchNorm. A convolution's bias is immediately subtracted away by normalisation. Harmless, but wasted parameters — set bias=False on a conv that feeds a BatchNorm.

Try it yourself

Print bn.running_var alongside running_mean in the script. Predict its steady-state values from the batch, then run twenty steps and check yourself.

What to learn next

Researcher — Mathematics and papers.

The transformation

For channel activations $x$ over a batch $\mathcal{B}$ (Ioffe and Szegedy, 2015, Batch Normalization: Accelerating Deep Network Training by Reducing Internal Covariate Shift):

$$ \hat{x} = \frac{x - \mu_{\mathcal{B}}}{\sqrt{\sigma^2_{\mathcal{B}} + \epsilon}}, \qquad y = \gamma \hat{x} + \beta $$

Where $\mu_{\mathcal{B}}$ and $\sigma^2_{\mathcal{B}}$ are the batch mean and variance per channel, $\epsilon$ a small constant for numerical safety (PyTorch default $10^{-5}$), and $\gamma, \beta$ learned per-channel scale and shift — the weight and bias parameters. The learned affine restores representational capacity that raw normalisation removes: the layer can undo itself if that is what the loss prefers.

Running statistics update as an exponential moving average with momentum $\alpha$ (PyTorch default 0.1, note the opposite naming convention to optimizer momentum):

$$ \mu_{run} \leftarrow (1 - \alpha)\,\mu_{run} + \alpha\,\mu_{\mathcal{B}} $$

PyTorch uses the unbiased variance estimate for the running update but biased variance for the in-batch normalisation — a documented asymmetry visible when reconciling train and eval outputs by hand.

Why it works is genuinely contested

The original paper attributed the benefit to reducing "internal covariate shift". Santurkar et al. (2018), How Does Batch Normalization Help Optimization?, found that inducing covariate shift after BN does not remove the benefit, and argued instead that BN smooths the optimisation landscape — reducing the Lipschitz constant of the loss and gradients. Bjorck et al. (2018) frame it as enabling larger learning rates. The honest summary: the empirical benefit is enormous and reproducible; the mechanism has several partial explanations and no consensus.

The batch-dependence tax

BN couples every sample's output to its batchmates. Consequences:

  • Train/inference mismatch is structural: training uses $\mu_{\mathcal{B}}$, inference uses $\mu_{run}$; the gap grows as batch size shrinks.
  • Small-batch degradation: quantified systematically by Wu and He (2018), Group Normalization, whose GroupNorm normalises over channel groups within each sample — batch-independent, and now default in detection and segmentation where memory forces small batches.
  • Distributed training: per-GPU statistics differ from global; SyncBatchNorm all-reduces $\mu, \sigma^2$ across devices at a communication cost.
  • Information leakage: in contrastive and metric learning, batchmates leak label information through the statistics — one motivation for BN-free designs.

Transformers standardised on LayerNorm (Ba et al., 2016) for exactly this independence; NFNets (Brock et al., 2021) demonstrated BN-free ResNets at state-of-the-art accuracy via scaled weight standardisation and gradient clipping.

Reading

  • Ioffe and Szegedy (2015) — the original.
  • Santurkar et al. (2018) — the mechanism challenged.
  • Wu and He (2018) — GroupNorm and the small-batch data.
  • Brock et al. (2021), High-Performance Large-Scale Image Recognition Without Normalization.

What to learn next