Course Content
Transformer Architecture Q&A
6 sections · 60 lessons
Why is layer norm used in Transformers instead of batch norm? Give a concrete reason?
What you need to know
The difference is which numbers get averaged. For a tensor of shape (B, T, d_model):
LayerNorm : mean and variance over d_model, separately for each (batch, position)BatchNorm : mean and variance over (batch, position), separately for each featureWorked by hand: one token, d_model = 4
1import torch23x = torch.tensor([[2.0, 4.0, 6.0, 8.0]]) # one token, d_model = 445ln = torch.nn.LayerNorm(4, elementwise_affine=False)6rms = torch.nn.RMSNorm(4, elementwise_affine=False)7print(ln(x)) # subtract mean 5, divide by std sqrt(5)8print(rms(x)) # no mean subtraction, divide by sqrt(mean of squares) = sqrt(30)9# tensor([[-1.3416, -0.4472, 0.4472, 1.3416]])10# tensor([[0.3651, 0.7303, 1.0954, 1.4606]])LayerNorm: mean is 5, variance (9 + 1 + 1 + 9) / 4 = 5, so (x - 5) / 2.236. RMSNorm: root-mean-square is sqrt((4 + 16 + 36 + 64) / 4) = sqrt(30) = 5.477, so x / 5.477. Both then multiply by a learned gain (LayerNorm also adds a learned bias). Neither needs any other token.
Why BatchNorm fails here
- Padding. Statistics over
(batch, position)include padding slots, and the amount of padding changes every batch. - Uneven positions. In a batch where one sequence has 2,000 tokens and the rest 200, position 1,500 appears once; its statistics are noise.
- Inference. A decoder generates one token at a time, often for one user. BatchNorm must use running averages from training, which do not match this distribution, and the mismatch shows up in generated text.
LayerNorm is a per-token function, so batch size 1 and batch size 256 give identical results for the same token.
RMSNorm and placement in 2026
RMSNorm drops the mean subtraction and bias. It is cheaper and works as well in practice, so Llama, Qwen, Mistral and most current LLMs use it. Placement matters as much as type: modern models use pre-norm (x + f(Norm(x))). Some also add QK-norm — normalising queries and keys before the dot product — to keep attention scores from growing too large in big training runs.
A real-life example
A speech-to-text team builds a streaming transcriber for a customer-support line. In production each call is transcribed live, one audio chunk at a time, batch size 1. An early prototype used a convolutional front-end with BatchNorm inside the Transformer layers, copied from a vision model.
Offline accuracy on batched test files was good. In the live system, accuracy dropped, and it varied with how many calls happened to be batched together on the server. The BatchNorm running statistics simply did not describe single-call streaming audio. Replacing BatchNorm with LayerNorm in the Transformer layers made results identical whether one call or thirty were processed together — the concrete property interviewers want you to name.
Follow-up questions to expect
- "Does BatchNorm ever appear in speech or vision Transformers?" — Some convolutional front-ends use it, where inputs are fixed-size frames; the Transformer blocks themselves use LayerNorm or RMSNorm.
- "Why is RMSNorm enough without centring?" — Experiments showed re-scaling, not re-centring, gives most of the benefit, and removing the mean saves compute.
- "Pre-norm or post-norm?" — Pre-norm trains more stably in deep models and is the default; post-norm needs careful warm-up, though some studies report slightly better final quality when it does train.