Practice / Transformer pieces

Attention backward

Ten problems on the backward pass of scaled dot-product attention: the gradients to V, A, the scores, Q and K, the projections in self-attention, the causal mask, why the scores are divided by √d, and the multi-head output projection, with worked solutions and the mistakes that transpose the wrong matrix.

Before you start

Scaled dot-product attention is three matrix products and a softmax: scores from queries and keys, weights from the scores, outputs from the weights and the values. Its backward pass is the same pieces run in reverse, and every step is either a matrix product, whose gradient is another matrix product with one factor transposed, or the softmax page's row result applied to each row. These ten problems derive each gradient, push them through the projections of self-attention, add the causal mask and the 1/d1/\sqrt d scaling, and finish with the multi-head output projection. The five mistakes are the ones that give a plausible shape or a plausible number: a missing transpose, a softmax Jacobian that mixes rows, a dropped 1/d1/\sqrt d, a path through XX left out, and a mask applied after the softmax.

  • The conventions are those of the previous pages: vectors are columns, Jacobians are in numerator layout, a gradient has the shape of the variable it is taken with respect to, and ⊙\odot is the elementwise product.
  • One head. Q∈Rn×dQ \in \mathbb{R}^{n\times d}, K∈Rm×dK \in \mathbb{R}^{m\times d} and V∈Rm×dvV \in \mathbb{R}^{m\times d_v}, with rows as positions: nn query positions, mm key positions, width dd for queries and keys and dvd_v for values.
  • S=QK⊤/dS = QK^\top/\sqrt d is n×mn\times m; A=softmax⁡(S)A = \operatorname{softmax}(S) is taken row by row; O=AVO = AV is n×dvn\times d_v.
  • LL is a scalar loss that depends on QQ, KK and VV only through OO. The upstream gradient is G=∇OLG = \nabla_O L, and G~:=∇AL\tilde G := \nabla_A L.
  • Indices: ii is a query position (a row of QQ, SS, AA, OO), jj a key position (a row of KK and VV, a column of SS and AA), ll a column of QQ and KK, and cc a column of VV and OO.
  • 1\mathbf{1} is the all-ones vector of whatever length the product needs. rowsum⁡(B):=B1\operatorname{rowsum}(B) := B\mathbf{1} is a column holding the row sums of a matrix BB, so rowsum⁡(B) 1⊤\operatorname{rowsum}(B)\,\mathbf{1}^\top copies each row's sum across that row.
  • The softmax page's row result: for one row, a=softmax⁡(s)a = \operatorname{softmax}(s) with upstream g~=∇aL\tilde g = \nabla_a L, the Jacobian J=diag⁡(a)−aa⊤J = \operatorname{diag}(a) - aa^\top is symmetric (the softmax page, Problem 3), so ∇sL=Jg~=a⊙g~−a (a⊤g~)\nabla_s L = J\tilde g = a\odot\tilde g - a\,(a^\top\tilde g), that is, ∇s=a⊙(g~−(a⊤g~)1)\nabla_s = a\odot(\tilde g - (a^\top\tilde g)\mathbf{1}).
  • Every matrix below has rows as positions, so every formula is the row formula stacked.
  • Self-attention: Q=XWQQ = XW_Q, K=XWKK = XW_K, V=XWVV = XW_V with X∈Rn×dmodelX \in \mathbb{R}^{n\times d_{\text{model}}}, so m=nm = n. HH is the number of heads and hh indexes them.

Builds on: Jacobians and the chain rule, The softmax Jacobian

