Attention Mechanics

QK normalisation

Shrink every query and key to the same length before comparing them, so attention scores can never grow without limit no matter how far the weights drift during training.

On this page 6
  1. What goes wrong without it
  2. The fix
  3. Why this is not the same as the division you already met
  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.

QK normalisation makes every query and key the same length before they are compared. Only the direction they point in matters.

Stand two people in a field and ask them to point at the water tank. One has long arms and one has short arms. Neither fact tells you anything useful. What matters is whether they are pointing the same way.

If you record arm length as well as direction, the long-armed person seems to be "pointing harder". That is a measurement problem, not a fact about the water tank.

Attention has this exact problem. QK normalisation removes the arm length and keeps only the direction.

What goes wrong without it

During training, a model's internal numbers drift. Not because anything is broken — drifting is how learning happens.

One thing that drifts is the size of the query and key vectors. They tend to grow. As they grow, the scores between them grow with them, because a longer arm gives a bigger number.

Once scores get large, the step that turns them into shares stops mixing. It hands almost everything to a single winner. And a step that has already decided cannot be nudged, so learning through that head stops.

Then it gets worse. Very large scores overflow the number format a graphics card uses. The whole training run then collapses into an error. This has happened to well-funded teams training very large models. It is one reason a run can die a week in.

The fix

Before comparing a query with a key, shrink both to a fixed length. Now the comparison depends only on direction.

A comparison of directions is bounded. Two arrows pointing the same way give the largest possible score. Pointing opposite ways gives the smallest. There is nowhere further to go.

There is a cost to that. A bounded score that is always small would make attention too blurry to be useful. So one learnable dial is added, which multiplies the whole thing. The model chooses how sharp it wants attention to be, and that choice has a ceiling by construction.

   without QK normalisation
   weights drift bigger  ->  scores grow  ->  softmax picks one winner
                                          ->  learning through this head stops
                                          ->  eventually, overflow

   with QK normalisation
   weights drift bigger  ->  scores UNCHANGED
                          the only thing that can sharpen attention is
                          one dial the model trains deliberately

Why this is not the same as the division you already met

Scaled dot-product attention divides scores by a fixed number based on head width. That handles the size that arises from counting more slots.

It does not handle weights that grow during training. That number is fixed at the start and never adapts.

QK normalisation handles the drift. The two work together and neither replaces the other.

The honest part

This is not free. Forcing every query and key to the same length throws away information. The model can no longer signal how strongly it cares. It can signal only what it cares about.

In practice that loss has proven small and the stability gain large. Several recent large models use it for that reason. It is not universal. Plenty of strong models do without it and control drift by other means.

Remember this

  • Shrink queries and keys to a fixed length, so only direction counts.
  • Scores then have a hard ceiling, no matter how the weights drift.
  • One learnable dial restores sharpness, under a limit chosen by the design rather than by accident.

What to learn next

Developer — Code and libraries.

Setup

bash
pip install numpy

Drift, and the fix, measured side by side

qk_norm.py
import numpy as np

rng = np.random.default_rng(0)
d_head, T = 64, 16

def softmax(x, axis=-1):
    x = x - x.max(axis=axis, keepdims=True)
    e = np.exp(x)
    return e / e.sum(axis=axis, keepdims=True)

def entropy(p):
    return float(-(p * np.log(p + 1e-12)).sum(axis=-1).mean())

def l2norm(x):
    return x / np.linalg.norm(x, axis=-1, keepdims=True)

base_q = rng.normal(size=(T, d_head))
base_k = rng.normal(size=(T, d_head))

print("as training pushes the query/key weights larger, plain attention saturates.")
print("(entropy is in nats; ln(16) = 2.77 means 'looking everywhere equally')")
print(f"{'weight scale':>13} {'max |logit|':>12} {'top weight':>11} {'entropy':>9} "
      f"{'mean softmax slope':>19}")
for s in (1, 2, 4, 8, 16):
    q, k = base_q * s, base_k * s
    logits = q @ k.T / np.sqrt(d_head)
    p = softmax(logits)
    slope = (p * (1 - p)).mean()         # average diagonal of the softmax Jacobian
    print(f"{s:>13} {np.abs(logits).max():>12.1f} {p.max():>11.4f} "
          f"{entropy(p):>9.3f} {slope:>19.2e}")

print("\nsame weights, with QK normalisation (unit-length q and k, learnable g):")
print(f"{'weight scale':>13} {'max |logit|':>12} {'top weight':>11} {'entropy':>9} "
      f"{'mean softmax slope':>19}")
g = 10.0                                  # the learnable scale, fixed here so you can see it
for s in (1, 2, 4, 8, 16):
    q, k = l2norm(base_q * s), l2norm(base_k * s)
    logits = g * (q @ k.T)                # cosine similarity, so every entry is in [-g, g]
    p = softmax(logits)
    slope = (p * (1 - p)).mean()
    print(f"{s:>13} {np.abs(logits).max():>12.1f} {p.max():>11.4f} "
          f"{entropy(p):>9.3f} {slope:>19.2e}")

