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 Q∈RTq×dk\mQ \in \R^{T_q \times d_k}, key matrix K∈RTk×dk\mK \in \R^{T_k \times d_k}, and value matrix V∈RTk×dv\mV \in \R^{T_k \times d_v},

S=QK⊤dk,A=softmax⁡(S),O=AV.(19.1)\mS=\frac{\mQ\mK^\T}{\sqrt{d_k}},\qquad \mA=\softmax(\mS),\qquad \mO=\mA\mV .\tag{19.1}

Rows of A\mA 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.

Listing 19.1 Scaled dot-product attention
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 dk\sqrt{d_k}

Assume query and key entries are independent, mean zero, and unit variance. The unscaled score is s=∑ℓ=1dkqℓkℓs=\sum_{\ell=1}^{d_k}q_\ell k_\ell. Its expectation is zero. Since the summands are independent,

Var⁡(s)=∑ℓ=1dkVar⁡(qℓkℓ)=∑ℓ=1dkE[qℓ2]E[kℓ2]=dk.(19.2)\Var(s)=\sum_{\ell=1}^{d_k}\Var(q_\ell k_\ell) =\sum_{\ell=1}^{d_k}\E[q_\ell^2]\E[k_\ell^2]=d_k .\tag{19.2}

So dot products grow in standard deviation like dk\sqrt{d_k}. Dividing by dk\sqrt{d_k} 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.

Listing 19.2 Numerically checking dot-product variance
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 tt to read only positions ≤t\le t, which is required for autoregressive language modeling. A padding mask blocks pad tokens in a batch. In code, disallowed scores become −∞-\infty, so their exponentials and softmax weights are zero.

Mask shape is a practical source of bugs. A self-attention causal mask has shape T×TT \times T 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.

Listing 19.3 Causal and padding masks
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 O=AV\mO=\mA\mV, ordinary matrix calculus gives

Vˉ=A⊤Oˉ,Aˉ=OˉV⊤.(19.3)\bar{\mV}=\mA^\T\bar{\mO},\qquad \bar{\mA}=\bar{\mO}\mV^\T .\tag{19.3}

Softmax is row-wise. For one row a=softmax⁡(s)\va=\softmax(\vs), the Jacobian-vector product is sˉ=a⊙(aˉ−(aˉ⊙a)1)\bar{\vs}=\va\odot(\bar{\va}-(\bar{\va}\odot\va)\one). Applied to every row,

Sˉ=A⊙(Aˉ−rowsum⁡(Aˉ⊙A)).(19.4)\bar{\mS}=\mA\odot\left(\bar{\mA} -\operatorname{rowsum}(\bar{\mA}\odot\mA)\right).\tag{19.4}

Masked positions receive zero gradient because changing a disallowed score cannot change the output. Finally, S=QK⊤/dk\mS=\mQ\mK^\T/\sqrt{d_k} gives

Qˉ=SˉKdk,Kˉ=Sˉ⊤Qdk.(19.5)\bar{\mQ}=\frac{\bar{\mS}\mK}{\sqrt{d_k}},\qquad \bar{\mK}=\frac{\bar{\mS}^\T\mQ}{\sqrt{d_k}} .\tag{19.5}
Listing 19.4 Attention backward pass
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 Qˉ\bar{\mQ}, Kˉ\bar{\mK}, and Vˉ\bar{\mV} 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, Q\mQ, K\mK, and V\mV 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 Q\mQ differs from the source of K\mK and V\mV.

The shape difference is the clue. Self-attention usually has the same query and key length, so A\mA is square. Cross-attention can have TqT_q decoder positions and TkT_k source positions, so A\mA is rectangular. The value width dvd_v controls the output width; the key/query width dkd_k 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].

Key equations
S=QK⊤/dk,A=softmax⁡(S),O=AV\mS=\mQ\mK^\T/\sqrt{d_k},\qquad \mA=\softmax(\mS),\qquad \mO=\mA\mV
Var⁡(∑ℓ=1dkqℓkℓ)=dk\Var\left(\sum_{\ell=1}^{d_k}q_\ell k_\ell\right)=d_k
Vˉ=A⊤Oˉ,Aˉ=OˉV⊤\bar{\mV}=\mA^\T\bar{\mO},\qquad \bar{\mA}=\bar{\mO}\mV^\T
Sˉ=A⊙(Aˉ−rowsum⁡(Aˉ⊙A))\bar{\mS}=\mA\odot\left(\bar{\mA} -\operatorname{rowsum}(\bar{\mA}\odot\mA)\right)
Qˉ=SˉK/dk,Kˉ=Sˉ⊤Q/dk\bar{\mQ}=\bar{\mS}\mK/\sqrt{d_k},\qquad \bar{\mK}=\bar{\mS}^\T\mQ/\sqrt{d_k}

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.

  1. Draw three key-value cards and one query.

  2. Compute dot products, divide by dk\sqrt{d_k}, and softmax them.

  3. Multiply the weights by values to get the output.

  4. 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

Exercise 19.1 ★ Soft lookup

Explain why attention can be viewed as a soft dictionary lookup. Which arrays play the roles of query, key, and value?

Exercise 19.2 ★★ Dot-product variance

Derive (19.2) under the independence and unit-variance assumptions, then explain the dk\sqrt{d_k} scale.

Exercise 19.3 ★★ Softmax row backward

For one row a=softmax⁡(s)\va=\softmax(\vs), derive sˉ=a⊙(aˉ−(aˉ⊙a)1)\bar{\vs}=\va\odot(\bar{\va}-(\bar{\va}\odot\va)\one).

Exercise 19.4 ★★★ Gradient-check masked attention

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