Skip to content

Scaled dot-product attention (forward and backward)

ModuleL5.1 · build · Python · Pass 5 · 3 to 4 h, plus your graded tests (rung R5)
You buildpython/tinyllm/xfmr/sdpa.py: sdpa_forward, sdpa_backward (by hand), and scaled_dot_product_attention, one autograd node; and your own oracle tests in python/tests/l5-1-sdpa/
Contractcourse/contracts/py/tinyllm/xfmr/sdpa.pyi
Testscourse/tests/L5.1/test_sdpa.py (what they check: section 4), golden values from torch 2.14.1 in course/fixtures/L5.1/sdpa_torch.npz (course/oracle/L5.1/sdpa_torch.py); your tests are graded by mutation, threshold 0.80 with every pitfall fault required
NeedsL0.1 Tensor and from_op · M09.2 stable softmax · reading: L4.3 Luong’s dot score, M08.3 and S-M08 (the softmax VJP), L0.2 (the dropout draw rule) (or --ref-deps)
Used byL5.3 multi-head attention calls it once per layer · later L7.5 and L7.6 (grouped-query and latent attention), L8.2 (attention over a KV cache), L9.3 (the C kernel is proven against sdpa_forward)
MilestoneMS-L5
Optional depthVaswani et al., “Attention Is All You Need” (2017), section 3.2.1; Dao et al., “FlashAttention” (2022), section 3.1 and appendix B.2 (the backward with D=rowsum⁡(dO∘O)D = \operatorname{rowsum}(dO \circ O))
  • Attention is a soft dictionary lookup: O=softmax⁡(QK⊤/d+mask) VO = \operatorname{softmax}(QK^\top/\sqrt{d} + \text{mask})\,V, a weighted average of values with weights from query-key similarity (test_hand_example_forward).
  • The 1/d1/\sqrt{d} keeps the scores’ variance at 1 whatever the width, so the softmax neither saturates nor flattens (test_score_variance_is_one, test_default_scale_is_inverse_sqrt_d).
  • The backward is four matrix products and the softmax VJP dS=P∘(dP−rowsum⁡(dP∘P))dS = P \circ (dP - \operatorname{rowsum}(dP \circ P)), and that row sum equals dO⋅OdO \cdot O (test_hand_example_backward, test_gradcheck_float64).
  • Attention is equivariant to permuting keys with their values: it has no notion of order (test_key_permutation_equivariance).
  • Blocked keys get weight and gradient exactly 0, and a fully blocked row is zeros, never NaN (test_masked_keys_get_no_weight_or_gradient, test_fully_masked_row_gives_zeros).
Terminal window
ol start L5.1 # stubs sdpa.py; prints your test path and rung (R5)
ol tests L5.1 # the course tests
# write your oracle tests in python/tests/l5-1-sdpa/, then:
ol check L5.1 # course tests and the mutation grade of your tests
ol diff L5.1 # after passing: your code against the reference

The sequence models of Part 4 read the source through an RNN: one step per token, and everything the decoder knows about token 3 has passed through every later state. Luong attention (L4.3) already reads the encoder states directly with a dot product. The transformer keeps only that read: every position attends to every position in one matrix product, so a sequence of length TT costs T2T^2 dot products but no sequential steps, and the gradient from position 50 to position 3 is one hop instead of 47. Every model you build from here on (the 2017 Transformer, GPT, BERT, the modern Llama block, the inference engine’s kernels) is built on this one function. Writing its backward by hand, rather than composing ops, makes it one node of your autograd graph instead of seven, and it is the exact computation the C kernel of L9.3 and FlashAttention reorganize.

SymbolMeaningType / shape
QQqueries, one row per position asking[..., Tq, d]
KKkeys, one row per position offering[..., Tk, d]
VVvalues, what each key position contributes[..., Tk, dv]
ddquery and key widthint
α\alpha (scale)score scale, default 1/d1/\sqrt{d}float
S=αQK⊤S = \alpha QK^\topscores[..., Tq, Tk]
MMmask, True = may attend (L5.2)bool, broadcast to SS
P=softmax⁡(S)P = \operatorname{softmax}(S)weights, each row sums to 1 (over keys)[..., Tq, Tk]
O=PVO = PVoutput[..., Tq, dv]
dO,dP,dS,dQ,dK,dVdO, dP, dS, dQ, dK, dVgradients of a scalar loss with respect to eachsame shapes as their arrays
Di=∑jPij dPijD_i = \sum_j P_{ij}\, dP_{ij}the row sum in the softmax VJP[..., Tq, 1]

