Skip to content
AI-grafen
EUniversityTransformer architecture· about 60 min· evolving, reviewed regularly· verified 2026-09-20· EN

LayerNorm, RMSNorm and pre-/post-norm

Be able to compare normalisation variants and placements in modern transformers.

Prerequisites

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:

RMSNorm(x)=x1d∑ixi2+ϵ⊙γ\text{RMSNorm}(x) = \frac{x}{\sqrt{\frac{1}{d}\sum_i x_i^2 + \epsilon}}\odot\gamma

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): x←LN(x+Sublayer(x))x \leftarrow \text{LN}(x + \text{Sublayer}(x))

Pre-norm (all the modern ones): x←x+Sublayer(LN(x))x \leftarrow x + \text{Sublayer}(\text{LN}(x))

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

All the sources and licences