Layer normalisation
Rescale each token's own numbers to a standard size, using only that token and nothing else in the batch, which is what keeps a deep stack of layers from drifting apart.
- 13 min read
- 3 reading levels
- Updated
Read these first
On this page 6
One lesson, three depths. Pick the one that fits you today — you can switch any time.
Beginner — No maths. Plain English.
Layer normalisation rescales each token's numbers so they sit in a standard range, using only that token's own numbers.
Think of a playlist where one song was recorded very loud and the next one very quiet. You spend the evening reaching for the volume. A good music app fixes this by adjusting each track to a standard loudness before playing it.
Notice what the app does not do. It does not make every song sound the same. The tune, the words, the instruments are all untouched. Only the overall loudness is standardised.
That is layer normalisation. It flattens the volume, not the music.
Why a deep model needs it
A transformer stacks dozens of layers, and each one adds to a shared running total. Nothing in that arrangement keeps the numbers in a reasonable range.
Left alone, they drift. Some tokens end up with values in the thousands, others in the thousandths. Every layer after that has to cope with wildly inconsistent input, and training becomes slow and fragile.
Normalising before each layer removes the problem at the source. Whatever arrives, it leaves at a standard size.
The two steps
Step one: recentre. Work out the average of this token's numbers and subtract it from every one of them. The new average is zero.
Step two: rescale. Work out how spread out the numbers are, and divide everything by that. The new spread is one.
Now every token, at every layer, arrives with the same average and the same spread.
There is a third piece. If normalisation were the last word, the model could never leave the standard range. That would be a real loss. So two small sets of learnable numbers follow: one that stretches, one that shifts. The model can undo the normalisation if it wants to, and normally it chooses something in between.
The detail everybody gets wrong
There is an older, more famous technique called batch normalisation. It standardises using the average across the whole batch being processed together.
Layer normalisation is not that. It uses one token's own numbers and nothing else. No other token in the sentence. No other sentence in the batch.
Why this matters, in practice:
- Batches change shape. A model serves one request now and thirty-two later. With batch statistics, a token's output would depend on who it was batched with. That is unacceptable for something users pay for.
- Sentences differ in length. Padding would pollute batch statistics.
- One example at a time is the normal case at serving time. Batch statistics need a batch.
Layer normalisation looks only at the token in front of it. Its answer is always the same for the same input. That predictability is worth a great deal.
BATCH NORM LAYER NORM
averages down the column averages across the row
(across examples) (within one token)
token A [ . . . . ] token A [ . . . . ] <- averaged here
token B [ . . . . ] token B [ . . . . ] <- and here
token C [ . . . . ] token C [ . . . . ] <- and here
^ ^ ^ ^
averaged hereThe small number nobody explains
Dividing by the spread has an obvious danger. What if every number in the token is identical? Then the spread is zero, and dividing by zero breaks everything.
So a tiny number is added to the spread before the division. It is far too small to matter for ordinary values, and it makes the flat case harmless. You will see it in every config file as eps or layer_norm_epsilon.
Remember this
- Each token is standardised using only its own numbers.
- Two learnable pieces let the model stretch and shift afterwards.
- Unlike batch normalisation, the answer never depends on the other examples in the batch.
What to learn next
- Pre-norm vs post-norm — where the norm goes matters more than which norm it is.
- RMSNorm — the simpler version most models use today.
- BatchNorm in PyTorch — the technique this one deliberately avoids.
Developer — Code and libraries.
Setup
pip install torchLayer normalisation, hand-written and verified
import torch
import torch.nn as nn
torch.manual_seed(0)
torch.set_printoptions(precision=4, sci_mode=False)
d_model = 6
x = torch.tensor([[[10., 12., 9., 40., 11., 10.], # one token with an outlier
[ 1., 1., 1., 1., 1., 1.]]]) # one token that is completely flat
def layer_norm(x, weight, bias, eps=1e-5):
mu = x.mean(dim=-1, keepdim=True) # last dim only: one token at a time
var = x.var(dim=-1, keepdim=True, unbiased=False) # biased variance, matching torch
return (x - mu) / torch.sqrt(var + eps) * weight + bias
ref = nn.LayerNorm(d_model)
with torch.no_grad(): # give it non-default gain and bias
ref.weight.copy_(torch.linspace(0.5, 1.5, d_model))
ref.bias.copy_(torch.linspace(-0.2, 0.2, d_model))
mine = layer_norm(x, ref.weight, ref.bias)
print("hand-written vs nn.LayerNorm, max difference:", float((mine - ref(x)).abs().max()))
plain = nn.LayerNorm(d_model, elementwise_affine=False)(x)
print("\ninput:\n", x)
print("after normalising (no gain, no bias):\n", plain)
print("mean of each token afterwards:", plain.mean(-1))
print("std of each token afterwards:", plain.std(-1, unbiased=False))
print("-> the flat token had zero variance; eps is what stops a divide by zero")
print("\nthe property that matters in production: batch independence")
one = x[:, :1] # the first token, entirely alone
batched = torch.cat([x, torch.randn(1, 2, d_model) * 100], dim=1)
ln = nn.LayerNorm(d_model)
print(" token normalised alone :", ln(one)[0, 0])
print(" same token in a big batch:", ln(batched)[0, 0])
print(" max difference:", float((ln(one)[0, 0] - ln(batched)[0, 0]).abs().max()))
bn = nn.BatchNorm1d(d_model, affine=False)
small = bn(x.reshape(-1, d_model))[0]
large = bn(batched.reshape(-1, d_model))[0]
print("\n BatchNorm on that same first token:")
print(" in a batch of 2 tokens:", small.detach())
print(" in a batch of 4 tokens:", large.detach())
print(" -> BatchNorm's answer depends on the token's neighbours. LayerNorm's never does.")
print(" -> transformers see batches of every shape, so they use LayerNorm.")hand-written vs nn.LayerNorm, max difference: 2.384185791015625e-07
input:
tensor([[[10., 12., 9., 40., 11., 10.],
[ 1., 1., 1., 1., 1., 1.]]])
after normalising (no gain, no bias):
tensor([[[-0.4818, -0.3011, -0.5721, 2.2281, -0.3914, -0.4818],
[ 0.0000, 0.0000, 0.0000, 0.0000, 0.0000, 0.0000]]])
mean of each token afterwards: tensor([[ 0.0000, 0.0000]])
std of each token afterwards: tensor([[1.0000, 0.0000]])
-> the flat token had zero variance; eps is what stops a divide by zero
the property that matters in production: batch independence
token normalised alone : tensor([-0.4818, -0.3011, -0.5721, 2.2281, -0.3914, -0.4818],
grad_fn=<SelectBackward0>)
same token in a big batch: tensor([-0.4818, -0.3011, -0.5721, 2.2281, -0.3914, -0.4818],
grad_fn=<SelectBackward0>)
max difference: 0.0
BatchNorm on that same first token:
in a batch of 2 tokens: tensor([1.0000, 1.0000, 1.0000, 1.0000, 1.0000, 1.0000])
in a batch of 4 tokens: tensor([-0.6771, -0.1172, 0.8672, 0.6828, 1.0347, 0.5789])
-> BatchNorm's answer depends on the token's neighbours. LayerNorm's never does.
-> transformers see batches of every shape, so they use LayerNorm.Five things in that output
2.38e-07 means the hand-written version is correct. Two details make it match. dim=-1 restricts the statistics to the last axis, and unbiased=False divides by n rather than n-1. Using the unbiased variance is a real and easily missed discrepancy.
The outlier survived. The input had a 40. among values near 10., and afterwards it is 2.2281 among values near -0.45. Normalisation rescaled the token; it did not flatten the structure inside it. This is the playlist volume point, made numerically.
The flat token became all zeros with a standard deviation of exactly 0.0000. Its variance was zero, so eps prevented a division by zero and the result is harmless. Without eps it would be nan, and that nan would spread through the entire batch.
LayerNorm's batch independence is 0.0, not "small". The same token alone and inside a batch containing values a hundred times larger produced bitwise identical output.
BatchNorm gave 1.0000 six times in a batch of two, then something completely different in a batch of four. With two rows, each feature's standard deviation is exactly half the gap between them. Every standardised value lands on plus or minus one. Add two more rows and the answer changes entirely. That behaviour is fine for image classification with large fixed batches. It is unusable for a text model serving one request at a time.
Where it appears in a real block
h = self.norm1(x) # normalise a copy
x = x + self.attention(h) # the stream itself is never normalised in place
h = self.norm2(x)
x = x + self.feedforward(h)Two separate nn.LayerNorm modules, each with its own gain and bias. Sharing one module between the two positions ties unrelated parameters together and is a silent bug.
Cost
For width d, LayerNorm holds 2d parameters: a gain and a bias. Across a 32-layer model with width 4096 and two norms per layer, that is 524,288 parameters. Negligible against eight billion.
The compute is also negligible, but the memory traffic is not. Normalisation reads and writes the whole activation tensor while doing almost no arithmetic, which makes it memory-bandwidth bound. This is why fused kernels that combine normalisation with the following matrix multiply are worth having. It is also why measured normalisation timings often surprise people. See RMSNorm for a measured example.
Common mistakes
Using unbiased=True when reimplementing. PyTorch uses the biased variance. The difference is a factor of n/(n-1). It is small for large widths, and large enough to fail an equality test.
Normalising over the wrong axis. nn.LayerNorm(d_model) normalises the last axis. For a (batch, tokens, channels) tensor that is right. For a (batch, channels, tokens) tensor it is wrong. That is a common error when porting code from vision.
Setting eps too small in half precision. 1e-12 underflows to zero in float16. Use 1e-5 or 1e-6, which is what the standard configs use.
Expecting model.eval() to change LayerNorm's behaviour. It does not, because there are no running statistics. BatchNorm does change, which is one more reason the two are not interchangeable. See train and eval mode.
Try it yourself
Set eps=0 and normalise the flat token. Confirm you get nan. Then increase eps to 1.0 and watch the outlier token's output shrink noticeably — eps is not entirely free. It is small enough not to matter at sensible values.
What to learn next
- Pre-norm vs post-norm — where the norm goes matters more than which norm it is.
- RMSNorm — the simpler version most models use today.
- BatchNorm in PyTorch — the technique this one deliberately avoids.
Researcher — Mathematics and papers.
Definition
For a token vector $x \in \mathbb{R}^d$:
$$ \mu = \frac{1}{d} \sum_{i=1}^{d} x_i, \qquad \sigma^2 = \frac{1}{d} \sum_{i=1}^{d} (x_i - \mu)^2 $$
$$ \operatorname{LN}(x)_i = \frac{x_i - \mu}{\sqrt{\sigma^2 + \epsilon}} \gamma_i + \beta_i $$
with $\gamma, \beta \in \mathbb{R}^d$ learned and $\epsilon$ a small constant, typically $10^{-5}$. Statistics are computed per token, so the operation is independent across both the batch and the sequence axes. Introduced by Ba, Kiros and Hinton (2016), arXiv:1607.06450.
Invariances
LayerNorm is invariant to two transformations of its input:
$$ \operatorname{LN}(a x + b\mathbf{1}) = \operatorname{LN}(x) \quad \text{for } a > 0,\ b \in \mathbb{R} $$
That is, invariance to per-token rescaling and to per-token shifts along the all-ones direction. Equivalently, it projects $x$ onto the hyperplane orthogonal to $\mathbf{1}$. It then projects onto a sphere of radius $\sqrt{d}$ in that hyperplane.
Zhang and Sennrich (2019) test which of the two invariances is doing the work. They conclude re-centering is dispensable, motivating RMSNorm.
Why the original explanation did not survive
Both BatchNorm and LayerNorm were introduced with an "internal covariate shift" justification. Santurkar et al. (2018), arXiv:1805.11604, showed empirically that BatchNorm does not reduce covariate shift. Injecting shift after the normalisation does not remove the benefit either. Their alternative account is that normalisation smooths the loss landscape. It reduces the Lipschitz constant of the loss and its gradient, permitting larger stable learning rates.
Xu et al. (2019), Understanding and Improving Layer Normalization, arXiv:1911.07013, isolate the gradient contribution specifically. The backward pass through the mean and variance terms projects the incoming gradient. That projection, not the forward standardisation, accounts for most of the observed benefit. They also report that the bias and gain frequently increase overfitting risk, and propose removing them.
The gradient
Writing $\hat{x} = (x - \mu)/\sqrt{\sigma^2 + \epsilon}$ and $g = \partial \mathcal{L} / \partial \hat{x}$:
$$ \frac{\partial \mathcal{L}}{\partial x} = \frac{1}{\sqrt{\sigma^2 + \epsilon}} \left( g - \frac{1}{d}\sum_j g_j - \hat{x} \cdot \frac{1}{d}\sum_j g_j \hat{x}_j \right) $$
The two subtracted terms remove the component of $g$ along $\mathbf{1}$ and along $\hat{x}$. The incoming gradient is therefore projected onto the subspace orthogonal to both. That follows directly from the forward invariances. No gradient flows in directions the forward pass ignores.
One practical consequence: the gradient magnitude scales as $1/\sigma$. A token whose activations happen to have small spread receives a proportionally larger gradient.
Placement, and the alternatives
Where the norm sits relative to the residual addition changes training behaviour more than which norm is used. That is covered separately in pre-norm vs post-norm.
Related variants worth knowing:
- RMSNorm (Zhang and Sennrich, 2019, arXiv:1910.07467) drops the mean subtraction and the bias.
- ScaleNorm (Nguyen and Salazar, 2019, arXiv:1910.05895) replaces the per-channel gain with a single learned scalar. It normalises to a fixed radius.
- DeepNorm (Wang et al., 2022, arXiv:2203.00555) keeps post-norm. The residual branch is scaled by a depth-dependent constant. It reports stable training to 1,000 layers.
- Normalisation-free designs (Brock et al., 2021, NFNets, arXiv:2102.06171) replace normalisation in vision. Careful initialisation and gradient clipping do the job instead. The equivalent has not displaced normalisation in language models.
Papers
- Ba, Kiros and Hinton, Layer Normalization, 2016 — arxiv.org/abs/1607.06450
- Santurkar et al., How Does Batch Normalization Help Optimization?, 2018 — arxiv.org/abs/1805.11604
- Xu et al., Understanding and Improving Layer Normalization, 2019 — arxiv.org/abs/1911.07013
- Zhang and Sennrich, Root Mean Square Layer Normalization, 2019 — arxiv.org/abs/1910.07467
- Wang et al., DeepNet, 2022 — arxiv.org/abs/2203.00555
What to learn next
- Pre-norm vs post-norm — where the norm goes matters more than which norm it is.
- RMSNorm — the simpler version most models use today.
- BatchNorm in PyTorch — the technique this one deliberately avoids.