print(f"\nthe bound is exact: |logit| can never exceed g = {g}")
print("because a dot product of two unit vectors lies in [-1, 1].")

print("\nlarger g buys back sharpness without ever risking overflow:")
for g in (1, 5, 10, 20, 50):
    q, k = l2norm(base_q), l2norm(base_k)
    p = softmax(g * (q @ k.T))
    print(f"  g={g:>3}  top weight={p.max():.4f}  entropy={entropy(p):.3f}")
Output
as training pushes the query/key weights larger, plain attention saturates.
(entropy is in nats; ln(16) = 2.77 means 'looking everywhere equally')
 weight scale  max |logit|  top weight   entropy  mean softmax slope
            1          2.6      0.4111     2.332            5.39e-02
            2         10.6      0.9957     0.676            2.14e-02
            4         42.3      1.0000     0.148            5.48e-03
            8        169.2      1.0000     0.037            1.36e-03
           16        676.9      1.0000     0.001            1.74e-05

same weights, with QK normalisation (unit-length q and k, learnable g):
 weight scale  max |logit|  top weight   entropy  mean softmax slope
            1          3.3      0.5237     2.127            5.10e-02
            2          3.3      0.5237     2.127            5.10e-02
            4          3.3      0.5237     2.127            5.10e-02
            8          3.3      0.5237     2.127            5.10e-02
           16          3.3      0.5237     2.127            5.10e-02

the bound is exact: |logit| can never exceed g = 10.0
because a dot product of two unit vectors lies in [-1, 1].

larger g buys back sharpness without ever risking overflow:
  g=  1  top weight=0.0844  entropy=2.765
  g=  5  top weight=0.2297  entropy=2.592
  g= 10  top weight=0.5237  entropy=2.127
  g= 20  top weight=0.9252  entropy=1.182
  g= 50  top weight=0.9999  entropy=0.469

What the two tables prove

The first table is a controlled experiment. The same base vectors, multiplied by a growing scale. Nothing about the directions changed. Yet the maximum logit went 2.6, 10.6, 42.3, 169.2, 676.9. It grows with the square of the scale, since both operands grew.

Entropy collapsed from 2.332 to 0.001. With 16 keys, the largest possible entropy is about 2.77. At scale 16 the head is looking at exactly one token and ignoring the other fifteen. It has stopped being an attention mechanism.

The gradient signal fell by a factor of three thousand. The mean softmax slope went from 5.39e-02 to 1.74e-05. That column is the derivative of the shares with respect to the scores. Near zero, changes to the query and key projections barely move the output, so those weights stop being trained. A saturated head is not blurry. It is stuck.

The second table is five identical rows. Every column is unchanged across a sixteen-fold weight scale. Normalisation removed the magnitude completely, exactly as advertised, and the slope stayed at 5.10e-02.

The last block shows the dial doing its job. At g=1 attention is nearly uniform and useless. At g=50 it is a hard selection. The model learns where on this line to sit, and every point on the line is safe.

In PyTorch

python
import torch, torch.nn as nn, torch.nn.functional as F

class QKNormAttention(nn.Module):
    def __init__(self, d_model, n_heads):
        super().__init__()
        self.h, self.dh = n_heads, d_model // n_heads
        self.qkv = nn.Linear(d_model, 3 * d_model, bias=False)
        self.o   = nn.Linear(d_model, d_model, bias=False)
        self.q_norm = nn.RMSNorm(self.dh)      # per-head, over the head width only
        self.k_norm = nn.RMSNorm(self.dh)

    def forward(self, x):
        B, T, C = x.shape
        q, k, v = self.qkv(x).chunk(3, dim=-1)
        q, k, v = (t.view(B, T, self.h, self.dh).transpose(1, 2) for t in (q, k, v))
        q, k = self.q_norm(q), self.k_norm(k)  # normalise AFTER the head split
        a = F.scaled_dot_product_attention(q, k, v, is_causal=True)
        return self.o(a.transpose(1, 2).reshape(B, T, C))

m = QKNormAttention(64, 4)
print(m(torch.randn(2, 10, 64)).shape)
Output
torch.Size([2, 10, 64])

Written against PyTorch 2.5.1. nn.RMSNorm was added in PyTorch 2.4; before that, write the two lines by hand. Two variants are in circulation and both are called QK-norm:

  • The original, from Henry et al. (2020): true unit-length normalisation, then multiply by a single learned scalar. Logits are bounded by that scalar.
  • The common modern form: an RMSNorm with a learned per-channel gain applied to q and k. Logits are bounded only up to the size of that gain, so the bound is soft rather than hard. This is the version most open models ship, because it slots into an existing normalisation implementation.

Know which one a codebase means before you compare results with it.

Where to put the normalisation

After splitting into heads, over the head width. Normalising over the full model width before the split couples the heads together and defeats the purpose.

The head width is what appears in the scaling denominator, and it is what should be normalised. Getting this wrong produces code that runs, trains, and quietly loses the stability benefit.

Common mistakes

Keeping the division by the square root of head width as well. With unit-length queries and keys, the dot product is already bounded. The extra division makes the effective scale smaller than intended. Most implementations drop it and let the learned gain absorb everything.