Problems

  1. ·

    Give the shapes of SS, AA and OO, and the number of multiply-adds to form SS. What does that cost look like as the sequence length n=mn = m grows?

  2. ··

    Compute ∇VL\nabla_V L.

  3. ··

    Compute G~=∇AL\tilde G = \nabla_A L.

  4. ···

    Compute ∇SL\nabla_S L from G~\tilde G and AA, using the softmax Jacobian row by row, and write it as one matrix expression with no m×mm\times m Jacobians.

  5. ··

    Compute ∇QL\nabla_Q L.

  6. ··

    Compute ∇KL\nabla_K L.

  7. ···

    Self-attention: Q=XWQQ = XW_Q, K=XWKK = XW_K, V=XWVV = XW_V. Compute ∇WQL\nabla_{W_Q} L, ∇WKL\nabla_{W_K} L, ∇WVL\nabla_{W_V} L and ∇XL\nabla_X L.

  8. ··

    A causal mask sets Sij=−∞S_{ij} = -\infty for j>ij > i before the softmax. What are the masked entries of AA, and what are the masked entries of ∇SL\nabla_S L? Show it from Problem 4.

  9. ··

    If the entries of qq and kk are independent with mean 00 and variance 11, what is Var⁡(q⊤k)\operatorname{Var}(q^\top k)? What does dividing by d\sqrt d achieve, and where does the factor appear in the backward pass?

  10. ···

    Multi-head: O=[O1  ⋯  OH] WOO = [O_1\;\cdots\;O_H]\,W_O with Oh∈Rn×dvO_h \in \mathbb{R}^{n\times d_v} and WO∈RHdv×dmodelW_O \in \mathbb{R}^{Hd_v\times d_{\text{model}}}. Compute ∇WOL\nabla_{W_O} L and ∇OhL\nabla_{O_h} L.

Worked solutions

Problem 1

Give the shapes of SS, AA and OO, and the number of multiply-adds to form SS. What does that cost look like as the sequence length n=mn = m grows?

  1. QK⊤QK^\top is (n×d)(d×m)=n×m(n\times d)(d\times m) = n\times m, and so is SS.K⊤K^\top is d×md\times m, and dividing by the scalar d\sqrt d keeps the shape.
  2. AA is n×mn\times m.The softmax is applied to each row separately and keeps each row's length.
  3. O=AVO = AV is (n×m)(m×dv)=n×dv(n\times m)(m\times d_v) = n\times d_v.The inner dimensions agree: each query position gets a weighted sum of the mm value rows.
  4. Sij=1d∑lQilKjlS_{ij} = \tfrac1{\sqrt d}\sum_l Q_{il}K_{jl} takes dd multiply-adds, and there are nmnm entries.Each score is the dot product of query row ii with key row jj, both of length dd.
  5. S,A:n×mS, A: n\times m; O:n×dvO: n\times d_v; SS costs nmdnmd multiply-adds, quadratic in the sequence lengthWith n=mn = m the cost is n2dn^2d: doubling the sequence length quadruples it. O=AVO = AV costs another nmdvnmd_v, and AA itself has n2n^2 entries to store for the backward pass, so time and memory both grow quadratically.

Problem 2

Compute ∇VL\nabla_V L.

  1. Oic=∑jAijVjcO_{ic} = \sum_j A_{ij}V_{jc}.Index form of O=AVO = AV, which is linear in VV for fixed AA.
  2. VjcV_{jc} appears in OicO_{ic} for every ii, with coefficient AijA_{ij}, and in no entry of another column.Column cc of OO is built from column cc of VV only, and every query row uses key row jj.
  3. ∂L/∂Vjc=∑iGicAij\partial L/\partial V_{jc} = \sum_i G_{ic}A_{ij}.LL depends on VV only through OO; by step 2 the chain rule sums over the rows ii of OO in column cc.
  4. ∑iAijGic=(A⊤G)jc\sum_i A_{ij}G_{ic} = (A^\top G)_{jc}, with shapes (m×n)(n×dv)=m×dv(m\times n)(n\times d_v) = m\times d_v.The sum runs over AA's row index, which a product can only contract if AA is transposed.
  5. ∇VL=A⊤G\nabla_V L = A^\top G (m×dvm\times d_v)The shape of VV. The pattern is general: for C=PRC = PR, the gradient with respect to the right factor is P⊤∇CLP^\top\nabla_C L.

