1. What is it?#
An operation that rescales activations so their magnitude stays in a predictable range.
LayerNorm: y = (x - mean) / sqrt(var + eps) · γ + β (over the feature dim)
RMSNorm: y = x / sqrt(mean(x²) + eps) · γ (no mean, no bias)Both are applied per-token across the hidden dimension. Both are memory-bound. Both appear twice per transformer layer.
2. Why does it exist?#
Without normalization, deep networks are hard to train: activations drift in scale, gradients explode or vanish. Normalization pins the scale at every layer, and pre-norm placement gives a clean residual path.
At inference the training benefits are already baked in — but the operation still runs, twice per layer, 160 times for an 80-layer model, and it touches the full activation tensor each time. Its cost is entirely a memory-traffic story.
3. Simple analogy#
A volume normalizer on a mixing desk. Each channel’s loudness is scaled to a standard level
before mixing, so no single instrument drowns the others regardless of how it was recorded. The
learned γ is the engineer’s per-channel preference applied afterward.
4. Tiny example#
package main
import (
"fmt"
"math"
)
func main() {
x := []float64{1, 2, 3, 4}
const eps = 1e-6
n := float64(len(x))
// LayerNorm
var mean, variance float64
for _, v := range x {
mean += v / n // 2.5
}
for _, v := range x {
variance += (v - mean) * (v - mean) / n // 1.25
}
fmt.Print("layernorm:")
for _, v := range x {
fmt.Printf(" %.4f", (v-mean)/math.Sqrt(variance+eps)) // gamma = 1
}
fmt.Println() // -1.3416 -0.4472 0.4472 1.3416 mean 0, std 1
// RMSNorm
var meanSq float64
for _, v := range x {
meanSq += v * v / n
}
rms := math.Sqrt(meanSq + eps) // sqrt(7.5) = 2.7386
fmt.Print("rmsnorm: ")
for _, v := range x {
fmt.Printf(" %.4f", v/rms)
}
fmt.Println() // 0.3651 0.7303 1.0954 1.4606 mean NOT 0, RMS = 1
}
RMSNorm doesn’t center the data. Empirically that doesn’t hurt, and it saves a reduction and a subtraction — which, for a memory-bound op, means one fewer pass over the tensor.
5. Technical explanation#
Why normalization is memory-bound#
For a (B, S, d) tensor in BF16:
Bytes: read 2·B·S·d + write 2·B·S·d + read γ (2d, cached) ≈ 4·B·S·d
FLOPs: ~5·B·S·d (square, sum, rsqrt, multiply, multiply)
Intensity = 5/4 = 1.25 FLOP/byteIntensity 1.25 against a ridge point of 296. Utterly memory-bound. The only optimization that matters is not making a separate pass over the data — i.e., fusion.
The naive implementation’s cost#
Unfused RMSNorm as separate PyTorch ops:
x.pow(2) read x, write tmp1
tmp1.mean(-1) read tmp1, write tmp2
tmp2 + eps read/write tmp2
torch.rsqrt(tmp2) read/write tmp2
x * tmp2 read x, read tmp2, write tmp3
tmp3 * gamma read tmp3, write out
→ ~6 passes over the dataA fused kernel does one pass:
load x tile into registers/shared
compute sum of squares (warp reduction)
rsqrt
multiply by gamma
store
→ 1 read + 1 write6x less memory traffic. For a 70B model at batch 64, that’s the difference between the norms costing 15% of your step time and 2.5%.
Fusion with neighbors#
The bigger win is fusing the norm with what surrounds it:
Typical sequence: residual_add → rmsnorm → linear
Fused: one kernel that reads x and residual, adds, normalizes,
and writes the normalized output (and optionally the new
residual for the next layer)
Even better: fuse the norm into the epilogue of the PREVIOUS matmul,
or the prologue of the NEXT oneProduction engines do “add_rms_norm” as a single kernel. It’s one of the first custom kernels any engine writes.
Numerical considerations#
Compute the reduction in FP32 even for BF16 tensors.
Summing 8192 squared BF16 values in BF16 loses precision badly
(BF16 has 8 mantissa bits; the sum's magnitude grows as d).Also: eps placement matters. sqrt(mean + eps) vs sqrt(mean) + eps differ, and different
frameworks have historically differed. When porting a model, match the reference exactly or you
get subtle quality drift.
Other normalizations you may meet#
BatchNorm normalizes over the batch. NOT used in transformers — it makes
inference batch-dependent, which is unacceptable for serving.
GroupNorm normalizes over channel groups. Used in diffusion U-Nets.
QK-Norm RMSNorm applied to q and k before attention. Improves stability
at scale; appears in newer models. Adds two more memory-bound ops.BatchNorm’s absence from transformers is worth understanding: at inference it uses running statistics (so it’s fine), but during training it couples examples in a batch. For serving, any op whose output depends on batch composition is a nightmare — results would change based on who else’s request was batched with yours.
6. Under the hood#
A fused RMSNorm CUDA kernel, structurally:
one thread block per token (row of length d)
each thread loads d/blockDim elements into registers
compute partial sum of squares
warp-level reduction (__shfl_down_sync)
block-level reduction via shared memory
thread 0 computes rsqrt, broadcasts via shared memory
each thread scales its elements by (rsqrt × gamma) and storesVectorized loads (float4 / half8) matter: loading 16 bytes per instruction instead of 2
reduces instruction count 8x and improves memory coalescing.
7. Performance implications#
Measured share of decode step time for a 70B model:
Unfused norms: 12-18%
Fused norms: 3-5%
Fused with residual: 2-3%That’s a 10-15% end-to-end win from fusing an operation that is 0.1% of your FLOPs. This is the canonical example of why memory-bound thinking matters.
8. Production implications#
- Verify your engine uses fused norms. Look for
rms_norm_kerneloradd_rms_norm_kernelin a profile, not a chain ofpow/mean/rsqrt/mul. - Match the reference implementation exactly when porting: eps value and placement, whether the reduction is in FP32, whether γ is applied before or after.
- QK-Norm adds cost. If a new model uses it, budget for two more memory-bound ops per layer.
- Never use BatchNorm in a serving path where batch composition varies.
9. Common mistakes#
Leaving norms unfused. 10-15% left on the table.
Reducing in BF16. Precision loss that shows up as quality drift.
Mismatched eps. Small but real output differences when porting.
Assuming norms are negligible because their FLOPs are. They’re 3-18% of time.
Using LayerNorm where the model specifies RMSNorm (or vice versa). Different parameter shapes will usually error; but a model with γ-only RMSNorm loaded into a LayerNorm with β=0 silently computes something different.
10. Hands-on exercise#
A. Measure the fusion win. Implement RMSNorm as (i) separate PyTorch ops, (ii)
torch.compiled, (iii) an existing fused kernel (e.g. from flash_attn or your engine).
Benchmark all three on a (64, 1, 8192) tensor. Report bytes/sec achieved for each.
B. Precision. Compute RMSNorm on a large BF16 tensor with BF16 accumulation and with FP32
accumulation. Measure the difference. At what d does it become significant?
C. Write the kernel. After Section VI, come back and write a fused RMSNorm CUDA or Triton kernel with vectorized loads. Compare to PyTorch’s.
D. Profile a real model. What percentage of decode step time is spent in normalization kernels? Is it fused?
11. Interview questions#
- Why is normalization memory-bound despite doing very few FLOPs?
- What is the difference between LayerNorm and RMSNorm, in parameters and in cost?
- Why is BatchNorm unsuitable for transformer inference?
- How much does fusing normalization typically save, and why is that surprising given its FLOP share?
- Why must the reduction be computed in FP32?
- What is QK-Norm and what does it cost?
12. Further reading#
- [ESTABLISHED] Zhang & Sennrich, “Root Mean Square Layer Normalization” (2019)
- [ESTABLISHED] Xiong et al., “On Layer Normalization in the Transformer Architecture” (2020) — pre-norm vs post-norm
- [REFERENCE]
flash_attn’slayer_normfused kernels; vLLM’slayernorm.cu - Next: 06 — Attention computation in practice