In 2015 Sergey Ioffe and Christian Szegedy introduced Batch Normalisation, and it almost immediately became standard in convolutional networks. It allowed much higher learning rates, reduced sensitivity to initialisation, and often improved accuracy. It also has subtle behaviour that causes some of the most confusing bugs in deep learning. Today we cover both.
The operation#
For a mini-batch $\{x_1, \dots, x_m\}$ of values of one feature (one neuron, or one channel in a CNN):
The learned scale $\gamma$ and shift $\beta$ let the network undo the normalisation if that is optimal โ so BatchNorm never reduces what the layer can represent. For convolutional layers, statistics are computed per channel over the batch and spatial positions.
Training versus inference#
At training time, BatchNorm uses the current batch's statistics. At inference time, a single example has no batch, and predictions must be deterministic, so BatchNorm uses running averages of $\mu$ and $\sigma^2$ accumulated during training:
Why does BatchNorm help?#
The original paper attributed its success to reducing internal covariate shift โ the change in each layer's input distribution as earlier layers update. Later work questioned this. Santurkar et al. (2018) showed that BatchNorm networks train well even when distribution shift is deliberately re-injected, and argued the main effect is a smoother optimisation landscape: gradients become more predictable (smaller Lipschitz constants), allowing larger learning rates. Other analyses emphasise that BatchNorm makes a layer's output invariant to the scale of its incoming weights, which interacts with weight decay to produce an effective learning-rate schedule.
Practical benefits are undisputed:
- Higher learning rates and faster convergence;
- Less sensitivity to initialisation;
- Regularisation โ batch statistics add noise, sometimes reducing the need for dropout.
Where to place it#
The original paper placed BatchNorm before the activation: Linear/Conv โ BN โ ReLU. Placing it after the activation also works in practice. A bias in the preceding layer is redundant (BN's $\beta$ replaces it), so set bias=False.
import torch.nn as nn
block = nn.Sequential(
nn.Conv2d(64, 128, kernel_size=3, padding=1, bias=False),
nn.BatchNorm2d(128),
nn.ReLU(inplace=True),
)BatchNorm from scratch#
import numpy as np
class BatchNorm1d:
def __init__(self, dim, momentum=0.1, eps=1e-5):
self.gamma, self.beta = np.ones(dim), np.zeros(dim)
self.run_mean, self.run_var = np.zeros(dim), np.ones(dim)
self.m, self.eps, self.training = momentum, eps, True
def __call__(self, x):
if self.training:
mu, var = x.mean(0), x.var(0)
self.run_mean = (1 - self.m) * self.run_mean + self.m * mu
self.run_var = (1 - self.m) * self.run_var + self.m * var * len(x) / (len(x) - 1)
else:
mu, var = self.run_mean, self.run_var
return self.gamma * (x - mu) / np.sqrt(var + self.eps) + self.beta
bn = BatchNorm1d(3)
for _ in range(200):
bn(np.random.default_rng().normal([5, -2, 100], [2, 0.5, 30], size=(64, 3)))
print("running mean:", bn.run_mean.round(2), " running var:", bn.run_var.round(2))
bn.training = False
print(bn(np.array([[5.0, -2.0, 100.0]])).round(3)) # ~0 after normalisationPitfalls and limitations#
- Small batches: with batch sizes of 1โ4 (common in detection, segmentation and 3-D medical imaging), batch statistics are too noisy. Use Group Normalisation or Layer Normalisation, or synchronise statistics across GPUs (SyncBatchNorm).
- Sequence models: variable sequence lengths and autoregressive decoding make batch statistics awkward โ transformers use LayerNorm instead.
- Train/test discrepancy: if the data distribution shifts, running statistics may no longer fit; re-estimating them on target-domain data is a simple adaptation technique.
- Batch dependence leaks information between examples in a batch, which can break certain contrastive-learning setups unless handled carefully.
- Fine-tuning with small batches often works better with BatchNorm layers frozen (kept in eval mode).