Problem 3

Compute G~=∇AL\tilde G = \nabla_A L.

  1. AijA_{ij} appears in OicO_{ic} for every cc, with coefficient VjcV_{jc}, and in no other row of OO.Problem 2, step 1: row ii of AA builds only row ii of OO.
  2. ∂L/∂Aij=∑cGicVjc\partial L/\partial A_{ij} = \sum_c G_{ic}V_{jc}.LL depends on AA only through OO, and step 1 limits the sum to row ii.
  3. ∑cGicVjc=(GV⊤)ij\sum_c G_{ic}V_{jc} = (GV^\top)_{ij}, with shapes (n×dv)(dv×m)=n×m(n\times d_v)(d_v\times m) = n\times m.The sum runs over VV's column index, so VV is transposed.
  4. G~=∇AL=GV⊤\tilde G = \nabla_A L = GV^\top (n×mn\times m)The shape of AA. The companion pattern to Problem 2: for C=PRC = PR, the gradient with respect to the left factor is (∇CL)R⊤(\nabla_C L)R^\top.

Problem 4

Compute ∇SL\nabla_S L from G~\tilde G and AA, using the softmax Jacobian row by row, and write it as one matrix expression with no m×mm\times m Jacobians.

Write sis_i, aia_i and g~i∈Rm\tilde g_i \in \mathbb{R}^m for row ii of SS, AA and G~\tilde G, as columns.

  1. ai=softmax⁡(si)a_i = \operatorname{softmax}(s_i), and aia_i depends on no other row of SS.The softmax is taken row by row, so row ii of SS reaches LL only through row ii of AA.
  2. ∇siL=(diag⁡(ai)−aiai⊤) g~i=ai⊙(g~i−(ai⊤g~i)1)\nabla_{s_i} L = (\operatorname{diag}(a_i) - a_ia_i^\top)\,\tilde g_i = a_i\odot(\tilde g_i - (a_i^\top\tilde g_i)\mathbf{1}).The softmax page's row result with a=aia = a_i and upstream ∇aiL=g~i\nabla_{a_i}L = \tilde g_i, row ii of G~\tilde G; by step 1, aia_i is the only path from sis_i to LL, so no other row's upstream enters.
  3. ai⊤g~i=∑jAijG~ija_i^\top\tilde g_i = \sum_j A_{ij}\tilde G_{ij} is entry ii of rowsum⁡(A⊙G~)\operatorname{rowsum}(A\odot\tilde G), n×1n\times 1.A dot product of two rows is the sum along that row of their elementwise product.
  4. Row ii of G~−rowsum⁡(A⊙G~) 1⊤\tilde G - \operatorname{rowsum}(A\odot\tilde G)\,\mathbf{1}^\top is (g~i−(ai⊤g~i)1)⊤\big(\tilde g_i - (a_i^\top\tilde g_i)\mathbf{1}\big)^\top, with shapes (n×1)(1×m)=n×m(n\times 1)(1\times m) = n\times m.The outer product with 1⊤\mathbf{1}^\top copies entry ii of the column across row ii, so each row subtracts its own scalar.
  5. ∇SL=A⊙(G~−rowsum⁡(A⊙G~) 1⊤)\nabla_S L = A\odot\big(\tilde G - \operatorname{rowsum}(A\odot\tilde G)\,\mathbf{1}^\top\big)n×mn\times m, the shape of SS: step 2 stacked, since an elementwise product acts row by row. Each row's Jacobian would be m×mm\times m, and the whole matrix's nm×nmnm\times nm; neither is built, and the cost is O(nm)O(nm). In array code the rowsum is a sum along the last axis with the dimension kept, so it broadcasts.

Problem 5

