Era 4 · Architectures and scale · 2015

26 Batch Normalization

Batch Normalization: Accelerating Deep Network Training by Reducing Internal Covariate Shift · Ioffe & Szegedy · ICML
🟧 read selectively~45 minoriginal ↗
The gist in 20 seconds. Normalize the input of every layer over the mini-batch (zero mean, unit variance) plus a learnable scale and shift. Activations stabilize → a larger learning rate, faster convergence, less dependence on initialization. All but a mandatory part of a CNN.

Context

Deep networks are sensitive to initialization and to the learning rate; the distribution of each layer's inputs drifts as the layers below it train, slowing convergence down. Ioffe and Szegedy offer a cure.

The idea and the mechanism

For each mini-batch we normalize the layer's input per feature, then scale and shift it with learnable γ, β (so the network can restore any distribution it wants). At inference, instead of batch statistics we use the running averages accumulated during training.

optimization The BN transform, and the argument about why it works

For a mini-batch B we compute the mean and variance of each feature, normalize, and apply the learnable parameters:

μB = mean(x),   σ²B = var(x)
x̂ = x − μB√(σ²B + ε),   y = γ x̂ + β

The learnable γ, β matter: without them the layer could not restore a useful distribution when it needs one — normalization must not impoverish the representation.

The argument about the cause. The authors explained the effect by a reduction in "internal covariate shift" (the drift of distributions between layers). An MIT paper (2018) later disputed that: the measurements showed BN does not so much remove the drift as smooth the loss landscape (make gradients more predictable), which is what lets you take a bigger step. So cite BN as a landmark technique, but be careful with its original explanation.

NumPy Forward Batch Normalization
import numpy as np

def batchnorm(x, gamma, beta, eps=1e-5):   # x: (batch, features)
    mu  = x.mean(0)
    var = x.var(0)
    xhat = (x - mu) / np.sqrt(var + eps)    # normalize over the batch
    return gamma * xhat + beta              # learnable scale/shift
# at inference mu, var are the running averages from training
batch of activations − μ, ÷ σ centred at 0, variance 1 γ·x̂ + β stable activations
The batch's activations are centred and scaled to a standard form, then the learnable γ, β give back whatever shape is needed. The distributions stop drifting.
Analogy. A teacher who, before every assignment, rescales the class's marks to a common scale (mean 0, spread 1). The next "layer" of markers then always works with a predictable range and is not thrown off when the previous ones started grading systematically higher or lower. γ and β are the right to bring your own scale back, if it really is needed.

Why it matters

It became all but mandatory in CNNs; it sped up training of deep networks several times over (the same accuracy in ~14× fewer steps) and cut the dependence on initialization. Transformers more often use LayerNorm (normalizing over the features of a single example rather than over the batch) — but the "normalize your activations" idea comes from here.

Connections

→ folded into27. ResNet

ResNet puts BatchNorm inside every residual block. Together, skip connections and BN are the two pillars very deep networks stand on: one rescues the gradient, the other stabilizes the activations.

↔ rival19. Dropout

Both regularize or stabilize, but in different ways, and they get on badly: BN relies on batch statistics, while dropout adds noise to them. With BN and large datasets, dropout was pushed noticeably out of convolutional networks.

↔ replaced in32. Transformer

Transformers use LayerNorm rather than BatchNorm — normalization over the features of a single example. The reason: in NLP batches have varying lengths and depending on batch statistics is harmful; normalizing "within the example" is more robust.

Questions worth asking

If the "covariate shift" explanation is wrong, why does BN help so much anyway?

Because the benefit is real, it is the mechanism that turned out to be different. According to the MIT paper, BN makes the loss landscape smoother — gradients become steadier and more predictable, which is what allows a larger learning rate and fast convergence. Same effect, but not the cause the authors claimed. A classic case of "works, just not for the stated reason".

Why does BN work badly with a small batch size?

μ and σ are estimated from the batch; on a batch of 2–4 examples they are noisy and unstable, and the normalization starts doing harm. Batch-independent alternatives were invented for exactly those cases: LayerNorm, GroupNorm, InstanceNorm. Dependence on batch size is BatchNorm's main practical weakness.

BN behaves differently in training and at inference — isn't that a source of bugs?

It certainly is. In training it uses the current batch's statistics, at inference the accumulated running averages. Forgetting to switch the model into eval mode is the classic mistake, and it gives you predictions that wobble. And if the distribution at inference differs from the training one, the running averages may not fit. BN's dual mode is a well-known source of subtle bugs.

What to read in the original

This write-up plus the key passages is enough: the transform itself and the train/inference modes. On "why it works", the critique (MIT 2018) is more useful reading than the original explanation.