Inside a Transformer Block

RMSNorm

Drop the recentring half of layer normalisation and keep only the resizing half - fewer parameters, less work, and no measurable loss in quality.

On this page 6
  1. Why this was worth doing
  2. What changes in practice
  3. Which models use which
  4. The honest part
  5. Remember this
  6. What to learn next

One lesson, three depths. Pick the one that fits you today — you can switch any time.

Beginner — No maths. Plain English.

RMSNorm is layer normalisation with the recentring step removed. Only the resizing remains.

Think about placing a photo on a page. There are two separate adjustments you might make. Resize it so it fits. Move it so it sits in the middle.

Layer normalisation does both. It shifts the numbers so their average lands on zero, then it resizes them to a standard spread.

RMSNorm asks a fair question: is the shifting doing any work? The answer, tested carefully in 2019, was no. So RMSNorm keeps the resizing and drops the shifting.

Why this was worth doing

Two savings, and neither is dramatic on its own.

Half the parameters. Layer normalisation carries two learnable pieces per slot, one for stretching and one for shifting. RMSNorm carries only the stretching one.

One less pass over the numbers. Computing the average, subtracting it, then measuring the spread is real work. Measuring the size directly is less.

Normalisation is not where a model spends most of its time. But it runs twice per layer, in every layer, on every token, for the entire life of the model. Small savings in that position add up.

The larger point is different, and better. Someone checked whether a step everyone had been using was actually necessary. It was not. That kind of subtraction is rarer and more valuable than addition.

What changes in practice

Almost nothing you can see. The two behave so similarly that models trained with either reach comparable quality.

There is one visible difference worth knowing. If a token's numbers are all identical, layer normalisation turns them into zero. Subtracting the average wipes them out. RMSNorm keeps them, because it never subtracts anything.

Whether that matters is not settled. It is the one place where the two genuinely disagree.

Which models use which

Layer normalisation: the original transformer, BERT, GPT-2, GPT-3, and most models up to about 2020.

RMSNorm: T5, the Llama family, and the great majority of open models since. If you open a recent config file and see rms_norm_eps, that is this.

The change spread quickly because it costs nothing to adopt and there is no evidence of a downside.

The honest part

The savings are small. Nobody switched to RMSNorm and saw their training bill halve.

And "faster" depends entirely on the code. A hand-written RMSNorm can easily be slower than a well-optimised layer normalisation. The built-in one may have years of tuning behind it. Yours does not. The developer section below shows exactly that happening, measured, on an ordinary laptop.

Speed claims about small operations are claims about implementations, not about mathematics. Measure on the hardware you will use.

Remember this

  • RMSNorm resizes but does not recentre.
  • It holds half the parameters of layer normalisation and does less work per token.
  • Quality is comparable, which is why almost every recent model uses it.

What to learn next

Developer — Code and libraries.

Setup

bash
pip install torch

nn.RMSNorm was added in PyTorch 2.4. On older versions, write the three lines yourself.

Verified against PyTorch, and timed honestly

rmsnorm.py
import time
import torch
import torch.nn as nn

torch.manual_seed(0)
torch.set_printoptions(precision=4, sci_mode=False)
d = 8

def rms_norm(x, weight, eps=1e-6):
    rms = x.pow(2).mean(dim=-1, keepdim=True).sqrt()   # no mean is subtracted anywhere
    return x / (rms + eps) * weight

x = torch.tensor([[3., 4., 0., 0., 0., 0., 0., 0.],    # rms = sqrt(25/8) = 1.7678
                  [5., 5., 5., 5., 5., 5., 5., 5.]])   # a token with a large offset

ref = nn.RMSNorm(d, eps=1e-6)
with torch.no_grad():
    ref.weight.copy_(torch.linspace(0.5, 1.5, d))
print("hand-written vs nn.RMSNorm, max difference:",
      float((rms_norm(x, ref.weight) - ref(x)).abs().max()))
print("torch version:", torch.__version__)

plain_rms = nn.RMSNorm(d, elementwise_affine=False)
plain_ln = nn.LayerNorm(d, elementwise_affine=False)
print("\ninput row 1:", x[0])
print("  RMSNorm  :", plain_rms(x)[0])
print("  LayerNorm:", plain_ln(x)[0])
print("input row 2:", x[1], " <- every value identical")
print("  RMSNorm  :", plain_rms(x)[1], " <- direction kept")
print("  LayerNorm:", plain_ln(x)[1], " <- flattened to zero")

print("\nwhat each one guarantees about its output:")
for name, fn in (("RMSNorm", plain_rms), ("LayerNorm", plain_ln)):
    y = fn(x)
    print(f"  {name:9s} mean per token: {y.mean(-1).tolist()}")
    print(f"  {name:9s} rms  per token: {y.pow(2).mean(-1).sqrt().tolist()}")

