LayerNorm, RMSNorm and pre-/post-norm
Be able to compare normalisation variants and placements in modern transformers.
Prerequisites
- DTransformers — the architecturerequired
- EBatch and layer normalisationrequired
Intuition
LayerNorm normalises over the features within each example: subtract the mean, divide by the standard deviation, scale and shift with learnt γ and β.
RMSNorm skips the mean centring and normalises only by the root mean square:
One statistic instead of two, no β. About 10–15 % faster, and empirically just as good — which is why Llama, Mistral and Qwen use RMSNorm.
Formal
The placement matters more than the variant.
Post-norm (the original transformer):
Pre-norm (all the modern ones):
The difference looks cosmetic but is decisive. In pre-norm there is a clean residual path from the input to the output that does not pass through any normalisation. The gradient can flow back along it without being rescaled layer by layer.
The consequence (Xiong et al. 2020): post-norm transformers require warm-up in order not to diverge, and become unstable at great depth. Pre-norm trains stably without warm-up and scales to hundreds of layers.
The price is a small quality loss at the same depth — post-norm gives a somewhat better result once it does converge. The solution in practice is pre-norm plus a final normalisation before the output.
Variants worth knowing about: DeepNorm (a scaled residual that makes very deep post-norm networks possible) and sandwich norm (normalisation both before and after the sublayer).
Code
import torch, torch.nn as nn
class RMSNorm(nn.Module):
def __init__(self, d, eps=1e-6):
super().__init__()
self.gamma = nn.Parameter(torch.ones(d)); self.eps = eps
def forward(self, x):
rms = x.pow(2).mean(-1, keepdim=True).add(self.eps).rsqrt()
return x * rms * self.gamma
class PreNormBlock(nn.Module):
"""The modern order: normalise BEFORE the sublayer, add to the residual."""
def __init__(self, d, n_heads):
super().__init__()
self.n1, self.n2 = RMSNorm(d), RMSNorm(d)
self.attn = nn.MultiheadAttention(d, n_heads, batch_first=True)
self.mlp = nn.Sequential(nn.Linear(d, 4*d), nn.GELU(), nn.Linear(4*d, d))
def forward(self, x, mask=None):
h = self.n1(x)
x = x + self.attn(h, h, h, attn_mask=mask, need_weights=False)[0]
return x + self.mlp(self.n2(x))
x = torch.randn(2, 8, 64)
print(PreNormBlock(64, 8)(x).shape) # torch.Size([2, 8, 64])
print(RMSNorm(64)(x).pow(2).mean(-1).mean().item()) # ≈ 1.0 (RMS normalised)
Note that the residual addition happens on x, not on the normalised h — that is precisely what preserves the clean residual path.
Mastery means
- Compares LayerNorm and RMSNorm
- Explains the difference between pre-norm and post-norm
- Knows why pre-norm became the standard
Sign in to do the exercises and build your mastery up.
Sources
- arXiv — On Layer Normalization in the Transformer Architecture — arXiv (open access; licence per article)
- arXiv — Root Mean Square Layer Normalization — arXiv (open access; licence per article)