A Python dict returns the value whose key equals the query. Attention returns a blend: score query qiq_i against every key kjk_j by the dot product (large when they point the same way), turn the scores into positive weights that sum to 1 with a softmax over the keys, and average the values with those weights: oi=∑jPijvjo_i = \sum_j P_{ij} v_j. Because the output is a weighted sum over a set of key-value pairs, reordering the pairs changes nothing: attention has no idea of position, which is why the transformer adds positional encodings (L5.4). Each query is computed independently of the others.

Why scale? If the entries of qq and kk are independent with mean 0 and variance 1, then q⋅k=∑t=1dqtktq \cdot k = \sum_{t=1}^{d} q_t k_t is a sum of dd terms of variance 1: variance dd. At d=64d = 64 the scores have standard deviation 8, the softmax is nearly one-hot, and its gradient (which is P(1−P)P(1-P)-shaped) nearly vanishes. Multiplying by 1/d1/\sqrt{d} brings the variance back to 1 for every width. test_score_variance_is_one measures exactly this.

A mask blocks query-key pairs (L5.2): set the blocked scores to −∞-\infty before the softmax. Then e−∞=0e^{-\infty} = 0, the blocked weights are exactly 0, and the open weights still sum to 1. Masking after the softmax (multiplying weights by 0/1) leaves rows that no longer sum to 1. A row with every key blocked is all −∞-\infty; the stable softmax of M09.2 defines its weights as zeros, so its output and gradients are zero instead of NaN. The softmax also subtracts each row’s maximum first, so scores of 10410^4 never overflow.

Given dOdO, the gradient of the loss with respect to the output, work backward through O=PVO = PV, then P=softmax⁡(S)P = \operatorname{softmax}(S), then S=αQK⊤S = \alpha QK^\top. For a matrix product C=ABC = AB, the vector-Jacobian products are dA=dC B⊤dA = dC\,B^\top and dB=A⊤dCdB = A^\top dC (M08.3). So

dV=P⊤dO,dP=dO V⊤.dV = P^\top dO, \qquad dP = dO\,V^\top .

For one row p=softmax⁡(s)p = \operatorname{softmax}(s), ∂pj/∂sk=pj(δjk−pk)\partial p_j / \partial s_k = p_j(\delta_{jk} - p_k), so dsk=∑jdpj pj(δjk−pk)=pk(dpk−∑jpjdpj)ds_k = \sum_j dp_j\, p_j(\delta_{jk} - p_k) = p_k(dp_k - \sum_j p_j dp_j) (S-M08):

dS=P∘(dP−D),Di=∑jPij dPij.dS = P \circ (dP - D), \qquad D_i = \sum_j P_{ij}\, dP_{ij} .

Finally S=αQK⊤S = \alpha QK^\top gives dQ=α dS KdQ = \alpha\, dS\,K and dK=α dS⊤QdK = \alpha\, dS^\top Q. Note the transpose in dKdK: dSdS is [Tq, Tk] and dKdK must be [Tk, d]. A useful identity: Di=∑jPij(dOi⋅vj)=dOi⋅oiD_i = \sum_j P_{ij}(dO_i \cdot v_j) = dO_i \cdot o_i, so the row sum can be computed from the output without storing dPdP. FlashAttention’s backward uses exactly this. Masked positions have P=0P = 0, so every gradient through them is 0: padding never learns.

Every leading dimension is an independent problem: [B, H, T, d] is B⋅HB \cdot H separate attentions, and numpy’s @ broadcasts over them. L5.3 reshapes [B, T, d_model] into heads and calls this function once.

During training, the 2017 Transformer drops attention weights: draw one uniform uu per weight (in C order, from the PCG32 the caller passes, the same rule as L0.2’s dropout), keep where u≥pu \ge p, and scale the kept ones by 1/(1−p)1/(1-p) so the expected weight is unchanged: O=(P∘m)VO = (P \circ m)V with m=keep/(1−p)m = \text{keep}/(1-p). The backward must use the same mm: dV=(P∘m)⊤dOdV = (P \circ m)^\top dO and dP=(dO V⊤)∘mdP = (dO\,V^\top) \circ m, and the softmax VJP is unchanged. The function returns PP before dropout, as a constant: the weights are for inspection (heat maps, the KV cache), and gradients flow through OO only.