print("\nparameter count for width 4096:")
print("  LayerNorm:", sum(p.numel() for p in nn.LayerNorm(4096).parameters()), "(gain + bias)")
print("  RMSNorm  :", sum(p.numel() for p in nn.RMSNorm(4096).parameters()), "(gain only)")

big = torch.randn(64, 512, 4096)
for name, mod in (("LayerNorm", nn.LayerNorm(4096)), ("RMSNorm", nn.RMSNorm(4096))):
    with torch.no_grad():
        mod(big)                                        # warm up caches
        t0 = time.perf_counter()
        for _ in range(5):
            mod(big)
        print(f"  {name:9s} {1000*(time.perf_counter()-t0)/5:7.1f} ms per call "
              f"(this CPU only; your numbers will differ)")
Output
hand-written vs nn.RMSNorm, max difference: 5.960464477539062e-07
torch version: 2.5.1+cu121

input row 1: tensor([3., 4., 0., 0., 0., 0., 0., 0.])
  RMSNorm  : tensor([1.6971, 2.2627, 0.0000, 0.0000, 0.0000, 0.0000, 0.0000, 0.0000])
  LayerNorm: tensor([ 1.3834,  2.0345, -0.5697, -0.5697, -0.5697, -0.5697, -0.5697, -0.5697])
input row 2: tensor([5., 5., 5., 5., 5., 5., 5., 5.])  <- every value identical
  RMSNorm  : tensor([1., 1., 1., 1., 1., 1., 1., 1.])  <- direction kept
  LayerNorm: tensor([0., 0., 0., 0., 0., 0., 0., 0.])  <- flattened to zero

what each one guarantees about its output:
  RMSNorm   mean per token: [0.4949747622013092, 1.0]
  RMSNorm   rms  per token: [1.0, 1.0]
  LayerNorm mean per token: [4.470348358154297e-08, 0.0]
  LayerNorm rms  per token: [0.9999979138374329, 0.0]

parameter count for width 4096:
  LayerNorm: 8192 (gain + bias)
  RMSNorm  : 4096 (gain only)
  LayerNorm    49.0 ms per call (this CPU only; your numbers will differ)
  RMSNorm     152.0 ms per call (this CPU only; your numbers will differ)

The timing result is the most useful line on this page

RMSNorm does strictly less arithmetic than LayerNorm. On this machine it took three times longer.

That is not an error, and it is worth understanding rather than explaining away. Both route to native operations. nn.RMSNorm calls torch.rms_norm, added in PyTorch 2.4. But nn.LayerNorm's CPU path has had years of tuning that the newer one has not. Normalisation does very little arithmetic per byte it touches. It is limited by memory bandwidth rather than multiplications. Kernel quality decides the result almost entirely.

Rerun this on a CUDA device, or inside a compiled graph via torch.compile, and the ordering typically flips. The lesson generalises: for small memory-bound operations, published FLOP savings do not predict wall-clock time. Measure on your target.

Zhang and Sennrich report a 7 to 64 percent speedup in the original paper. It was measured on their models and their hardware. It is a real result about a real implementation, not a property of the formula.

The other numbers

5.96e-07 confirms the three-line version is correct. The whole operation is: square, mean over the last axis, square root, divide, multiply by the gain.

Row two is where the two genuinely differ. An input of eight identical 5.s becomes all 1.s under RMSNorm and all 0.s under LayerNorm. LayerNorm subtracts the mean, and for a constant vector the mean is the whole vector. RMSNorm preserves the direction and standardises only the length.

The guarantee table says what each one actually promises. RMSNorm guarantees the root-mean-square is 1 and says nothing about the mean. Row one gives 0.4949, row two gives 1.0. LayerNorm guarantees mean 0 and standard deviation 1, and for a constant input the output is degenerate.

Using it

python
import torch.nn as nn
norm = nn.RMSNorm(4096, eps=1e-5)          # PyTorch 2.4+

Written against PyTorch 2.5.1. The signature in the current 2.13 documentation is RMSNorm(normalized_shape, eps=None, elementwise_affine=True, device=None, dtype=None). Leaving eps as None uses the machine epsilon of the computation dtype. That is far smaller than the 1e-5 or 1e-6 real configs specify. Set it explicitly to match a checkpoint.

Before PyTorch 2.4:

python
class RMSNorm(nn.Module):
    def __init__(self, d, eps=1e-6):
        super().__init__()
        self.eps = eps
        self.weight = nn.Parameter(torch.ones(d))
    def forward(self, x):
        rms = x.pow(2).mean(-1, keepdim=True).add(self.eps).rsqrt()
        return x * rms * self.weight            # note: eps inside the sqrt, not outside

There are two conventions for where eps goes — inside the square root or added to the computed root. PyTorch and the Llama reference implementation both put it inside. The difference is invisible at normal magnitudes and shows up when reproducing a checkpoint's exact outputs.

