Why architecture matters here
The instability is structural, not accidental, and the argument is worth following carefully. Consider a single attention head with head dimension d. At initialization, q and k entries are roughly zero-mean and unit variance, so their dot product has variance d and standard deviation √d. Dividing by √d returns the logit to unit scale — this is the entire justification for the scaled-dot-product design, and it is sound. But it holds only at initialization. Training updates W_q and W_k, and nothing in the loss penalizes their growth: weight decay is usually excluded from or weak on these matrices, and there is a genuine incentive to grow them, because sharper attention lowers loss in the short run. The √d divisor is a constant. The thing it corrects for is not.
The scaling is multiplicative, which is why the drift becomes a cliff. Logit magnitude scales roughly with ‖W_q x‖ · ‖W_k x‖, so a 3× growth in each projection's norm yields a 9× growth in logits. And the process is self-reinforcing: larger logits produce sharper attention, sharper attention lowers loss on the current batch, gradient descent rewards it, and the weights grow further — a positive feedback loop with no restoring force anywhere in the architecture. What looks like a sudden spike at step 40,000 is the visible end of a slow exponential running since step one.
Softmax then converts drift into catastrophe, sharply. Softmax entropy falls smoothly as logits grow — attention gets progressively more peaked, which is fine and often desirable — until the max logit reaches roughly 30-50, at which point in bfloat16 the largest entry rounds to 1.0 and the rest to 0. The transition from 'sharp' to 'one-hot' is abrupt in a way the underlying continuous quantity is not. Once one-hot, the softmax Jacobian is approximately zero, so no gradient flows back through that head to W_q or W_k. The head is frozen. Worse, it is frozen in a configuration chosen by whatever the weights happened to be at saturation, and if the loss spikes and the optimizer responds with a large update, the head can be knocked into a wildly wrong one-hot pattern it can never learn its way out of.
This is why the problem scales with model size and shows up in exactly the runs where it hurts most. Bigger models have more heads and more layers, so some head saturating becomes near-certain — you are sampling the tail of a distribution many more times. Longer training gives the exponential more time to run. Larger learning rates accelerate weight growth. And low-precision training lowers the logit threshold at which saturation bites. Every axis along which the field has scaled over the past five years makes this failure more likely, which is precisely why QK-Norm went from a curiosity in ViT-22B to standard equipment across Gemma, OLMo, Chameleon, and most serious training stacks — the small models where you could ignore it are not the models anyone is training.
The architecture: every piece explained
Projection, then normalization, then RoPE (top row). The order in the diagram is the whole design and each step is load-bearing. The hidden state x arrives already normalized by the block's pre-norm — which is worth noting, because it shows that pre-norm does not solve this problem: it bounds the input to the projection, not the projection's own weights, and the weights are what grow. W_q and W_k project to q and k. Then QK-Norm applies RMSNorm to each: q̂ = g_q · q / rms(q). Now ‖q̂‖ is fixed at roughly g_q·√d no matter what W_q did during training. The logit q̂·k̂/√d is bounded by construction. The unbounded quantity has been made bounded, and the fix is at the source rather than a clamp downstream.
Why RoPE comes after (top-right). This ordering is not stylistic. RoPE applies a rotation to q and k, and rotations preserve norm — so applying RoPE after normalization leaves the norm exactly as QK-Norm set it, and the bound survives. Reversing the order still technically works for the norm (rotation is norm-preserving in both directions) but breaks something subtler: RMSNorm's learnable per-dimension gain would then be applied to rotated coordinates, where the dimension index no longer corresponds to a fixed frequency pair. The gain vector would be learning a position-dependent thing, which is not what it is for. Normalize in the projection's coordinate frame, then rotate — every production implementation does this and it is worth knowing why rather than copying it.
Per-head versus per-tensor (the RMSNorm box). The normalization can be applied over the full projected tensor or independently per head, and the difference matters more than it looks. Per-head is the standard choice: each head gets its own norm and its own learnable gain, so each can select its own effective attention temperature — one head can be sharp and near-one-hot for induction-style copying while another stays diffuse for averaging. Per-tensor normalization couples all heads to one scale and removes that freedom, which measurably costs quality. The per-head gain is what preserves expressiveness: the model is not forced into a fixed temperature, it is merely forced to choose one explicitly rather than obtain it by growing weights without bound.
What it costs and what it interacts with (lower row). The cost is two RMSNorms per attention layer. These are memory-bandwidth-bound elementwise operations on tensors already in registers or L2, so measured overhead lands around 1-2% of step time — and it is frequently repaid, because stable runs tolerate higher learning rates and need less warmup. Two interactions are worth knowing. First, attention sinks: models commonly learn to dump attention on the first token as a no-op, which requires a large logit for that position; QK-Norm bounds logits and therefore makes the sink harder to express, which is one reason QK-Norm pairs naturally with an explicit learnable sink or an off-by-one softmax. Second, quantization: bounded activations quantize dramatically better than unbounded ones, so QK-Norm is a quiet gift to anyone who later has to serve the model in int8 — the outliers that make attention activations hard to quantize are exactly the logit blowups QK-Norm prevents.
End-to-end flow
The forward pass. A hidden state x of shape [batch, seq, d_model] enters the attention block and passes through the block's pre-norm. It is projected: q = x W_q, k = x W_k, v = x W_v, each reshaped to [batch, heads, seq, d_head]. Now QK-Norm fires on q and k — but not on v, which is the detail most first implementations get wrong. V does not participate in the dot product; it is averaged with the attention weights, and normalizing it would change the output distribution for no stability benefit. Only the two tensors that multiply together to form a logit need bounding.
Normalize, rotate, score. For each head independently, RMSNorm computes rms(q) = √(mean(q²) + ε) over the head dimension and returns q̂ = g_q · q / rms(q), with g_q a learnable per-dimension gain initialized to 1. Same for k. RoPE then rotates q̂ and k̂ by position — norm-preserving, so the bound holds. The logits are q̂ k̂ᵀ / √d. Where an unnormalized model at step 40,000 might produce a max logit of 80, this produces something in the range of 5-15, tuned by the learned gains. Softmax over that is peaked but not saturated: the top token might take 0.7 of the mass, leaving 0.3 distributed — and crucially, leaving a nonzero Jacobian.
The backward pass. This is where the benefit is actually collected. Gradient flows back through the softmax, and because the softmax is not saturated, the Jacobian is well-conditioned and real gradient reaches q̂ and k̂. It then flows through RMSNorm, which has a useful property: its gradient is orthogonal to the input direction, meaning it passes back information about direction while discarding information about magnitude. That is precisely the desired behavior — the model can learn where to attend without being able to learn to attend harder by growing weights. The feedback loop from section one is severed: growing W_q no longer increases logits, so there is no longer a gradient incentive to grow it, and the exponential never starts.
Over a full run. The observable difference is in the metrics rather than in any single step. Max attention logit in an unnormalized run climbs steadily — 5 at step 1k, 15 at 10k, 40 at 30k — and the loss spike arrives shortly after saturation. With QK-Norm the max logit is flat for the entire run: 8 at step 1k, 8 at step 100k, because it is bounded by g·√d and g moves slowly under weight decay. Softmax entropy declines gently and stabilizes instead of collapsing to zero. The loss curve has no spikes. The practical dividend is that the whole apparatus of spike defense — the LR reductions, the batch skipping, the checkpoint rollbacks, the babysitting — becomes unnecessary, and you can often raise the learning rate on top, because the instability that forced you to be conservative is gone.