Chapter 19
Scaled Dot-Product Attention
Queries, keys, and values; why divide by the square root of d; masking; and the backward pass.
Attention lets each token read the parts of a sequence that matter for its current computation. In a transformer, this is the operation that replaces a fixed context window with a learned, data-dependent lookup over previous states. This chapter derives scaled dot-product attention, its masks, and the backward pass used by NumPy tests.
19.1 A soft dictionary lookup
Think of keys and values as a dictionary. A query asks a question, each key receives a similarity score, softmax turns scores into weights, and the output is a weighted average of values. For query matrix , key matrix , and value matrix ,
Rows of sum to one, so each output row is a convex combination of value rows. Unlike a hard dictionary lookup, every key can contribute, and the weights are differentiable with respect to the query and key vectors.
The names matter. A query does not carry the information that will be returned; it describes what information is needed. A key describes when a value should be retrieved. A value is the vector that is actually mixed into the output. In the next chapter, learned projections will create all three from hidden states, but attention itself only needs the three arrays.
Keeping these roles separate makes the backward pass easier to read.
def attention_forward(Q, K, V, mask=None, return_cache=False):
"""Scaled dot-product attention: softmax(QK^T / sqrt(d_k)) V."""
scale = 1.0 / np.sqrt(Q.shape[-1])
scores = (Q @ np.swapaxes(K, -1, -2)) * scale
weights = masked_softmax(scores, mask)
output = weights @ V
if not return_cache:
return output
return output, (Q, K, V, weights, scale, mask)
The dot product is the compatibility function. If a query and a key point in similar directions, their score is large and that value receives more weight. If all scores are equal, softmax returns a uniform average.
The operation is permutation-aware only through the vectors it receives. If no positional information has been added upstream, swapping two key-value rows simply swaps their weights and leaves the weighted sum consistent with that swap. Positional encodings are therefore not decoration; they tell attention where tokens sit in the sequence.
19.2 Why divide by
Assume query and key entries are independent, mean zero, and unit variance. The unscaled score is . Its expectation is zero. Since the summands are independent,
So dot products grow in standard deviation like . Dividing by keeps score variance near one, which keeps softmax away from extreme saturation at initialization. The chapter tests verify this by Monte Carlo for several widths.
Saturation matters because softmax gradients shrink when one score dominates the row. Without scaling, increasing the key/query width makes large random scores more common, so early training can behave as if attention had already made hard choices. The scale is not a learned temperature here; it is a variance correction built into the layer.
def dot_product_variance(width, samples=50_000, seed=19):
"""Monte Carlo Var(q dot k) for independent unit-variance entries."""
rng = np.random.default_rng(seed)
Q = rng.standard_normal((samples, width))
K = rng.standard_normal((samples, width))
return float(np.var(np.sum(Q * K, axis=1)))
Masks remove illegal dictionary entries before softmax. A causal mask allows position to read only positions , which is required for autoregressive language modeling. A padding mask blocks pad tokens in a batch. In code, disallowed scores become , so their exponentials and softmax weights are zero.
Mask shape is a practical source of bugs. A self-attention causal mask has shape and can broadcast across the batch. A padding mask usually has one boolean per key position in each example, so every query in that example blocks the same padded keys. When both are needed, combine them with logical and before softmax.
def causal_mask(length):
"""True where a query position may attend to a key position."""
positions = np.arange(length)
return positions[:, None] >= positions[None, :]
def padding_mask(valid):
"""Convert a boolean (B, T) validity array to a (B, 1, T) key mask."""
return np.asarray(valid, dtype=bool)[:, None, :]
19.3 Backward pass
Let bars denote gradients of a scalar loss. From , ordinary matrix calculus gives
Softmax is row-wise. For one row , the Jacobian-vector product is . Applied to every row,
Masked positions receive zero gradient because changing a disallowed score cannot change the output. Finally, gives
def attention_backward(grad_output, cache):
"""Backward pass for scaled dot-product attention."""
Q, K, V, weights, scale, mask = cache
grad_V = np.swapaxes(weights, -1, -2) @ grad_output
grad_weights = grad_output @ np.swapaxes(V, -1, -2)
row_dot = np.sum(grad_weights * weights, axis=-1, keepdims=True)
grad_scores = weights * (grad_weights - row_dot)
if mask is not None:
grad_scores = np.where(mask, grad_scores, 0.0)
grad_Q = (grad_scores @ K) * scale
grad_K = (np.swapaxes(grad_scores, -1, -2) @ Q) * scale
return grad_Q, grad_K, grad_V
The tests check , , and against finite differences with a causal mask. That is important: the forward pass is short, but a sign error in the softmax row term silently corrupts training.
The order of these gradients mirrors the forward graph. First split the output gradient between the weights and values. Then move through the row-wise softmax, where rows do not interact. Last, move through the score matrix multiplication, which sends one contribution to queries and the transposed contribution to keys. This structure is also why optimized kernels can recompute or stream pieces of attention without changing the derivative.
19.4 Self-attention and cross-attention
In self-attention, , , and are projections of the same sequence. Decoder language models add a causal mask so a token cannot read the future. This is the attention used in the next-token objective from Chapter 18.
In cross-attention, queries come from one sequence and keys and values come from another. A decoder can query encoder states, or a text model can query image features. The equations are unchanged; only the source of differs from the source of and .
The shape difference is the clue. Self-attention usually has the same query and key length, so is square. Cross-attention can have decoder positions and source positions, so is rectangular. The value width controls the output width; the key/query width controls the scoring space.
|
In practice
|
Scaled dot-product attention is the core operation introduced in the Transformer [vaswani2017attention], building on earlier attention mechanisms for sequence-to-sequence models [bahdanau2014neural]. Decoder-only LLMs use causal self-attention in every block. Padding masks still matter for batched variable-length examples, packed training streams, and encoder-style models. Efficient exact kernels such as FlashAttention reorganize the same equations to reduce memory traffic, not to change the mathematical result [dao2022flashattention], [dao2023flashattention2]. |
19.5 Teach it
The one-sentence version. Attention is a differentiable lookup: queries score keys, softmax makes weights, and weights average values.
An analogy. A query is a search phrase, keys are document titles, and values are the document contents; attention reads a blend instead of choosing one document.
At the board.
-
Draw three key-value cards and one query.
-
Compute dot products, divide by , and softmax them.
-
Multiply the weights by values to get the output.
-
Add a causal mask and erase future cards before softmax.
Misconceptions to address.
-
"The values decide the weights." Queries and keys decide weights; values are averaged.
-
"The scale is arbitrary." It controls score variance before softmax.
-
"A mask is applied after softmax." It is applied to scores before softmax.
Check for understanding. If a padding mask blocks a key, what should its attention weight and score gradient be?
19.6 Exercises
Explain why attention can be viewed as a soft dictionary lookup. Which arrays play the roles of query, key, and value?
Derive (19.2) under the independence and unit-variance assumptions, then explain the scale.
For one row , derive .
Implement the forward and backward passes for causal scaled dot-product attention and check gradients with finite differences.
References
-
[bahdanau2014neural] D. Bahdanau, K. Cho, and Y. Bengio. Neural Machine Translation by Jointly Learning to Align and Translate. 2014. arXiv:1409.0473
-
[dao2022flashattention] T. Dao et al. FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness. 2022. arXiv:2205.14135
-
[dao2023flashattention2] T. Dao. FlashAttention-2: Faster Attention with Better Parallelism and Work Partitioning. 2023. arXiv:2307.08691
-
[vaswani2017attention] A. Vaswani et al. Attention Is All You Need. 2017. arXiv:1706.03762