The whole operation, with shapes attached

For one attention head over a sequence of N tokens, with head dimension d_k for queries and keys and d_v for values (almost always d_k = d_v), the operation is:

Q: [N, d_k]   K: [N, d_k]   V: [N, d_v]
S = Q K^T          → [N, N]     raw scores
S' = S / √d_k       → [N, N]     scaled
S'' = S' + M       → [N, N]     M is the additive mask
A = softmax(S'', axis=-1)  → [N, N]  rows sum to 1
O = A V            → [N, d_v]   output

Row i of S holds the affinity of query i against every key; row i of A is a probability distribution over the N positions; row i of O is that distribution’s weighted average of the value vectors. In a real implementation everything carries two leading axes, [B, h, N, d_k], and the two matmuls are batched over B·h independent problems. The softmax is strictly row-wise — keys never compete across queries — and that is why the whole thing parallelises so cleanly.

Advertisement

Where the √d_k comes from

Treat the components of a query and a key as independent, zero-mean, unit-variance draws — roughly true at initialisation. A score is a sum of d_k such products, and independent variances add:

s = Σ_r q_r k_r ,  r = 1..d_k
E[q_r k_r] = 0 ,  Var(q_r k_r) = 1
Var(s) = d_k    →   std(s) = √d_k

So the spread of the logits grows with head width, and the softmax sees a different input scale at d_k = 32 than at d_k = 128. Dividing by √d_k makes the score distribution width-invariant: one architectural knob stops leaking into another. Note why the exponent is one half and not one. Dividing by d_k would shrink logit spread as 1/√d_k, driving every attention distribution toward uniform as the model widens — the opposite failure. The exact factor that cancels what the sum introduced is √d_k, nothing else.

Advertisement

What saturation does to the gradient

The cost of unscaled logits shows up in the backward pass. Softmax has Jacobian ∂p/∂s = diag(p) − p p^T, so an upstream gradient g on the weights becomes

∂L/∂s_j = p_j (g_j − Σ_i p_i g_i)

Every term carries a factor of p_j. If one weight is 0.99999 and the rest share 10-5, the small entries get gradients scaled by 10-5 and the large one gets g_j minus an average it itself dominates — also near zero. The layer stops learning which key to attend to; it can only reinforce the key it already picked.

How close is that to reality? At d_k = 128 the raw logit spread is √128 ≈ 11.3, so a one-standard-deviation gap between the top two scores means a probability ratio of e^11.3 ≈ 8.2×10^4 — effectively one-hot from the first step. Scaled, the same gap becomes 1.0 and the ratio is e ≈ 2.72: a preference, not a verdict.

A worked example on four dimensions

Take d_k = 4, one query, three keys:

q  = [ 1, -1,  2, 0.5]
k_1 = [ 2,  0,  1, -1]   k_2 = [0, 1, -1, 2]   k_3 = [1, 1, 1, 1]

raw    s = [ 3.5, -2.0,  2.5]
scaled s/√4 = [1.75, -1.0, 1.25]

softmax(raw)    = [0.729, 0.003, 0.268]
softmax(scaled) = [0.599, 0.038, 0.363]

Both rows rank the keys identically — scaling is monotone, it never changes which key wins. What it changes is how much probability mass the losers keep, and therefore how much gradient they receive. The raw row gives key 2 a weight of 0.003; the scaled row gives it 0.038, twelve times more signal to work with. At d_k = 4 the effect is a nudge. Redo the same arithmetic with d_k = 128 and typical logits near ±11 and the unscaled row is numerically one-hot, with the losers' gradients rounded to zero rather than merely small.

Masking is arithmetic on scores, never on weights

Causal and padding masks are applied additively to the scaled scores, before the softmax — a lower-triangular M with 0 on allowed positions and a large negative constant elsewhere. Zeroing entries of A after the softmax is the classic wrong fix: the rows no longer sum to 1, so the output is a shrunken, unnormalised average whose magnitude depends on how many positions were masked.

Two traps follow from the choice of constant. In fp16 the largest finite value is 65504, so the popular -1e9 becomes -inf on cast — use finfo(dtype).min or a value like -1e4 that survives. And a row that is entirely masked (a fully padded query position) yields -inf − (-inf) = NaN under max-subtraction, which then poisons every downstream tensor. Guard those rows explicitly rather than hunting the NaN later.

Numerical stability: subtract the row max

Every correct softmax computes the shifted form, using the identity that a constant shift cancels between numerator and denominator:

m_i = max_j S''_ij
p_ij = exp(S''_ij − m_i) / Σ_j exp(S''_ij − m_i)

Because every exponent is now ≤ 0, exp can never overflow — it can only underflow to 0, which is harmless since the largest term is exactly 1 and the denominator is therefore at least 1. Without the shift the ceiling is real: exp overflows above ln(3.4×10^38) ≈ 88.7 in fp32 and above ln(65504) ≈ 11.1 in fp16 — a threshold raw d_k = 128 logits reach routinely. Max-subtraction, not scaling, is what prevents overflow; scaling is about gradients. This same two-pass structure is what FlashAttention makes single-pass with a running max and a rescaled running sum.