Proposition 2I.8.P0229 of 86 in the corpus
The normalisation formula is incomplete until its axes are named.
Batch normalisation shares statistics between examples; layer normalisation shares them between features of one example. Their forward and backward passes inherit that choice.
The same arithmetic, two different functions
Take a matrix with one example per row. For dense batch normalisation, each feature column supplies its own mean and variance across the examples. For layer normalisation over features, each example supplies its statistics across the entries of its row.
The figure highlights one group in each case. A column pools information between examples. A row does not. Writing only “subtract the mean” leaves this distinction out of the model.
For any group of scalar entries, define
The mean and variance are computed over the same explicitly chosen group.
The forward variance in (I.8.1) has divisor , not . The positive keeps the denominator nonzero. It also means the standardised values have variance , not exactly one. If every input in a group is the same, all are zero.
The model then learns an affine transformation. For batch normalisation, and belong to a feature and are shared across its examples. For layer normalisation, each normalised feature can have its own . Standardising activations does not force the layer’s final outputs to have zero mean or unit variance.
Batch normalisation has two forward passes
Ordinary batch normalisation uses the current group’s statistics during training and frozen running estimates at inference:
Evaluation stops asking the neighbouring examples how to scale this one.
An exponential update might be , where is the weight given to the new batch. This state update is separate from gradient descent on and . Frameworks do not all name their averaging coefficient the same way.
The variance estimator must also be recorded. PyTorch’s documented default uses the biased variance in the training forward pass and an unbiased estimate for the running-variance update. A textbook implementation using the same variance for both is a different convention, not automatically an implementation bug.
The checkpoint therefore contains more than weights: it needs the running statistics and the correct evaluation mode. Keeping training mode at inference can make one prediction depend on the other examples served with it. Some explicit inference policies use batch statistics; they are a different policy and must be evaluated as such.
Layer normalisation has no corresponding switch of statistics. It computes them from each example at both training and inference. This does not make every layer in the network mode-independent: dropout still has its own switch.
The backward pass must differentiate the statistics
Let and . Since , differentiation gives
Applying the quotient rule to now gives the whole Jacobian:
Every entry participates in the shared mean and denominator.
For batch normalisation, write . Multiply (I.8.3) by and sum over :
Two reductions are enough; the dense Jacobian need not be stored.
The bars are means over the group. The parameter gradients are and . For layer normalisation with per-feature scales, first set , then use . Pulling a single outside that expression would be wrong.
At inference, with frozen batch statistics, the derivative instead is . There is no dependency through other examples because the statistics are constants. Training and evaluation have different Jacobians too.
The derivation and a finite-difference test are worked through in I.8.B02.
What the operation removes
Adding a constant to every member of a group leaves (I.8.1) unchanged. This shift invariance explains why the input gradients in (I.8.4) sum to zero.
Positive rescaling is subtler. Replacing by changes the denominator to . For , this is equivalent to replacing by in the original computation. Scale invariance is exact at when , and only approximate when is negligible relative to the variance. Negative scaling also reverses the sign of the standardised entries.
These invariances change the parameterisation and its derivatives. They do not prove that normalisation improves every optimisation problem. Ioffe and Szegedy motivated batch normalisation through changing activation distributions. Santurkar and colleagues later challenged that explanation and studied its effect on optimisation smoothness. Treat the mechanism as an operation we can differentiate, not as a slogan that the data “stay the same.”
Reproduce the forward pass
This scalar-feature example is pure Python. It uses population-divisor variance for the training group and explicitly supplied frozen statistics for evaluation. It does not estimate a running variance from one batch.
from math import sqrt
def batch_norm(x, *, stats=None, gamma=1.0, beta=0.0, eps=1e-5):
if stats is None:
mu = sum(x) / len(x)
var = sum((v - mu)**2 for v in x) / len(x)
else:
mu, var = stats
return [gamma * (v - mu) / sqrt(var + eps) + beta for v in x]
print([round(v, 4) for v in batch_norm([1.0, 2.0, 3.0])])
print([round(v, 4) for v in batch_norm([3.0], stats=(2.0, 2.0/3.0))])
It prints [-1.2247, 0.0, 1.2247] and [1.2247].
The complete reproduction test is attached to
I.8.B01.
For a channels-first tensor of shape , spatial batch normalisation usually reduces over separately per channel. Layer normalisation has no single universal image convention: one must specify which trailing dimensions, or which rearranged feature axis, form its group. See the shape audit in I.8.B04 before translating the matrix drawing into image code.
Sources
Depends on
Used by
Nothing yet.
Problems using this
- I.8.B01 — One feature, two forward passesnumeric▲△△
- I.8.B02 — Differentiate the statistics, not just the numeratorsymbolic▲▲▲
- I.8.B03 — The same example changes sign in another batchcounterexample▲▲▲
- I.8.B04 — Count the groups before counting the parametersshape▲▲△
- I.8.X01 — Epsilon breaks exact scale invariancesymbolic▲▲△
- I.8.X02 — The same inputs in training and evaluationnumeric▲△△
- I.8.X05 — Batch size is not the reduction sizecounterexample▲▲△
- I.8.X06 — A saved variance has an estimator conventioncounterexample▲▲△
- I.8.X08 — An axis bug that preserves every tensor dimensionshape▲▲△