Layer Normalization, annotated
How to read this page
- Any dotted word explains itself on hover, focus or tap; so does every symbol in every equation.
- The batch-versus-layer grid in §3 is the one picture to remember: the two methods compute the same statistics along different directions of the same table.
Each idea climbs the ladder: everyday picture, tiny example, diagram, math, why it matters today. The deep networks lesson implements batch norm and layer norm from scratch and shows why Transformers use the latter.
Abstract · original
“In this paper, we transpose batch normalization into layer normalization by computing the mean and variance used for normalization from all of the summed inputs to the neurons in a layer on a single training case.”Ba, Kiros and Hinton (2016), Abstract
Everyday picture
Two ways to grade on a curve. Batch normalization grades each question by comparing everyone in the class on that question: to know if your answer is good, it needs the rest of the class present. Layer normalization grades each student by comparing their answers to all the questions against their own average: it needs nobody else. That independence is the whole point.
What the paper claims
- Normalization statistics can come from all the units in one layer, for one example, instead of from one unit across a batch.
- So it works with any batch size, including 1, and does exactly the same computation during training and at test time.
- It fits recurrent networks naturally, one normalization per time step, and stabilizes their hidden states.
- It speeds up training on six tasks, mostly with recurrent networks.
Why it matters today
Every Transformer, and so every large language model, normalizes with layer norm or its close cousin RMSNorm, typically twice per block. The paper's own experiments are about RNNs; its lasting impact came a year later, when the Transformer adopted it.
1 Introduction · original
Everyday picture
Deep networks are like a relay of translators, each passing their version to the next. If one translator starts shouting (numbers growing) or whispering (numbers shrinking), everyone after them struggles. Normalization is a volume knob at every hand-off that resets the message to a standard level. Batch normalization already did this and made training much faster, but its knob needs a whole batch of examples to set the level.
Tiny example
Batch norm estimates a neuron's average from the current batch. With a batch of 256 that estimate is decent; with a batch of 2 it is little more than the average of two random numbers; with a batch of 1 the spread is zero and the formula divides by zero. Recurrent networks make it worse: a sentence of 40 words needs separate statistics for each of 40 time steps, and a test sentence of 60 words has no statistics at all for steps 41 to 60.
Why it matters today
Language models train on variable-length sequences, sharded across many GPUs, often with small per-device batches. A normalization that never looks across the batch is simply easier to live with.
2 Background: batch normalization · original
Everyday picture
A neuron adds up its weighted inputs into one number, its summed input. Batch norm watches that number across all the examples in the batch, subtracts the batch average, divides by the batch spread, and then lets a learned gain set the spread it actually wants.
Tiny example
One neuron's summed inputs across a batch of three examples are 1, 5 and 9. The mean is 5 and the standard deviation is √((16 + 0 + 16) / 3) = 3.266. So the normalized values are (1 − 5)/3.266 = −1.225, then 0, then +1.225. With gain g = 1 they stay as they are.
In words: “each neuron's summed input is the dot product of its weights with the previous layer's outputs; batch norm subtracts that neuron's average over the data, divides by its spread over the data, and multiplies by a learned gain.”
With the numbers: for the batch (1, 5, 9): μ = 5, σ = 3.266, so ā = (1/3.266) × (1 − 5) = −1.225 for the first example.
In Python:
import math
# one neuron's summed input a_i, over a batch of 3 examples
a = [1, 5, 9]
# μ_i = E_x[a_i]
mu = sum(a) / len(a)
# σ_i
sigma = math.sqrt(sum((a_x - mu) ** 2 for a_x in a) / len(a))
mu, round(sigma, 3) # → (5.0, 3.266)
g = 1
# ā_i for the first example
round(g / sigma * (a[0] - mu), 3) # → -1.225
The expectation E is over all the training data, which would need a pass through the whole dataset after every weight update. So in practice μ and σ are estimated from the current mini-batch. That estimate is what ties batch norm to the batch size.
Why it matters today
Batch norm is still standard in convolutional image networks (the ResNet companion uses it after every convolution), where batches are large and every example has the same shape.
3 Layer normalization · original
“Unlike batch normalization, layer normaliztion does not impose any constraint on the size of a mini-batch and it can be used in the pure online regime with batch size 1.”Ba, Kiros and Hinton (2016), §3 (spelling as in the original)
Everyday picture
Picture the numbers as a table: one row per example in the batch, one column per neuron in the layer. Batch norm computes its average and spread down each column (one neuron, many examples). Layer norm computes them along each row (one example, all its neurons). Same arithmetic, different direction.
Try it: the same table, two directions
Summed inputs a (4 examples × 5 neurons)
After normalizing
Hover, tab to or tap a cell to see which numbers its average and spread come from.
Reading it: the left table holds the summed inputs for a batch of 4 examples (rows) and a layer of 5 neurons (columns). Pick a cell. In layer norm mode, its whole row lights up: the average and spread come only from that one example's own 5 neurons, so every row ends up with mean 0 and spread 1 by itself, and the other examples are irrelevant. Switch to batch norm and its column lights up instead: the statistics come from the same neuron across the 4 examples, so a cell's normalized value depends on what else happened to be in the batch. Try example 3: its values (5, 5, 5, 6, 4) are nearly flat, so layer norm stretches its small differences to full size, while batch norm judges them against the other examples.
In words: “for one example, average the summed inputs of all H neurons in the layer, and measure their spread around that average; every neuron in the layer then shares these two numbers.”
With the numbers: for a layer of H = 4 neurons with a = (2, 4, 6, 8): μ = 20 / 4 = 5, σ = √((9 + 1 + 1 + 9) / 4) = √5 = 2.236, so the normalized values are (−1.342, −0.447, 0.447, 1.342).
In Python:
import math
# the H summed inputs of one layer, one example
a = [2, 4, 6, 8]
H = len(a)
# μ^l
mu = sum(a) / H
# σ^l
sigma = math.sqrt(sum((a_i - mu) ** 2 for a_i in a) / H)
mu, round(sigma, 3) # → (5.0, 2.236)
[round((a_i - mu) / sigma, 3) for a_i in a] # → [-1.342, -0.447, 0.447, 1.342]
As with batch norm, each neuron then gets its own learned bias b and gain g, applied after normalizing and before the nonlinearity (equation 5 in §5 shows the shared form). And because nothing depends on the batch, training and testing run the identical computation: no running averages to store, no train/test mismatch.
Why it matters today
In a Transformer, each token's vector is one “row”: layer norm normalizes each token across its dmodel features (for example 512), independently of every other token and every other sentence in the batch. Real implementations also add a tiny ε under the square root so the division is safe.
3.1 Layer-normalized recurrent networks · original
Everyday picture
A recurrent network is a note-taker who reads one word at a time and rewrites a one-page summary (the hidden state) after each word. If the handwriting gets a little bigger with every rewrite, after a hundred words it no longer fits the page; a little smaller, and it fades to nothing. Layer norm resets the summary to a standard size after every word, whatever the sentence length.
Tiny example
Suppose the summed inputs at every step are 1.2 times those of the step before. After 50 steps they are 1.250 ≈ 9,100 times larger. Layer norm divides by the current step's own spread, so the normalized values stay the same size at step 1 and at step 50.
Hover or tap a part. The dashed line carries hₜ back around to become the next step's hₜ₋₁.
Reading it: read from the bottom up. The new word xt and the previous summary ht−1 are each multiplied by their own weight matrix and added into the summed inputs at. The yellow boxes are layer norm: compute this step's own mean and spread across all the hidden units, normalize, then apply the learned gain g and bias b. Only then does the nonlinearity produce the new summary ht, which the dashed line feeds back in at the next step. The statistics are recomputed at every step from that step alone, while g and b are shared across all steps. That is why any sentence length works.
In words: “mix the previous summary and the new input into summed inputs, normalize them using this step's own mean and spread, apply a learned per-unit gain and bias, then the nonlinearity.”
With the numbers: if at = (2, 4, 6, 8) then μt = 5 and σt = 2.236; with g = (1, 1, 1, 1) and b = 0 the normalized vector is (−1.342, −0.447, 0.447, 1.342), and tanh of it gives ht ≈ (−0.872, −0.420, 0.420, 0.872).
In Python:
import math
# a^t
a = [2, 4, 6, 8]
# μ^t
mu = sum(a) / len(a)
# σ^t
sigma = math.sqrt(sum((a_i - mu) ** 2 for a_i in a) / len(a))
g, b = [1, 1, 1, 1], [0, 0, 0, 0]
# f = tanh
h = [math.tanh(g_i / sigma * (a_i - mu) + b_i) for g_i, a_i, b_i in zip(g, a, b)]
[round(h_i, 3) for h_i in h] # → [-0.872, -0.42, 0.42, 0.872]
The paper points out the key stability property: layer norm makes the recurrent layer invariant to rescaling all its summed inputs, so hidden states can no longer grow or shrink step after step, the root of exploding and vanishing gradients in RNNs.
Why it matters today
“Normalize per position, share the gain and bias across positions” is exactly what a Transformer does per token. The step from “per time step” to “per token” is a small one.
4 Related work · original
Everyday picture
Several ways to keep a network's numbers at a sensible level had appeared by 2016. Recurrent batch norm kept separate statistics for every time step and needed careful initialization of its gain (0.1 worked, 1.0 did not). Weight normalization divides each neuron's weight vector by its length instead of looking at activations at all. Layer norm is different from both: it is not a re-parameterization of the network but a genuinely different computation, which is why its invariances in §5 differ.
Why it matters today
The idea space keeps growing (group norm for small-batch vision, RMSNorm for language models), but these three axes, across the batch, across the layer, or on the weights, remain the vocabulary for describing any of them.
5 Analysis · original
5.1 Invariance under weights and data transformations · original
Everyday picture
If a scale reads in kilograms and you switch it to pounds, a comparison of relative weights (“this one is twice that one”) is unaffected. A normalization method is invariant to a change if the network's output does not move when you make it. Invariances matter because they mean the network's behaviour cannot be thrown off by that kind of change during training.
Tiny example
Take the example a = (2, 4, 6, 8) and multiply the whole input by 3: a becomes (6, 12, 18, 24), μ becomes 15 and σ becomes 6.708, and the normalized values are still (−1.342, −0.447, 0.447, 1.342). Layer norm cannot tell the difference. Batch norm can: it only rescales if every example in the batch is rescaled together.
In words: “all three methods normalize a neuron's summed input with some mean and spread, then apply a gain and bias; for layer norm, scaling one example's input by δ scales its mean and spread by δ too, so the δ cancels and the output is unchanged.”
With the numbers: δ = 3: (6 − 15) / 6.708 = −1.342, the same as (2 − 5) / 2.236.
In Python:
import math
delta = 3
# every input scaled by δ
x = [delta * a_i for a_i in [2, 4, 6, 8]]
# μ' = δμ
mu = sum(x) / len(x)
# σ' = δσ
sigma = math.sqrt(sum((x_i - mu) ** 2 for x_i in x) / len(x))
mu, round(sigma, 3) # → (15.0, 6.708)
# same as (2 − 5) / 2.236
round((x[0] - mu) / sigma, 3) # → -1.342
| Rescale the weight matrix | Re-centre the weight matrix | Rescale one weight vector | Rescale the dataset | Re-centre the dataset | Rescale one example | |
|---|---|---|---|---|---|---|
| Batch norm | yes | no | yes | yes | yes | no |
| Weight norm | yes | no | yes | no | no | no |
| Layer norm | yes | yes | no | yes | no | yes |
Reading it: the bars are the four layer-normalized outputs for the example a = (2, 4, 6, 8). Drag scale δ, which multiplies the whole example, and shift, which adds the same amount to every summed input (as re-centring the weight matrix does): the bars do not move, because the mean absorbs the shift and the spread absorbs the scale. That is the “yes” in the “rescale one example” and “re-centre the weight matrix” columns. Now drag scale neuron 1 only, which rescales a single neuron's weights: the bars change, because one neuron now pulls the shared mean and spread. That is the “no” in the “rescale one weight vector” column.
5.2 Geometry of parameter space during learning · original
Everyday picture
If a normalized neuron's output only depends on the direction of its weight vector, not its length, then a long weight vector is like a long lever: the same push moves the tip a smaller angle. As a weight vector grows during training, the same gradient step turns it less. The paper calls this an implicit reduction of the learning rate, a built-in “settling down” that stabilizes training.
Tiny example
Direction is what matters, so compare angles. A weight vector (1, 0) nudged by (0, 0.1) turns by atan(0.1 / 1) ≈ 5.7°. Double its length to (2, 0) and the same nudge turns it by atan(0.1 / 2) ≈ 2.9°. The output depends only on the angle, so the doubled weights learn at about half the speed.
In words: “measure the size of a small parameter change by how much it changes the model's predictions; for small changes that is approximately a quadratic form in the change, whose matrix, the Fisher information, describes how sensitive predictions are in each direction.”
With the numbers: the paper's result is that in normalized models this sensitivity along a weight vector scales with 1/‖w‖: doubling ‖w‖ from 1 to 2 halves the curvature along it, matching the 5.7° versus 2.9° example. To see the approximation itself, take a coin-like model that predicts y = 1 with probability θ = 0.5 and nudge θ by δ = 0.1: DKL = 0.5 ln(0.5 / 0.6) + 0.5 ln(0.5 / 0.4) = 0.0204. A coin's Fisher information is 1 / (θ(1 − θ)) = 4, so ½ × 0.1 × 4 × 0.1 = 0.020: nearly the same.
In Python:
import math
round(math.degrees(math.atan(0.1 / 1)), 1), round(math.degrees(math.atan(0.1 / 2)), 1) # → (5.7, 2.9)
# a coin: P(y = 1) = θ, nudged by δ
theta, delta = 0.5, 0.1
D_KL = theta * math.log(theta / (theta + delta)) + (1 - theta) * math.log((1 - theta) / (1 - theta - delta))
round(D_KL, 4) # → 0.0204
# the coin's Fisher information F(θ)
F = 1 / (theta * (1 - theta))
# ½ δᵀ F(θ) δ
F, round(0.5 * delta * F * delta, 3) # → (4.0, 0.02)
The same analysis shows that learning the gain g (the neuron's output scale) depends only on the size of the prediction error, not on the scale of the input, so it is robust to inputs of any size.
Why it matters today
This interplay between weight length and effective learning rate is one reason weight decay still matters in normalized networks: shrinking the weights keeps their effective learning rate from quietly decaying.
6 Experimental results · original
Everyday picture
Six tasks, chosen to be where batch norm is awkward: mostly recurrent networks with variable-length sequences. In every case the gain and bias start at 1 and 0, and the comparison is the same model with and without layer norm.
| § | Task and model | What layer norm did |
|---|---|---|
| 6.1 | Matching images and captions (order embeddings, GRU text encoder, COCO) | Reached its best validation model in 60% of the baseline's time; caption-retrieval recall@1 rose from 46.6 to 48.5 and image-retrieval recall@1 from 37.8 to 38.9 |
| 6.2 | Question answering (attentive reader LSTM, CNN news corpus) | Trained faster and ended better than both the baseline and recurrent batch norm, and did not need the special 0.1 gain initialization recurrent batch norm relies on |
| 6.3 | Skip-thought sentence vectors (BookCorpus, 2,400-dimensional encoder) | Faster training and better downstream scores after 1 million iterations, e.g. movie-review sentiment accuracy 77.3 → 79.5 |
| 6.4 | Generating MNIST digits (DRAW, 256 LSTM units) | Converged almost twice as fast; after 200 epochs, 82.09 against 82.36 nats of test negative log-likelihood (lower is better) |
| 6.5 | Handwriting generation (3 LSTM layers, sequences of about 700 steps, mini-batches of 8) | Reached a comparable final likelihood much faster; the long sequences and tiny batches are exactly where stable hidden states matter |
| 6.6 | Permutation-invariant MNIST (784-1000-1000-10 feed-forward net) | Robust to batch size, including a batch of 4, and faster than batch norm applied to all layers |
| 6.7 | Convolutional networks (preliminary) | Faster than no normalization, but batch norm was better |
Tiny example: why batch norm hates small batches
Hover or tap to read the noise at each batch size.
Reading it: this chart is computed from simple statistics, not taken from the paper. If a neuron's summed input varies across examples with spread 1, then batch norm's estimate of its mean from a batch of B examples wobbles by about 1/√B from batch to batch (blue, solid). At B = 256 that is ±0.06, negligible; at B = 4 it is ±0.5, so the same example is normalized differently depending on who shares its batch. Layer norm (red, dashed) does not use the batch at all, so its line is flat at zero. The paper's §6.6 saw this in practice: with a batch of 4, layer norm still trained well.
Why it matters today
§6.7 is the honest footnote: in convolutional networks, units near the image border behave very differently from those in the centre, so pooling statistics across all of a layer's units fits them poorly. That is why image CNNs kept batch norm, while sequence models, and then Transformers, standardized on layer norm.
7 Conclusion · original
“Empirically, we showed that recurrent neural networks benefit the most from the proposed method especially for long sequences and small mini-batches.”Ba, Kiros and Hinton (2016), §7
A short paper with a simple idea: compute the same statistics along the other axis of the table. Its authors aimed it at RNNs. The architecture that made it indispensable, the Transformer, arrived a year later.
What changed since 2016
| In the paper | Common today | Why | Where to learn more |
|---|---|---|---|
| Layer norm inside RNNs | Layer norm around every Transformer sub-layer | Each token is normalized on its own, whatever the batch or sequence length | Transformer companion |
| Normalize after the sum (post-norm, as the original Transformer did) | Pre-norm: x + Sublayer(LayerNorm(x)) | Leaves a clean residual path, so very deep stacks train stably | transformer lesson |
| Subtract the mean, divide by the spread | RMSNorm (Zhang and Sennrich, 2019): skip the mean, divide by the root-mean-square | Cheaper, and works about as well; used by many large language models | deep nets lesson |
| Gain and bias per unit | Often gain only, no bias | Fewer parameters with no loss in quality at scale |
In words: “divide each summed input by the root-mean-square of the whole layer's summed inputs, then apply the gain; no mean is subtracted.”
With the numbers: for a = (2, 4, 6, 8): the mean of squares is (4 + 16 + 36 + 64) / 4 = 30, its root is 5.477, so RMSNorm gives (0.365, 0.730, 1.095, 1.461) with g = 1. Unlike layer norm, the values keep their positive offset.
In Python:
import math
a, g = [2, 4, 6, 8], [1, 1, 1, 1]
H = len(a)
# √((1/H) Σ_j a_j²)
rms = math.sqrt(sum(a_j ** 2 for a_j in a) / H)
round(rms, 3) # → 5.477
[round(a_i / rms * g_i, 3) for a_i, g_i in zip(a, g)] # → [0.365, 0.73, 1.095, 1.461]
Glossary
Every term with hover guidance on this page, in one place.