Taj
← all writing

LayerNorm vs RMSNorm

· 8 min read

LayerNorm

yi=γi xi−μσ2+ϵ+βiy_i = \gamma_i \, \frac{x_i - \mu}{\sqrt{\sigma^2 + \epsilon}} + \beta_i

RMSNorm

yi=γi xiRMS⁡(x)y_i = \gamma_i \, \frac{x_i}{\operatorname{RMS}(x)}

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 dd 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 dd 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

μ=1d∑j=1dxjσ2=1d∑j=1d(xj−μ)2\begin{gathered} \mu = \frac{1}{d}\sum_{j=1}^{d} x_j \\[6pt] \sigma^2 = \frac{1}{d}\sum_{j=1}^{d} (x_j - \mu)^2 \end{gathered}
yi=γi xi−μσ2+ϵ+βiy_i = \gamma_i \, \frac{x_i - \mu}{\sqrt{\sigma^2 + \epsilon}} + \beta_i

Read it left to right:

  • xx is one token's hidden vector, and xix_i is its ii-th feature. dd is the number of features, for example 768 in GPT-2 small.
  • μ\mu is the mean of this token's features. Subtracting it re-centers the vector around 0.
  • σ2\sigma^2 is the variance of the features. Dividing by σ\sigma re-scales them so their spread is 1.
  • ϵ\epsilon is a tiny constant, usually around 10−510^{-5}, that stops a division by zero when every feature is equal.
  • γ\gamma (gain) and β\beta (bias) are learned per-feature parameters, starting at 1 and 0. They let the model undo the normalization wherever that helps.

Before γ\gamma and β\beta are applied, the output always has mean 0 and variance 1, whatever the input looked like.

RMSNorm: re-scale only

RMS⁡(x)=1d∑j=1dxj2+ϵyi=γi xiRMS⁡(x)\begin{gathered} \operatorname{RMS}(x) = \sqrt{\frac{1}{d}\sum_{j=1}^{d} x_j^2 + \epsilon} \\[6pt] y_i = \gamma_i \, \frac{x_i}{\operatorname{RMS}(x)} \end{gathered}
  • RMS⁡(x)\operatorname{RMS}(x) is the root mean square: square every feature, average them, take the square root. It measures the vector's typical magnitude.
  • There is no μ\mu. The vector is never re-centered, only divided by its magnitude.
  • There is no β\beta, only the learned gain γ\gamma. 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 ϵ\epsilon, the mean square splits into variance plus squared mean:

RMS⁡(x)2=σ2+μ2\operatorname{RMS}(x)^2 = \sigma^2 + \mu^2

So when a token's features already average to about 0, RMS⁡(x)≈σ\operatorname{RMS}(x) \approx \sigma and RMSNorm gives almost exactly the same answer as LayerNorm. When the mean is large, μ2\mu^2 dominates the denominator. Every feature then gets divided by roughly ∣μ∣|\mu|, and all the outputs crowd together near ±1\pm 1. 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:

  1. At a mean shift of 0, the LayerNorm and RMSNorm rows are identical.
  2. Still at shift 0, drag the variance. The raw row stretches, but neither output moves: both are scale-invariant.
  3. 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.
seriesmeanstdrms
raw input-0.0020.9360.936
layernorm0.0001.0001.000
rmsnorm-0.0021.0001.000
Drag the mean shift. LayerNorm keeps a mean of 0 and a std of 1 no matter what. RMSNorm only fixes the rms at 1, so a big shift squashes every value toward the same number.

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.

layer_norm.py
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 ϵ\epsilon). The same function on [102, 104, 106, 108] would give the exact same output, because the shift is subtracted away.

rms_norm.py
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:

shift_experiment.py
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 γ\gamma, with no β\beta.
  • 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.

LayerNormRMSNorm
Subtracts the meanyesno
Divides bystd, σ\sigmaRMS⁡(x)\operatorname{RMS}(x)
Learned parametersγ,β\gamma, \beta (2d)γ\gamma (d)
Output guaranteemean 0, variance 1RMS 1
Scale-invariantyesyes
Shift-invariantyesno
Used inoriginal Transformer, BERT, GPT-2T5, 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:

pytorch
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.