Compute ∇QL\nabla_Q L.

  1. Sij=1d∑lQilKjlS_{ij} = \tfrac1{\sqrt d}\sum_l Q_{il}K_{jl}.Index form of S=QK⊤/dS = QK^\top/\sqrt d, which is linear in QQ for fixed KK.
  2. QilQ_{il} appears in SijS_{ij} for every jj, with coefficient Kjl/dK_{jl}/\sqrt d, and in no other row of SS.Query row ii is dotted with every key row, and builds only row ii of SS.
  3. ∂L/∂Qil=1d∑j(∇SL)ijKjl=1d((∇SL)K)il\partial L/\partial Q_{il} = \tfrac1{\sqrt d}\sum_j (\nabla_S L)_{ij}K_{jl} = \tfrac1{\sqrt d}\big((\nabla_S L)K\big)_{il}.LL depends on QQ only through SS; step 2 limits the chain rule to row ii. The sum runs over KK's row index, which is also the column index of ∇SL\nabla_S L, so no transpose is needed.
  4. ∇QL=1d(∇SL)K\nabla_Q L = \tfrac1{\sqrt d}(\nabla_S L)K (n×dn\times d)(n×m)(m×d)(n\times m)(m\times d), the shape of QQ. Problem 3's left-factor pattern with C=SC = S, P=Q/dP = Q/\sqrt d and R=K⊤R = K^\top gives ∇PL=(∇SL)K\nabla_P L = (\nabla_S L)K; since P=Q/dP = Q/\sqrt d, each QilQ_{il} enters LL through PilP_{il} times 1d\tfrac1{\sqrt d}, so ∇QL=1d∇PL\nabla_Q L = \tfrac1{\sqrt d}\nabla_P L.

Problem 6

Compute ∇KL\nabla_K L.

  1. KjlK_{jl} appears in SijS_{ij} for every ii, with coefficient Qil/dQ_{il}/\sqrt d, and in no other column of SS.Problem 5, step 1: key row jj is dotted with every query row and builds only column jj of SS.
  2. ∂L/∂Kjl=1d∑i(∇SL)ijQil\partial L/\partial K_{jl} = \tfrac1{\sqrt d}\sum_i (\nabla_S L)_{ij}Q_{il}.LL depends on KK only through SS, and step 1 limits the chain rule to column jj.
  3. ∑i(∇SL)ijQil=((∇SL)⊤Q)jl\sum_i (\nabla_S L)_{ij}Q_{il} = \big((\nabla_S L)^\top Q\big)_{jl}, with shapes (m×n)(n×d)=m×d(m\times n)(n\times d) = m\times d.The sum runs over the row index of ∇SL\nabla_S L, so it is transposed.
  4. ∇KL=1d(∇SL)⊤Q\nabla_K L = \tfrac1{\sqrt d}(\nabla_S L)^\top Q (m×dm\times d)The shape of KK. Equivalently, S⊤=KQ⊤/dS^\top = KQ^\top/\sqrt d has KK as its left factor, and ∇S⊤L=(∇SL)⊤\nabla_{S^\top}L = (\nabla_S L)^\top.

Problem 7

Self-attention: Q=XWQQ = XW_Q, K=XWKK = XW_K, V=XWVV = XW_V. Compute ∇WQL\nabla_{W_Q} L, ∇WKL\nabla_{W_K} L, ∇WVL\nabla_{W_V} L and ∇XL\nabla_X L.

