๐Ÿ”— Deep Learning ยท Lecture 13 of 38

Batch Normalisation: Faster, More Stable Training

BatchNorm normalises each feature using mini-batch statistics, then rescales it with learned parameters. We derive the forward pass, explain training-versus-inference behaviour, debate why it works, and list its pitfalls.

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):

$$ \mu_B = \frac{1}{m}\sum_{i=1}^{m}x_i, \qquad \sigma_B^2 = \frac{1}{m}\sum_{i=1}^{m}(x_i - \mu_B)^2 $$
$$ \hat{x}_i = \frac{x_i - \mu_B}{\sqrt{\sigma_B^2 + \epsilon}}, \qquad y_i = \gamma\,\hat{x}_i + \beta $$

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:

$$ \mu_{\text{run}} \leftarrow (1 - \alpha)\,\mu_{\text{run}} + \alpha\,\mu_B $$

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.

python
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#

python
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 normalisation

Pitfalls 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).
JA
Written by

Janin A Apurba

B.Sc. in CSE, AUST ยท Advanced ICT Officer, CNRS-UNHCR. Teaching AI, ML and Deep Learning to the next generation of engineers and researchers.

Keep learning

Related lectures

๐Ÿ”— Deep Learning

Beyond BatchNorm: Layer, Group, Instance and RMS Normalisation

Normalisation layers differ only in which axes they average over โ€” yet that choice decides where they work. We compare LayerNorm, GroupNorm, InstanceNorm and RMSNorm, and the pre-norm versus post-norm debate in transformers.

Intermediateโฑ 4 min#110
๐Ÿ”— Deep Learning

Convolutional Neural Networks: The Core Ideas

Convolutions exploit the structure of images through local connectivity, weight sharing and translation equivariance. We define the convolution operation, count parameters, and build a CNN that learns hierarchical features.

Beginnerโฑ 5 min#113
๐Ÿ”— Deep Learning

Padding, Stride, Pooling and Receptive Fields

The geometry of convolutional layers determines output sizes, computational cost and what each unit can see. We derive the output-size formula, compare pooling types, and compute receptive fields.

Beginnerโฑ 5 min#114