Initialising the learned scale too small. At a scale near 1, attention starts almost uniform and the model spends early training doing nothing useful. Initialising near the value that the ordinary square-root scaling would have produced is the standard choice.

Normalising the values as well. Only queries and keys feed the dot product. Values are mixed linearly afterwards, and normalising them discards magnitude information the model actually uses.

Adding it to a model that was not diverging. QK normalisation costs a little throughput and a little expressivity. It is a fix for a real symptom: growing logits, collapsing attention entropy, loss spikes. Log those first.

Try it yourself

Add a nan check. Set the weight scale to 400 and compute the plain logits in float16, then in float32. Watch where each one overflows. Then repeat with normalisation and confirm float16 survives at every scale. That is the practical difference in one experiment.

What to learn next

Researcher — Mathematics and papers.

Definition

The original formulation, Henry et al. (2020), Query-Key Normalization for Transformers, Findings of EMNLP 2020, arXiv:2010.04245:

$$ A = \operatorname{softmax}!\left( g \cdot \frac{Q}{|Q|_2} \left( \frac{K}{|K|_2} \right)^{!\top} \right) $$

with the $L_2$ norms taken along the head dimension and $g$ a single learned scalar per head. Each logit is then $g \cos\theta_{ij}$ for the angle $\theta_{ij}$ between query $i$ and key $j$, so

$$ -g \le s_{ij} \le g $$

exactly. The $1/\sqrt{d_k}$ factor is dropped, since the width-dependent growth it corrected no longer exists.

The paper's own framing is that this makes softmax "less prone to arbitrary saturation without sacrificing expressivity". It reports an average gain of 0.928 BLEU across five low-resource translation pairs. The initialisation of $g$ matters in practice. Too small and attention starts uniform. Too large and it starts already saturated. Consult the implementation you are copying rather than assuming a value.

What it prevents

Attention logits are $s_{ij} = x_i^\top W_Q W_K^\top x_j / \sqrt{d_k}$, so $|s_{ij}| \le |x_i| |x_j| \sigma_{\max}(W_Q W_K^\top) / \sqrt{d_k}$. Nothing in standard training bounds the residual-stream norm or the spectral norm of the QK circuit. Both are empirically observed to grow. Dehghani et al. (2023), arXiv:2302.05442, report divergence at scale. It was traced to attention logits reaching order $10^4$ within a few thousand steps. QK normalisation removed it.

Zhai et al. (2023), Stabilizing Transformer Training by Preventing Attention Entropy Collapse, arXiv:2303.06296, give the sharper framing. Define per-row attention entropy

$$ H_i = -\sum_j A_{ij} \log A_{ij} $$

They show empirically that a sustained drop in mean $H_i$ toward zero precedes loss divergence. Entropy is therefore a usable early-warning signal. Their own proposal, $\sigma$Reparam, bounds the spectral norm of the projections rather than normalising the activations. That is a different intervention on the same quantity.

The two variants, and why the distinction matters

Most current open implementations apply RMSNorm with a learned per-channel gain $\gamma \in \mathbb{R}^{d_k}$ to $q$ and $k$. That is not plain $L_2$ normalisation with a scalar. Then

$$ |s_{ij}| \le \frac{d_k |\gamma|_\infty^2}{\sqrt{d_k}} = \sqrt{d_k}\, |\gamma|_\infty^2 $$

which is a bound that grows if $\gamma$ grows. The elementwise-affine version is therefore not an unconditional guarantee; it is a strong prior toward bounded logits. Papers frequently cite Henry et al. while implementing the affine version. When reproducing a result, read the code.

Cost

Two extra normalisations per attention layer, each $O(T d_k)$ per head. That is negligible against the $O(T^2 d_k)$ of attention and the $O(T d^2)$ of the projections. The real cost is throughput. The normalisation sits between the projection and the attention kernel, and can prevent fusion of the two.

Expressivity cost is genuine but hard to measure. The head loses access to $|q_i|$ as a per-token confidence signal. In an unnormalised head that signal modulates row sharpness independently of direction.

  • Cosine attention in Swin Transformer V2 (Liu et al., 2022, arXiv:2111.09883) is the same construction. It was motivated by activation-magnitude blow-up at 3B parameters.
  • Logit soft-capping, $s \mapsto c \tanh(s/c)$, used in the Gemma 2 technical report, bounds smoothly and keeps a nonzero gradient everywhere. It is incompatible with several fused attention kernels, which limited its uptake.
  • nGPT (Loshchilov et al., 2024, arXiv:2410.01131) normalises every vector in the network to the unit hypersphere. QK normalisation is one component of that.

An honest caveat about adoption

QK normalisation is common in models trained after roughly 2023. Many strong models trained before and since omit it. Ablations that isolate it from co-occurring changes are rare in technical reports. Those changes include learning-rate schedule, initialisation scale and normalisation placement. Treat it as a well-motivated stabiliser with a clear mechanism. Controlled evidence about its effect on final quality is thin.

Papers

What to learn next