Here m=nm = n, WQ,WK∈Rdmodel×dW_Q, W_K \in \mathbb{R}^{d_{\text{model}}\times d} and WV∈Rdmodel×dvW_V \in \mathbb{R}^{d_{\text{model}}\times d_v}. ∇QL\nabla_Q L, ∇KL\nabla_K L and ∇VL\nabla_V L are Problems 5, 6 and 2, each taken with the other two inputs held fixed.

  1. ∇WQL=X⊤∇QL\nabla_{W_Q}L = X^\top\nabla_Q L, (dmodel×n)(n×d)=dmodel×d(d_{\text{model}}\times n)(n\times d) = d_{\text{model}}\times d.WQW_Q reaches LL only through Q=XWQQ = XW_Q, where it is the right factor: Problem 2's pattern.
  2. ∇WKL=X⊤∇KL\nabla_{W_K}L = X^\top\nabla_K L and ∇WVL=X⊤∇VL\nabla_{W_V}L = X^\top\nabla_V L, dmodel×dd_{\text{model}}\times d and dmodel×dvd_{\text{model}}\times d_v.The same argument for K=XWKK = XW_K and V=XWVV = XW_V; each weight matrix feeds one projection only.
  3. The part of ∇XL\nabla_X L through QQ is ∇QL WQ⊤\nabla_Q L\,W_Q^\top, (n×d)(d×dmodel)=n×dmodel(n\times d)(d\times d_{\text{model}}) = n\times d_{\text{model}}.In Q=XWQQ = XW_Q, XX is the left factor: Problem 3's pattern.
  4. The parts through KK and VV are ∇KL WK⊤\nabla_K L\,W_K^\top and ∇VL WV⊤\nabla_V L\,W_V^\top, both n×dmodeln\times d_{\text{model}}.The same pattern; ∇KL\nabla_K L is n×dn\times d because m=nm = n.
  5. ∇WQL=X⊤∇QL\nabla_{W_Q}L = X^\top\nabla_Q L; ∇WKL=X⊤∇KL\nabla_{W_K}L = X^\top\nabla_K L; ∇WVL=X⊤∇VL\nabla_{W_V}L = X^\top\nabla_V L; ∇XL=∇QL WQ⊤+∇KL WK⊤+∇VL WV⊤\nabla_X L = \nabla_Q L\,W_Q^\top + \nabla_K L\,W_K^\top + \nabla_V L\,W_V^\topXX feeds all three projections, and the chain rule adds the contributions of every path from a variable to LL. Each gradient has the shape of its variable.

Problem 8

A causal mask sets Sij=−∞S_{ij} = -\infty for j>ij > i before the softmax. What are the masked entries of AA, and what are the masked entries of ∇SL\nabla_S L? Show it from Problem 4.