One query, three keys, d=4d = 4 so α=1/2\alpha = 1/2. q=(2,0,0,0)q = (2, 0, 0, 0); keys (1,0,0,0)(1, 0, 0, 0), (0,1,0,0)(0, 1, 0, 0), (2,0,0,0)(2, 0, 0, 0); values (1,0)(1, 0), (0,1)(0, 1), (1,1)(1, 1).

StepComputationValue
SS12(2,0,4)\frac12(2, 0, 4)(1,0,2)(1, 0, 2)
PP(e,1,e2)/(1+e+e2)(e, 1, e^2)/(1 + e + e^2), 1+e+e2=11.1071 + e + e^2 = 11.107(0.244728,0.090031,0.665241)(0.244728, 0.090031, 0.665241)
OO0.244728(1,0)+0.090031(0,1)+0.665241(1,1)0.244728(1, 0) + 0.090031(0, 1) + 0.665241(1, 1)(0.909969,0.755272)(0.909969, 0.755272)

Backward with dO=(1,0)dO = (1, 0):

StepComputationValue
dVdVP⊤dOP^\top dO: row jj is (Pj,0)(P_j, 0)(0.244728,0)(0.244728, 0), (0.090031,0)(0.090031, 0), (0.665241,0)(0.665241, 0)
dPdPdO⋅vjdO \cdot v_j(1,0,1)(1, 0, 1)
DD0.244728+0.6652410.244728 + 0.665241, which is dO⋅OdO \cdot O0.9099690.909969
dSdSP∘(dP−D)P \circ (dP - D)(0.022033,−0.081925,0.059892)(0.022033, -0.081925, 0.059892)
dQdQ12∑jdSjkj=12(0.022033+2⋅0.059892,−0.081925,0,0)\frac12 \sum_j dS_j k_j = \frac12(0.022033 + 2 \cdot 0.059892, -0.081925, 0, 0)(0.070909,−0.040963,0,0)(0.070909, -0.040963, 0, 0)
dKdK12dSj q\frac12 dS_j\, q: row jj is (dSj,0,0,0)(dS_j, 0, 0, 0)(0.022033,0,0,0)(0.022033, 0, 0, 0), (−0.081925,0,0,0)(-0.081925, 0, 0, 0), (0.059892,0,0,0)(0.059892, 0, 0, 0)

The third key scored highest and gets most of the weight; raising its score further (moving qq toward it) raises O1O_1, which is why dS3>0dS_3 > 0. These numbers are the first two tests.

def sdpa_forward(q, k, v, mask=None, scale=None) -> tuple[NDArray, NDArray]: ... # (O, P)
def sdpa_backward(q, k, v, p, dout, scale=None, dropout_mult=None) -> tuple[NDArray, NDArray, NDArray]: ...
def scaled_dot_product_attention(q: Tensor, k: Tensor, v: Tensor, mask=None, dropout_p=0.0,
scale=None, rng=None) -> tuple[Tensor, Tensor]: ... # (out, weights)

scaled_dot_product_attention computes the forward with sdpa_forward, applies dropout if asked, and returns from_op(out, [q, k, v], vjp) whose vjp is sdpa_backward: one node of the graph. Masks are bool (True = may attend) and must broadcast to the scores without changing their shape. float32 stays float32.

TestKINDChecksWhy it matters downstream
test_hand_example_forwardunit, smokesection 3’s PP and OOyou and the tests agree on the definition
test_hand_example_backwardunit, smokesection 3’s dVdV, dSdS, dQdQ, dKdKthe backward by hand
test_golden_torchgoldentorch’s output, weights, and float64 gradients on five cases (cross attention, causal, padding, custom scale, 3-D)L5.3’s torch comparison builds on it
test_golden_torch_float32goldenfloat32 in, float32 out, within the reduction boundmodels train in float32
test_gradcheck_float64gradcheckevery gradient element against central differences, with masksthe hand backward is the derivative
test_gradcheck_through_dropoutgradcheckthe same with a replayed dropout masktraining-mode gradients
test_key_permutation_equivarianceproperty, smokepermuting keys and values permutes only the weights’ columnswhy positions are added in L5.4
test_queries_are_independentpropertypermuted or single queries give the same rowsL8.2 decodes one query at a time
test_masked_keys_get_no_weight_or_gradientboundaryblocked weights and gradients exactly 0, rows sum to 1padding never learns
test_fully_masked_row_gives_zerosboundaryzeros, never NaNempty sequences in a batch
test_large_scores_do_not_overflowboundaryscores of 10410^4 give a clean one-hotsharp attention in trained models
test_default_scale_is_inverse_sqrt_dunitdefault equals 1/d1/\sqrt{d} explicitlyhead width changes do not change the math
test_score_variance_is_onestatisticalscaled score differences have variance 2 at d=4d = 4 and 64section 2.1’s argument
test_dropout_draws_and_scalingunitone uniform per weight, keep u≥pu \ge p, scale 1/(1−p)1/(1-p); p=0p = 0 draws nothingthe same stream as L0.2’s dropout
test_weights_are_a_constantunitout requires grad, weights does notno gradient through inspection
test_batched_equals_per_itemproperty[B, H, T, d] equals B⋅HB \cdot H 2-D calls, gradients tooheads are independent
test_rejects_bad_argumentsboundarymismatched shapes, float masks, bad scale, bad dropout rate, dropout without rngwiring bugs fail loudly