Common mistakes

Initialising the gain to zero. It must start at one, or the layer outputs zeros and nothing downstream can learn. nn.RMSNorm does this correctly; a hand-written class often does not.

Keeping a bias. RMSNorm has no bias. Adding one back reintroduces the parameters you were removing and is not what any published RMSNorm model does.

Computing in half precision. x.pow(2) overflows float16 for values above about 256. Cast to float32 for the statistics and cast back. PyTorch's implementation handles this; a hand-written one in a mixed-precision training loop may not.

Swapping LayerNorm for RMSNorm in a pretrained checkpoint. The weights were trained with mean subtraction in place. Removing it changes every activation. This is an architecture choice made before training, not after.

Try it yourself

Add a large constant to every value of a token — say 100 — and pass it through both. LayerNorm's output is unchanged, because it subtracts the mean. RMSNorm's output changes a great deal, because the constant dominates the magnitude. That single experiment is the whole difference between the two.

What to learn next

Researcher — Mathematics and papers.

Definition

$$ \operatorname{RMS}(x) = \sqrt{\frac{1}{d} \sum_{i=1}^{d} x_i^2}, \qquad \operatorname{RMSNorm}(x)_i = \frac{x_i}{\operatorname{RMS}(x) + \epsilon} \gamma_i $$

with $\gamma \in \mathbb{R}^d$ learned. Compared with LayerNorm, the mean subtraction and the bias $\beta$ are both removed. PyTorch's documented form places $\epsilon$ inside the root:

$$ y_i = \frac{x_i}{\sqrt{\epsilon + \frac{1}{d}\sum_j x_j^2}} \gamma_i $$

Zhang and Sennrich (2019), Root Mean Square Layer Normalization, NeurIPS 2019, arXiv:1910.07467.

What is kept and what is discarded

LayerNorm has two invariances: to per-token shifts along $\mathbf{1}$, and to per-token rescaling. RMSNorm keeps only rescaling invariance:

$$ \operatorname{RMSNorm}(a x) = \operatorname{RMSNorm}(x) \ \text{ for } a > 0, \qquad \operatorname{RMSNorm}(x + b\mathbf{1}) \ne \operatorname{RMSNorm}(x) $$

The paper's central empirical claim is that the discarded invariance is not what produces LayerNorm's benefit. They report comparable quality across machine translation, image classification and question answering. Running-time reductions are 7 to 64 percent, depending on model and hardware. That range is wide precisely because the saving is implementation-dependent.

Geometrically, RMSNorm projects onto a sphere of radius $\sqrt{d}$ in $\mathbb{R}^d$. LayerNorm projects onto the same-radius sphere inside the hyperplane orthogonal to $\mathbf{1}$. Its image is a $(d{-}2)$-sphere rather than a $(d{-}1)$-sphere.

The gradient

With $r = \operatorname{RMS}(x)$ and $g = \partial \mathcal{L} / \partial y$:

$$ \frac{\partial \mathcal{L}}{\partial x} = \frac{\gamma \odot g}{r} - \frac{x}{d\, r^3} \left\langle \gamma \odot g,\ x \right\rangle $$

The second term removes the component of the incoming gradient along $x$ itself, mirroring the forward scale invariance. LayerNorm's gradient additionally removes the component along $\mathbf{1}$. So RMSNorm projects out one direction where LayerNorm projects out two.

Cost

Parameters: $d$ against $2d$. For a 32-layer, width-4096 model with two norms per layer, that is 262,144 against 524,288. Under 0.01 percent of an 8B model either way. This is not the reason to switch.

Compute: both are $O(d)$ per token and dominated by memory traffic, not arithmetic. RMSNorm needs one reduction pass. LayerNorm needs the mean and the variance. That is one pass with the sum-of-squares identity, and two in a naive implementation. The measured advantage therefore depends entirely on kernel quality. That is why the reported range spans an order of magnitude. It is also why the CPU measurement above runs the other way.

Adoption and the state of the evidence

T5 (Raffel et al., 2020) used RMSNorm without an affine gain. The Llama family, Mistral, Gemma, Qwen and most subsequent open models use RMSNorm with a gain. Llama 3 8B: rms_norm_eps $= 10^{-5}$.

The evidence base is thinner than the adoption rate suggests. Controlled ablations that hold everything else fixed are uncommon. RMSNorm is usually adopted alongside SwiGLU, rotary positions and pre-norm placement, and reported as a package. Narang et al. (2021), arXiv:2102.11972, is one of the few unified comparisons. RMSNorm is among the modifications that did reproduce there.

The reasonable summary is short. A well-motivated simplification with a clear mechanism. Comparable quality wherever tested. Efficiency gains that are real but implementation-specific.

Papers

What to learn next