Tokenwerk · The LLM textbook

Chapter 25 · V · Modern language models · 7 minutes

RMSNorm, RoPE and SwiGLU

Three targeted changes to the decoder. You learn the motivation, the calculation and the limits of each technique.

Modernizing through controlled swaps

We don't change all components at once and then claim that each one helped. The reference code offers a modern configuration; for scientific comparisons you add individual switches and keep the parameter or compute budget as comparable as possible. Below, we derive the basic calculations ourselves.

RMSNorm

RMSNorm normalizes by the square root of the mean square of the components. Unlike LayerNorm, it doesn't subtract a mean: First compute [3,4]: squaring gives [9,16], averaging gives 12.5 and taking the square root gives about 3.536. Divide each original component by this number. Only then comes the general notation:

RMSNorm⁡(x)i=γixi1D∑jxj2+ϵ.\operatorname{RMSNorm}(x)_i=\gamma_i\frac{x_i}{\sqrt{\frac1D\sum_jx_j^2+\epsilon}}.

Here x is the list of numbers for one position, i the selected component and D its length. The index j runs over all D components. The sum of the squares, divided by D, is their mean square. Epsilon is again a small safety constant. Gamma is a trainable multiplier per component. RMS stands for "root mean square".

For x=[3,4]x=[3,4] the RMS is (9+16)/2≈3.536\sqrt{(9+16)/2}\approx3.536. With a scale of one, this gives about [0.849,1.131][0.849,1.131]. The mean of the output is not zero. That is the essential difference from a centering normalization.

python
class RMSNorm(nn.Module):
    def __init__(self, dim, eps=1e-6):
        super().__init__()
        self.weight = nn.Parameter(torch.ones(dim))
        self.eps = eps

    def forward(self, x):
        xf = x.float()
        y = xf * torch.rsqrt(xf.square().mean(-1, keepdim=True) + self.eps)
        return y.to(x.dtype) * self.weight.to(x.dtype)

Here the statistic is computed in float32. That helps with low-precision activations. RMSNorm is not extra storage for facts; it changes the scaling and the optimization properties of the computation path. Original paper on RMSNorm.

Position as rotation

RoPE doesn't add a learned position table to the input embedding. It rotates pairs of components in Q and K depending on their position. Think of [1,0] as an arrow pointing to the right. A quarter turn turns it into [0,1], an arrow pointing up. The length stays 1.

Sine and cosine give the coordinates of a rotated arrow of length 1: cosine the horizontal coordinate, sine the vertical one. At zero degrees they are 1 and 0; at 90 degrees they are 0 and 1. In code, angles are measured in radians: a full turn is 2π, a quarter turn π/2. π (pi) is about 3.14159. For a pair [x1,x2][x_1,x_2] and angle θ\theta:

Rθx=[x1cos⁡θ−x2sin⁡θ, x1sin⁡θ+x2cos⁡θ].R_\theta x=[x_1\cos\theta-x_2\sin\theta,\ x_1\sin\theta+x_2\cos\theta].

R stands for rotation, theta (θ\theta) for the rotation angle, and x₁ and x₂ for the two original components. For a quarter turn, plug in cos θ = 0 and sin θ = 1: [1,0] becomes [0,1]. The vector norm is the length of the arrow, here the square root of the sum of the two squared components. A rotation preserves this vector norm. Because of the rotations, Q at position mm and K at position nn have a dot product that depends on the relative angle, that is, on n−mn-m. This explains why relative position relationships enter the attention scores.

For the feature pairs, we use different frequencies. A common form is ωi=base−2i/d\omega_i=base^{-2i/d}. Here i is the number of the component pair, starting at 0, d the full head width and base a fixed positive base greater than 1. Omega (ωi\omega_i) is the rotation per position step for this pair. The negative exponent means a reciprocal: base to the power of −1 is 1 / base. For example, base = 100 and d = 4 give frequency 1 for pair 0 and frequency 0.1 for pair 1. The angle at position t is tωit\omega_i. At position 3, in the example, that is 3 and 0.3 radians respectively. High and low frequencies capture different ranges of distance. Our implementation uses neighboring even/odd components as a pair; other implementations may arrange the features differently.

python
even, odd = x[..., 0::2], x[..., 1::2]
y_even = even * cos - odd * sin
y_odd  = even * sin + odd * cos
y = torch.stack((y_even, y_odd), dim=-1).flatten(-2)

cos and sin must broadcast correctly over the batch and head axes. The head dimension must be even. In the reference model, we rotate Q and K, not V. RoFormer / RoPE.

RoPE doesn't automatically make context unlimited

The formula can compute positions beyond the training window. That doesn't mean reliable performance there. The model can work worse at unfamiliar distances. Longer contexts need suitable training data, adjustments and evaluation. Common adjustments scale the RoPE frequencies (e.g. position interpolation or YaRN) and briefly continue training on longer texts. In the course project, even the modern decoder stays limited to its configured context length.

SwiGLU

A classic MLP has two projections with a nonlinearity in between. SwiGLU uses two input projections, multiplies them component by component and then projects back:

SwiGLU⁡(x)=Wdown(SiLU⁡(Wgatex)⊙Wupx).\operatorname{SwiGLU}(x)=W_{down}\left(\operatorname{SiLU}(W_{gate}x)\odot W_{up}x\right).

The symbol ⊙\odot means component-wise multiplication: [2,3] and [4,5] give [8,15], not a dot product. W_gate and W_up are two different learned weight tables, each of which computes a wider list from x. W_down then makes the list as wide as x again. "Gate" means just that: one path of numbers changes the strength of the other.

SiLU is xσ(x)x\sigma(x). Sigma (σ\sigma) stands for a function here, not for the standard deviation from LayerNorm. This function is called sigmoid and computes σ(x)=1/(1+exp⁡(−x))\sigma(x)=1/(1+\exp(-x)). It maps any number to a value between 0 and 1. At x = 0 it gives 1/2; at x = 1 about 0.731. SiLU multiplies this value by x. So for x = 1, SiLU gives about 0.731, and for x = −1 about −0.269. Example for a single gate output of 1 and an up output of 2: their product after SiLU is about 1.462, before W_down is applied. One path acts as a learned gate, the other supplies values. This gate is not a binary switch; its components can act continuously and also negatively.

python
h = F.silu(self.gate(x)) * self.up(x)
out = self.down(h)

There are three large matrices instead of two. For a comparable parameter count, the hidden width is therefore often chosen smaller than the classic 4D4D width, roughly 8D/38D/3, and then rounded to a practical block size. The reference project rounds to a multiple of 32. GLU variants.

A modern training run

bash
python train.py --data runs/data --out runs/modern \
  --modern --dim 128 --layers 4 --heads 4 --kv-heads 2 \
  --context 128 --steps 2000

This also enables GQA, which we explain in the next chapter. This small implementation illustrates established building blocks, but it contains no distributed training pipeline or highly optimized inference engine.

What you should compare

Record the actual parameter count, processed tokens, runtime, validation loss and fixed samples. If several parts are swapped at once, you can only interpret the overall comparison. To determine the effect of one component, you need an ablation: the same experimental setup with exactly this one difference.