The oracle is the formula written out in numpy (scores, masked softmax, weighted sum, in float64); check the section 3 numbers, every mask kind, and a custom scale against it. For the backward, write your own central differences of sum(oracle(q, k, v) * g) and compare them with the gradients of the Tensor op and of sdpa_backward, with and without dropout (replay the same uniforms for the oracle’s mask). Import only contract modules (tinyllm.xfmr.sdpa, tinyllm.autograd.tensor).

PitfallSymptomCaught by
1. scaling by 1/d1/d instead of 1/d1/\sqrt{d}flat attention that sharpens with width; torch disagreestest_hand_example_forward, test_score_variance_is_one (mutant s01)
2. the softmax over the query axiscolumns sum to 1 instead of rowstest_hand_example_forward (mutant s02)
3. masking after the softmaxrows sum to less than 1test_masked_keys_get_no_weight_or_gradient (mutant s03)
4. the softmax VJP without the −D-D termgradients of a sum-to-one output that do not sum to 0test_hand_example_backward, test_gradcheck_float64 (mutant s04)
5. dK=dS QdK = dS\,Q without the transposea shape error for Tq≠TkT_q \ne T_k, wrong values otherwisetest_batched_equals_per_item (mutant s05)
6. forgetting the scale in dQdQgradients d\sqrt{d} times too largetest_gradcheck_float64 (mutant s06)
7. a backward that ignores the dropout masktraining gradients of a different functiontest_gradcheck_through_dropout (mutants s07, s08)
8. keeping u<pu < pdrops the wrong 1 - p of the weightstest_dropout_draws_and_scaling (mutant s09)
9. masking with −109-10^9a fully masked row attends uniformly to paddingtest_fully_masked_row_gives_zeros (mutant s10)
DirectionModuleHow it uses this
BackL0.1from_op(out, [q, k, v], vjp): the whole attention is one node
BackM09.2softmax with the max subtracted and an all −∞-\infty row mapped to zeros
BackL4.3the unscaled dot score this module scales (reading)
BackM08.3the matrix-product and softmax VJPs of section 2.3 (reading)
ForwardL5.2the masks passed as mask
ForwardL5.3multi-head attention: [B, H, T, d_head] in one call
ForwardL7.5grouped-query attention shares keys and values across heads
ForwardL7.6latent attention reconstructs keys and values from a small cache
ForwardL8.2queries of one step against the cached keys, with q_offset masks
ForwardL9.3the C kernel (online softmax, tiled) is tested against sdpa_forward
Your pieceProduction equivalentWhat it addsWhere to look
sdpa_forwardtorch.nn.functional.scaled_dot_product_attentiondispatch to fused kernels (Flash, memory-efficient, math) by shape and deviceaten/src/ATen/native/transformers/attention.cpp
sdpa_backwardFlashAttention-2 backwardrecomputes PP tile by tile from saved row statistics instead of storing it; D=rowsum⁡(dO∘O)D = \operatorname{rowsum}(dO \circ O)flash_attn/flash_attn_triton.py, Dao (2023)
the softmax over all keysthe online softmaxone pass over the keys with a running max and sum, the basis of L9.3Milakov and Gimelshein, “Online normalizer calculation for softmax” (2018)
scaled_dot_product_attentionvLLM PagedAttentionkeys and values in fixed-size blocks of a shared pool (L8.3)vllm/attention/