Here m=nm = n, so SS is square and the mask keeps the diagonal and everything below it.

  1. Aij=eSij/∑j′≤ieSij′A_{ij} = e^{S_{ij}}\big/\sum_{j'\le i}e^{S_{ij'}} for every jj.Row-wise softmax; the masked terms of the denominator are e−∞=0e^{-\infty} = 0, so only j′≤ij' \le i remain, and the j′=ij' = i term makes the sum positive.
  2. Aij=0A_{ij} = 0 for j>ij > i, and the unmasked entries of each row sum to 11.The numerator is e−∞=0e^{-\infty} = 0; the softmax renormalises over the positions a query may see.
  3. (∇SL)ij=Aij(G~ij−∑j′Aij′G~ij′)(\nabla_S L)_{ij} = A_{ij}\big(\tilde G_{ij} - \sum_{j'}A_{ij'}\tilde G_{ij'}\big).Entry (i,j)(i,j) of Problem 4's formula.
  4. For j>ij > i the bracket is finite and the factor AijA_{ij} is 00, so (∇SL)ij=0(\nabla_S L)_{ij} = 0.G~=GV⊤\tilde G = GV^\top is finite, so the bracket is too; zero times a finite number is exactly zero, not a small number.
  5. The row sum ∑j′Aij′G~ij′\sum_{j'}A_{ij'}\tilde G_{ij'} has zero terms at the masked j′j'.Step 2: those Aij′A_{ij'} are 00.
  6. masked Aij=0A_{ij} = 0, and since ∇SL=A⊙(⋯ )\nabla_S L = A\odot(\cdots), masked (∇SL)ij=0(\nabla_S L)_{ij} = 0 exactly; the unmasked entries follow Problem 4 with the masked columns contributing nothing to the row sumsSo no gradient reaches QQ or KK through a masked score: in Problems 5 and 6 those entries of ∇SL\nabla_S L multiply by zero. In code the −∞-\infty is often a large negative number such as −109-10^9; its exponential underflows to 00 in floating point, so the same holds.

Problem 9

If the entries of qq and kk are independent with mean 00 and variance 11, what is Var⁡(q⊤k)\operatorname{Var}(q^\top k)? What does dividing by d\sqrt d achieve, and where does the factor appear in the backward pass?

Here q,k∈Rdq, k \in \mathbb{R}^d are one query row and one key row as columns, so q⊤k=∑lqlklq^\top k = \sum_l q_lk_l is a score before scaling.

  1. E[qlkl]=E[ql] E[kl]=0\mathbb{E}[q_lk_l] = \mathbb{E}[q_l]\,\mathbb{E}[k_l] = 0.qlq_l and klk_l are independent, so the expectation of the product factors, and each mean is 00.
  2. Var⁡(qlkl)=E[ql2kl2]−02=E[ql2] E[kl2]=1\operatorname{Var}(q_lk_l) = \mathbb{E}[q_l^2k_l^2] - 0^2 = \mathbb{E}[q_l^2]\,\mathbb{E}[k_l^2] = 1.Independence again; E[ql2]=Var⁡(ql)+E[ql]2=1\mathbb{E}[q_l^2] = \operatorname{Var}(q_l) + \mathbb{E}[q_l]^2 = 1, and likewise for klk_l.
  3. Var⁡(∑lqlkl)=∑lVar⁡(qlkl)=d\operatorname{Var}\big(\sum_l q_lk_l\big) = \sum_l\operatorname{Var}(q_lk_l) = d.The terms for different ll are functions of disjoint sets of independent entries, so they are independent and their variances add.
  4. Var⁡(q⊤k/d)=d/d=1\operatorname{Var}(q^\top k/\sqrt d) = d/d = 1.Scaling a random variable by λ\lambda scales its variance by λ2\lambda^2, here λ2=1/d\lambda^2 = 1/d.
  5. Without the scaling the scores have standard deviation d\sqrt d, 88 at d=64d = 64, so the softmax of a row is concentrated on a few keys, often one.The softmax weights of two scores differ by the factor exp⁡(Δ)\exp(\Delta), where Δ\Delta is the gap between them. At this spread the top two scores of a row are typically 22 to 33 units apart, and Δ=3\Delta = 3 already gives exp⁡(3)≈20\exp(3) \approx 20.
  6. The more concentrated a row aa, the smaller diag⁡(a)−aa⊤\operatorname{diag}(a) - aa^\top, so ∇SL\nabla_S L shrinks and less gradient reaches QQ and KK.At a one-hot row, a standard basis vector uu, diag⁡(u)−uu⊤=0\operatorname{diag}(u) - uu^\top = 0, and the entries aj(1−aj)a_j(1 - a_j) and −ajaj′-a_ja_{j'} go to 00 as aa approaches uu.
  7. In the backward pass d\sqrt d appears only where SS is differentiated with respect to QQ and KK.∇VL\nabla_V L, G~\tilde G and ∇SL\nabla_S L are computed from AA, GG and VV (Problems 2 to 4), and none of their formulas contains dd.
  8. Var⁡(q⊤k)=d\operatorname{Var}(q^\top k) = d; dividing by d\sqrt d makes the scores order 1 so the softmax does not saturate; the same 1d\tfrac1{\sqrt d} multiplies ∇QL\nabla_Q L and ∇KL\nabla_K L (Problems 5–6) and nothing elseThe projection gradients of Problem 7 inherit it through ∇QL\nabla_Q L and ∇KL\nabla_K L and add no second factor. Trained queries and keys are not independent unit-variance vectors, but the scaling keeps the scores at order 1 at initialisation, where saturation would stop learning before it starts.

Problem 10

Multi-head: O=[O1  ⋯  OH] WOO = [O_1\;\cdots\;O_H]\,W_O with Oh∈Rn×dvO_h \in \mathbb{R}^{n\times d_v} and WO∈RHdv×dmodelW_O \in \mathbb{R}^{Hd_v\times d_{\text{model}}}. Compute ∇WOL\nabla_{W_O} L and ∇OhL\nabla_{O_h} L.

In this problem OO is the output of the multi-head layer, n×dmodeln\times d_{\text{model}}, G=∇OLG = \nabla_O L has the same shape, and OhO_h, for h=1,…,Hh = 1, \ldots, H, is head hh's output from the earlier problems. Write Ocat=[O1  ⋯  OH]O_{\text{cat}} = [O_1\;\cdots\;O_H], n×Hdvn\times Hd_v.

  1. ∇WOL=Ocat⊤G\nabla_{W_O}L = O_{\text{cat}}^\top G, (Hdv×n)(n×dmodel)=Hdv×dmodel(Hd_v\times n)(n\times d_{\text{model}}) = Hd_v\times d_{\text{model}}.O=OcatWOO = O_{\text{cat}}W_O with WOW_O the right factor: Problem 2's pattern.
  2. ∇OcatL=GWO⊤\nabla_{O_{\text{cat}}}L = GW_O^\top, (n×dmodel)(dmodel×Hdv)=n×Hdv(n\times d_{\text{model}})(d_{\text{model}}\times Hd_v) = n\times Hd_v.OcatO_{\text{cat}} is the left factor: Problem 3's pattern.
  3. Column (h−1)dv+c(h-1)d_v + c of OcatO_{\text{cat}} is column cc of OhO_h, for c=1,…,dvc = 1, \ldots, d_v.Concatenation places O1O_1 in columns 11 to dvd_v, O2O_2 in columns dv+1d_v+1 to 2dv2d_v, and so on; each entry of OhO_h is exactly one entry of OcatO_{\text{cat}}.
  4. ∇WOL=[O1  ⋯  OH]⊤G\nabla_{W_O}L = [O_1\;\cdots\;O_H]^\top G; ∇OhL\nabla_{O_h}L = columns (h−1)dv+1(h-1)d_v+1 to hdvhd_v of GWO⊤GW_O^\topn×dvn\times d_v, the shape of OhO_h. Concatenation copies entries without mixing them, so its backward pass slices the gradient into the same blocks. From there each head runs Problems 2 to 6 on its own block, and the heads' contributions to XX add, as the three paths did in Problem 7.

Where this goes wrong

1. Gradient to V without the transpose

O=AVO = AV puts AA on the left, and it is tempting to keep it there on the way back.

  1. Oic=∑jAijVjcO_{ic} = \sum_j A_{ij}V_{jc} and G=∇OLG = \nabla_O LRight so far: Problem 2, step 1.
  2. “VV is multiplied by AA, so its gradient is AA times the gradient coming back.”The shortcut that causes the mistake: the scalar rule, where the derivative of avav is aa, applied to matrices without asking which index the sum runs over.
  3. ∇VL=AG\nabla_V L = AGThe shapes (n×m)(n×dv)(n\times m)(n\times d_v) do not multiply unless n=mn = m, and even then it is wrong. The index form, ∂L/∂Vjc=∑iAijGic\partial L/\partial V_{jc} = \sum_i A_{ij}G_{ic}, puts the sum on AA's row index, so the product is A⊤GA^\top G. In self-attention n=mn = m, so the wrong version runs and trains on the wrong gradient.

2. Softmax Jacobian applied to the whole matrix

The softmax page's Jacobian is written for one vector, and flattening a matrix into a vector makes it look as though it applies at once. Write vec⁡(B)\operatorname{vec}(B) for the column of length nmnm that stacks the rows of an n×mn\times m matrix BB one after another (row-major).

  1. ∂a/∂s=diag⁡(a)−aa⊤\partial a/\partial s = \operatorname{diag}(a) - aa^\top for a=softmax⁡(s)a = \operatorname{softmax}(s)Right so far: the softmax page, Problem 3, for one row.
  2. “Flatten SS and AA to vectors and use the same Jacobian.”The shortcut that causes the mistake: treating the row-wise softmax as one softmax over all nmnm scores.
  3. ∇SL=(diag⁡(a)−aa⊤)g~\nabla_S L = (\operatorname{diag}(a) - aa^\top)\tilde g with a=vec⁡(A)a = \operatorname{vec}(A) and g~=vec⁡(G~)\tilde g = \operatorname{vec}(\tilde G)The softmax is per row: entries in different rows do not interact, so the true nm×nmnm\times nm Jacobian is block-diagonal, with the row Jacobians diag⁡(ai)−aiai⊤\operatorname{diag}(a_i) - a_ia_i^\top as blocks, and the row formula of Problem 4 is all there is. The flattened version subtracts a (a⊤g~)a\,(a^\top\tilde g), where a⊤g~a^\top\tilde g sums over every row, so each row is corrected by the whole matrix's total instead of its own.

3. Losing the 1/√d on the way back

The scaling is a fixed number with no parameter in it, and it is easy to decide that the backward pass can ignore it.

  1. S=QK⊤/dS = QK^\top/\sqrt d and ∇SL\nabla_S L from Problem 4Right so far.
  2. “1/d1/\sqrt d is a constant, not something we train, so it needs no gradient.”The analogy that causes the mistake: treating the scaling as part of the data, when it is part of the function SS of QQ and KK. A constant needs no gradient of its own, but it still multiplies the gradient of everything it scales.
  3. ∇QL=(∇SL)K\nabla_Q L = (\nabla_S L)KSS contains the factor, so its derivative with respect to QQ does too: 1d(∇SL)K\tfrac1{\sqrt d}(\nabla_S L)K (Problem 5). The error makes ∇QL\nabla_Q L (and ∇KL\nabla_K L, if the factor is dropped there too) d\sqrt d times too large, 88 times at d=64d = 64, which under plain SGD acts like a larger learning rate for WQW_Q (and WKW_K) and passes a check that only compares directions.

4. Only one path from X

Problem 7 has one input and three projections, and it is easy to follow only the one the problem starts with.

  1. ∇QL WQ⊤\nabla_Q L\,W_Q^\top is the gradient reaching XX through QQRight so far: Problem 7, step 3.
  2. “XX is the query input, so its gradient comes back through the queries.”The analogy that causes the mistake: cross-attention, where the queries come from XX and the keys and values from another sequence, so the query path really is the only one.
  3. ∇XL=∇QL WQ⊤\nabla_X L = \nabla_Q L\,W_Q^\topIn self-attention XX enters through QQ, KK and VV, and the chain rule sums over every path: the terms ∇KL WK⊤\nabla_K L\,W_K^\top and ∇VL WV⊤\nabla_V L\,W_V^\top are missing. The shape is right, so nothing fails; the layers below simply train on part of their gradient.

5. Masking A instead of S

Problem 8 shows the masked weights are zero, and zeroing them directly looks like a shortcut to the same place. Write MM for the n×nn\times n matrix with Mij=1M_{ij} = 1 for j≤ij \le i and 00 otherwise.

  1. The causal mask must make Aij=0A_{ij} = 0 for j>ij > iRight so far: Problem 8, step 2.
  2. “Compute the full softmax, then zero the entries a query may not see.”The shortcut that causes the mistake: masking the output of the softmax instead of its input.
  3. A←A⊙MA \leftarrow A\odot M after the softmaxRows no longer sum to 11, so the forward pass is wrong before the backward pass starts. Worse, each denominator still contains eSije^{S_{ij}} for the future positions j>ij > i, so the kept weights depend on later keys and gradient flows into them: the model can see the future. Mask the scores with −∞-\infty and the softmax renormalises.

Print this set: attention-backward.pdf (problems, answers, and worked solutions on separate pages).