Why divide the scores by the square root of the key width?
Ask a room of ML people why the scores are divided by $\sqrt{d_k}$ and most will give the answer everyone gives: to keep the variance of the scores at one. Ask the follow-up, “and why does that matter?”, and the room gets quieter. The first answer is not wrong.
It is also not an answer, and a good interviewer knows it. This is exactly the kind of nuance that slides when an agent writes the attention block and the tests come back green, which is most days now, and exactly the kind a PhD student or a serious engineer should be able to produce without pausing.
The complete answer has two halves. The first half is a two-line variance calculation. The second half is about the derivative of the softmax, and it is the half that actually explains why the network would fail to learn.
This part is about both halves, and then about the follow-up questions that separate a memorized answer from an understood one: whether you could just initialize smaller, what happens after training moves the weights, why softmax at all, and why more than one head.
The 30-second version
If the components of $q$ and $k$ are independent with mean zero and unit variance, their dot product has variance $d_k$, so the scores have standard deviation $\sqrt{d_k}$ and grow with head width. Large scores push the softmax toward a one-hot distribution. The softmax Jacobian is $\text{diag}(p) - pp^\top$, which goes to zero as $p$ becomes one-hot, so almost no gradient flows back through the scores into $W_Q$ and $W_K$. Dividing by $\sqrt{d_k}$ restores unit variance, keeps the softmax soft at initialization, and keeps its gradient alive. It is a temperature chosen so that logits are $O(1)$ regardless of $d_k$.What is attention actually computing?
Before the square root, the thing it is attached to. For one query $q$ and a set of keys $k_1, \ldots, k_T$ with values $v_1, \ldots, v_T$:
\[\text{Attn}(q) = \sum_{j=1}^{T} \alpha_j v_j, \qquad \alpha_j = \frac{\exp(q \cdot k_j / \sqrt{d_k})}{\sum_{j'} \exp(q \cdot k_{j'} / \sqrt{d_k})}\]A similarity between the query and each key, turned into a probability distribution, used to average the values. A soft dictionary lookup: the query is what you are looking for, the keys are what each entry is about, the values are what you get back, and instead of retrieving one entry you retrieve all of them in proportion to how well they match.
The 2017 paper had a choice of similarity function. The previous generation of attention, from Bahdanau’s translation work, used a small neural network: $\text{score}(q, k) = w^\top \tanh(W_1 q + W_2 k)$. Additive attention, and it works. The paper picked the dot product instead for a reason that has nothing to do with expressiveness: $QK^\top$ is one matrix multiplication, which is the one operation GPUs are unreasonably good at.
Then they noted, in a footnote that is now the most-quoted footnote in the field, that the dot product performs worse than additive attention for large $d_k$ unless you scale it, and suspected the cause was softmax saturation. That footnote is the entire answer; the rest of this post is unpacking it.
Half one: why do the scores grow with the width?
Take a query and a key whose components are independent random variables with mean zero and variance one. This is not an arbitrary assumption: after a normalization layer and a standard weight initialization, the entries of $q = W_Q x$ and $k = W_K x$ have roughly this distribution at the start of training. Their dot product is
\[q \cdot k = \sum_{i=1}^{d_k} q_i k_i\]Each term has mean $\mathbb{E}[q_i k_i] = \mathbb{E}[q_i]\,\mathbb{E}[k_i] = 0$ and variance $\mathbb{E}[q_i^2]\,\mathbb{E}[k_i^2] = 1$, using independence. The terms are independent of each other, so variances add:
\[\text{Var}(q \cdot k) = d_k, \qquad \text{std}(q \cdot k) = \sqrt{d_k}\]The scores are not $O(1)$. They are $O(\sqrt{d_k})$, and a head of width 64 produces scores about eight times larger than a head of width 1. Divide by $\sqrt{d_k}$ and the variance is back to one for every width. That is the whole of half one, and it is what most people say and stop.
Half two: what do large scores do to the softmax’s gradient?
Here is why it matters. The softmax is $p_i = e^{z_i} / \sum_j e^{z_j}$. Its Jacobian, which is what the gradient has to pass through on its way from the loss back to the scores, is
\[\frac{\partial p_i}{\partial z_j} = p_i\,(\delta_{ij} - p_j), \qquad J = \text{diag}(p) - p\,p^\top\]Look at what happens as the scores grow. Multiply all the $z$ by a large constant and one entry wins: $p$ becomes one-hot, say $p = e_1$. Then $\text{diag}(p) = \text{diag}(1, 0, \ldots, 0)$ and $pp^\top$ is the same matrix, so $J = 0$. Every entry. The softmax has become a hard argmax, and an argmax has zero derivative almost everywhere.
The gradient from the loss back to the scores is $J^\top$ times the upstream gradient, so it is zero too, and nothing reaches $W_Q$ or $W_K$. The attention pattern is frozen at whatever the random initialization produced. The model can still learn, through the values and the FFN, but the one thing attention is for, learning where to look, does not happen.
This is the failure the footnote is describing. It is not that large scores are numerically unstable, though in low precision they can be. It is that a saturated softmax is a dead gradient.
You can check both halves in ten lines. The first block reproduces the $\sqrt{d_k}$ growth; the second scales a fixed set of 16 scores up and watches the Jacobian die:
import numpy as np
rng = np.random.default_rng(0)
for dk in [8, 64, 512]:
q = rng.standard_normal((10000, dk)); k = rng.standard_normal((10000, dk))
s = (q * k).sum(-1)
print(dk, s.std().round(2), (s / np.sqrt(dk)).std().round(2))
# 8 2.84 1.00
# 64 8.09 1.01
# 512 22.46 0.99
def softmax(z):
e = np.exp(z - z.max()); return e / e.sum()
def jac_norm(z):
p = softmax(z); return np.linalg.norm(np.diag(p) - np.outer(p, p))
z = rng.standard_normal(16)
for scale in [1, 8, 32, 128]:
print(scale, softmax(scale * z).max().round(3), f"{jac_norm(scale * z):.1e}")
# 1 0.213 3.0e-01
# 8 0.941 1.0e-01
# 32 1.000 2.2e-05
# 128 1.000 2.5e-20
A scale of 32 already leaves the gradient five orders of magnitude smaller than it should be. Unscaled attention with $d_k = 1024$ sits at exactly that scale.
Could you just initialize smaller, and what stops the logits growing later?
This is where the interview gets interesting, because every one of these is a natural next question and each has a real answer.
The 30-second version
At initialization, yes: scaling the logits by a constant is a reparameterization, and shrinking the weights by $d_k^{-1/4}$ each would give identical scores. The difference shows up in training. Optimizers like Adam take steps of roughly the learning rate per parameter regardless of the parameter's scale, so small weights get relatively huge updates and the logits leave the good regime immediately. Keeping the weights at their natural scale and putting the constant in the forward pass keeps the effective learning rate sane. Maximal update parameterization takes this line of thought further and argues the right scale is actually $1/d_k$, because after training $q$ and $k$ become correlated and their dot product grows like $d_k$, not $\sqrt{d_k}$.That last point is worth holding onto. The independence assumption in half one is an initialization story. Training deliberately makes queries and keys correlated: that is what “learning to attend” means. Once they are, the dot product scales with $d_k$, which is why Tensor Programs V (Yang et al., 2022) uses $1/d_k$ scaling in its width-transferable parameterization.
In a standard parameterization the $1/\sqrt{d_k}$ convention persists because everything else has been tuned around it.
The 30-second version
Nothing, and they do. Attention logit growth is a documented cause of training instability in large models: the weights drift, $\lVert q \rVert \lVert k \rVert$ grows, the softmax saturates mid-training, and the loss spikes. The modern fix is QK-norm: apply a LayerNorm or RMSNorm to $q$ and $k$ per head before the dot product, so the logit is bounded by a learned scale times $\cos\theta$ no matter what the weights do. ViT-22B needed it to train at all, and it is standard in recent open models.The scaling is a promise about the start of training. QK-norm, from Henry et al. (2020) and made necessary at scale by Dehghani et al. (2023), turns it into a promise for the whole run.
Wortsman et al. (2023) showed that attention logit growth is one of the two reproducible instabilities you can see in small models and predict for large ones, and that QK-norm removes it. Part 6 puts this in the list of things that changed since 2017.
The 30-second version
Softmax gives a convex combination: non-negative weights that sum to one, so the output is an average of values and its scale does not depend on how many tokens are attended to. It is differentiable, it can be made arbitrarily sharp (retrieve one token) or flat (average everything) by the same logits, and its exponential form is what lets FlashAttention compute it blockwise with a running maximum. Drop the exponential for a kernel feature map and you get linear attention, which is $O(T)$ but loses the ability to sharply pick out one token. Replace it with a sigmoid and the weights no longer compete, which works with care but changes the scale of the output with sequence length.Why more than one head?
The second half of the classic attention interview is multi-head, and the usual answer (“different heads learn different things”) is true but is an observation, not a reason.
The 30-second version
One head produces one distribution per query, so its output is a single convex combination of values. If a token needs information from two places, one head has to blend them, and blending loses both. $h$ heads give $h$ independent distributions, each a lookup in a $d/h$-dimensional subspace, concatenated so the token gets $h$ separate answers rather than one average. The cost is the same as one wide head: the projections are the same size, just block-structured. The typical head width of 64 to 128 is a compromise between the rank of each head's score matrix and the number of separate places a token can look.The cleanest way to see the limitation of one head: the output for a query lives in the convex hull of the value vectors. It is an average. An average of “the subject of this sentence” and “the previous token” is neither. Two heads can return both, side by side, and let $W_O$ and the FFN decide what to do with the pair.
There is also a rank argument. A single head’s score matrix is $Q K^\top$ with $Q, K \in \mathbb{R}^{T \times d_k}$, so its rank is at most $d_k$. Bhojanapalli et al. (2020) showed that when $d_k$ is smaller than the sequence length, there are attention patterns a head simply cannot represent, which is an argument against making heads too narrow.
Meanwhile Michel et al. (2019) found that many trained heads can be removed at inference with little loss, and Voita et al. (2019) found the survivors are interpretable: positional heads, syntactic heads, rare-token heads. The design space between those two results is why head width has stayed near 64 to 128 for nearly a decade of scaling.
How does the causal mask work, and what does attention cost?
Two smaller things that come up as follow-ups.
The causal mask adds $-\infty$ to every score with $j > i$ before the softmax. Because $e^{-\infty} = 0$, those positions get exactly zero weight and drop out of the normalization too. It is done to the scores rather than to the output because the normalization must not count masked positions. It is also why the same forward pass trains on all $T$ positions at once: position $i$’s output was provably computed from tokens $\le i$.
The cost is in Part 1’s accounting: $O(T^2 d)$ per layer for scores and weighted sums, against $O(T d^2)$ for the projections. But the thing that actually hurts is not the FLOPs, it is that a naive implementation writes a $T \times T$ matrix per head to memory and reads it back for the softmax.
At $T = 8192$ and 32 heads in bf16 that is four gigabytes per layer per sequence. FlashAttention computes exactly the same softmax without ever materializing the matrix, using a running maximum and running sum so that blocks of keys can be processed one at a time. Same math, a different order of operations, and the reason long contexts are possible. Part 6 has the details.
Rapid fire: can you do these from memory?
- Derive the variance of $q \cdot k$ under independent unit-variance components. State the assumption you needed.
- Write the softmax Jacobian and show it vanishes for a one-hot distribution.
- Why is unscaled attention "unlearnable" rather than merely "unstable"?
- Why put the scale in the forward pass instead of shrinking the initialization?
- Why does μP prefer $1/d_k$? What assumption from question 1 breaks after training?
- What is QK-norm and what failure does it prevent?
- Why is a single head's output an average, and why is that a limitation?
- Where in the computation is the causal mask applied, and why there?
Part 3 turns to the two components that let you stack a hundred of these blocks: the residual connection and the normalization, and why the order you put them in matters more than it looks.