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.
- 7 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.
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 insteadThe 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
- nn.Embedding and index errors — the next standard layer with a sharp edge.
- Transfer learning — where the frozen-backbone BatchNorm trap bites hardest.
- model.train() and model.eval() — the switch this whole lesson depends on.
Developer — Code and libraries.
Setup
pip install torchWritten and tested against torch 2.5 on CPU. The inputs below are fixed numbers, so this output is exact.
Watch the notebook being written
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()])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:
bn.train()
bn(torch.tensor([[10., 200., 3.]])) # a batch of oneValueError: 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
| Layer | Input shape | Typical 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
- nn.Embedding and index errors — the next standard layer with a sharp edge.
- Transfer learning — where the frozen-backbone BatchNorm trap bites hardest.
- model.train() and model.eval() — the switch this whole lesson depends on.
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;
SyncBatchNormall-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
- nn.Embedding and index errors — the next standard layer with a sharp edge.
- Transfer learning — where the frozen-backbone BatchNorm trap bites hardest.
- model.train() and model.eval() — the switch this whole lesson depends on.