LayerNorm vs RMSNorm
· 8 min read
LayerNorm
RMSNorm
I've been teaching myself how small language models work, and one detail kept showing up. GPT-2 and BERT normalize their activations with LayerNorm. LLaMA, Mistral and Gemma use RMSNorm instead. Both are one line of math, and the difference between them is a single subtraction. This is the explanation I wish I'd read first.
Why normalize at all
Inside a transformer, every token is a vector of numbers (its features), and that vector passes through dozens of layers. Each layer adds its output back onto the vector through a residual connection, so nothing naturally keeps its size in check. Left alone, the values drift bigger or smaller layer after layer, and training becomes unstable.
A normalization layer fixes that by rescaling each token's vector to a predictable size before it goes into attention or the MLP. Both LayerNorm and RMSNorm work on one token at a time, across its features. Unlike BatchNorm, they never look at other examples in the batch, so they behave the same during training and inference and for any batch size.
LayerNorm: re-center, then re-scale
Read it left to right:
- is one token's hidden vector, and is its -th feature. is the number of features, for example 768 in GPT-2 small.
- is the mean of this token's features. Subtracting it re-centers the vector around 0.
- is the variance of the features. Dividing by re-scales them so their spread is 1.
- is a tiny constant, usually around , that stops a division by zero when every feature is equal.
- (gain) and (bias) are learned per-feature parameters, starting at 1 and 0. They let the model undo the normalization wherever that helps.
Before and are applied, the output always has mean 0 and variance 1, whatever the input looked like.
RMSNorm: re-scale only
- is the root mean square: square every feature, average them, take the square root. It measures the vector's typical magnitude.
- There is no . The vector is never re-centered, only divided by its magnitude.
- There is no , only the learned gain . That's half the parameters of LayerNorm.
The output is guaranteed to have an RMS of 1, but its mean is not forced to 0. That's the entire difference, and the next formula shows when it matters.
The identity that connects them
Leaving out , the mean square splits into variance plus squared mean:
So when a token's features already average to about 0, and RMSNorm gives almost exactly the same answer as LayerNorm. When the mean is large, dominates the denominator. Every feature then gets divided by roughly , and all the outputs crowd together near . The differences between features, which carry the actual information, get squashed.
See it
Below are 50 features of one made-up token, before and after each normalization. Each dot is one feature, and the tick marks the mean. Try these:
- At a mean shift of 0, the LayerNorm and RMSNorm rows are identical.
- Still at shift 0, drag the variance. The raw row stretches, but neither output moves: both are scale-invariant.
- Push the mean shift to 5. LayerNorm doesn't care, because it subtracts the mean. RMSNorm's dots bunch up near 1 as the shift eats the RMS.
| series | mean | std | rms |
|---|---|---|---|
| raw input | -0.002 | 0.936 | 0.936 |
| layernorm | 0.000 | 1.000 | 1.000 |
| rmsnorm | -0.002 | 1.000 | 1.000 |
Run it yourself
Both are a few lines of numpy. These blocks run Python right in your browser: press Run. The first run takes a few seconds to download Python.
import numpy as np
def layer_norm(x, gamma, beta, eps=1e-5):
mu = x.mean(axis=-1, keepdims=True) # mean of the features
var = ((x - mu) ** 2).mean(axis=-1, keepdims=True) # variance of the features
x_hat = (x - mu) / np.sqrt(var + eps) # re-center, then re-scale
return gamma * x_hat + beta # learned scale and shift
x = np.array([2.0, 4.0, 6.0, 8.0])
d = x.shape[-1]
y = layer_norm(x, gamma=np.ones(d), beta=np.zeros(d))
print("input: ", x)
print("output:", y.round(4))
print(f"output mean: {y.mean():.4f} output std: {y.std():.4f}")The output has mean 0 and std 1 (to within ). The same function on [102, 104, 106, 108] would give the exact same output, because the shift is subtracted away.
import numpy as np
def rms_norm(x, gamma, eps=1e-5):
rms = np.sqrt((x ** 2).mean(axis=-1, keepdims=True) + eps) # typical magnitude
return gamma * (x / rms) # re-scale only
x = np.array([2.0, 4.0, 6.0, 8.0])
y = rms_norm(x, gamma=np.ones(x.shape[-1]))
print("input: ", x)
print("output:", y.round(4))
print(f"output mean: {y.mean():.4f} output rms: {np.sqrt((y ** 2).mean()):.4f}")The RMS is now 1, but the mean is not 0. The input was all positive, so the output is too. RMSNorm kept the shape and only changed the scale.
Now the experiment from the visualization. Take one random vector, add a growing constant to every feature, and watch what each norm does to the spread:
import numpy as np
def layer_norm(x, eps=1e-5):
return (x - x.mean()) / np.sqrt(x.var() + eps)
def rms_norm(x, eps=1e-5):
return x / np.sqrt((x ** 2).mean() + eps)
rng = np.random.default_rng(0)
x = rng.normal(0.0, 1.0, size=512) # one token's hidden vector
print(f"{'shift':>5} | {'LN std':>6} | {'RMS std':>7} | {'RMS mean':>8}")
for shift in [0, 1, 3, 10]:
xs = x + shift
ln, rn = layer_norm(xs), rms_norm(xs)
print(f"{shift:>5} | {ln.std():>6.3f} | {rn.std():>7.3f} | {rn.mean():>8.3f}")
# The identity that ties them together: RMS^2 = variance + mean^2
xs = x + 3
print()
print("RMS^2 == var + mean^2:", np.isclose((xs ** 2).mean(), xs.var() + xs.mean() ** 2))LayerNorm's output std stays at 1 for every shift. RMSNorm's std falls toward 0 while its mean climbs toward 1, which is the collapse you saw in the chart.
So why did LLMs switch to RMSNorm?
RMSNorm was introduced by Zhang and Sennrich in 2019. Their argument was that LayerNorm's benefit comes mostly from re-scaling, and that re-centering adds little. Dropping it keeps the model just as trainable, with less work:
- Less compute. One reduction (the mean of squares) instead of a mean followed by a variance, and no subtraction. Normalization runs twice in every transformer block, so this adds up.
- Fewer parameters. Just , with no .
- Same quality in practice. The paper reported accuracy comparable to LayerNorm with faster training, and later model families kept it.
The trade-off is the one the chart shows: RMSNorm isn't shift-invariant. It relies on the model not piling a large shared offset onto every feature, and in practice trained models are fine with that.
| LayerNorm | RMSNorm | |
|---|---|---|
| Subtracts the mean | yes | no |
| Divides by | std, | |
| Learned parameters | (2d) | (d) |
| Output guarantee | mean 0, variance 1 | RMS 1 |
| Scale-invariant | yes | yes |
| Shift-invariant | yes | no |
| Used in | original Transformer, BERT, GPT-2 | T5, LLaMA, Mistral, Gemma |
In PyTorch
You rarely write these by hand. Both ship with PyTorch. This one needs PyTorch installed, so it doesn't run in the browser:
import torch
from torch import nn
x = torch.randn(2, 8, 512) # (batch, tokens, features)
layer_norm = nn.LayerNorm(512)
rms_norm = nn.RMSNorm(512) # PyTorch 2.4+
print(layer_norm(x).shape, rms_norm(x).shape)The short version
LayerNorm re-centers and re-scales. RMSNorm only re-scales. When a token's features average to about 0, they give the same answer. RMSNorm is cheaper and trains just as well, which is why most modern LLMs use it.