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.
- 14 min read
- 3 reading levels
- Updated
Read these first
On this page 6
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 deliberatelyWhy 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
- Walking through one transformer block — where the attention layer you have now taken apart actually sits.
- Layer normalisation — the normalisation this one is built on top of.
- Monitoring and drift — spotting instability before a run dies.
Developer — Code and libraries.
Setup
pip install numpyDrift, and the fix, measured side by side
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}")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.469What 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
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)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
qandk. 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
- Walking through one transformer block — where the attention layer you have now taken apart actually sits.
- Layer normalisation — the normalisation this one is built on top of.
- Monitoring and drift — spotting instability before a run dies.
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.
Related bounding methods
- 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
- Henry et al., Query-Key Normalization for Transformers, 2020 — arxiv.org/abs/2010.04245
- Liu et al., Swin Transformer V2, 2022 — arxiv.org/abs/2111.09883
- Dehghani et al., Scaling Vision Transformers to 22 Billion Parameters, 2023 — arxiv.org/abs/2302.05442
- Zhai et al., Stabilizing Transformer Training by Preventing Attention Entropy Collapse, 2023 — arxiv.org/abs/2303.06296
- Loshchilov et al., nGPT: Normalized Transformer, 2024 — arxiv.org/abs/2410.01131
What to learn next
- Walking through one transformer block — where the attention layer you have now taken apart actually sits.
- Layer normalisation — the normalisation this one is built on top of.
- Monitoring and drift — spotting instability before a run dies.