Practice / Transformer pieces
Attention backward
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 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 , a path through 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 is the elementwise product.
- One head. , and , with rows as positions: query positions, key positions, width for queries and keys and for values.
- is ; is taken row by row; is .
- is a scalar loss that depends on , and only through . The upstream gradient is , and .
- Indices: is a query position (a row of , , , ), a key position (a row of and , a column of and ), a column of and , and a column of and .
- is the all-ones vector of whatever length the product needs. is a column holding the row sums of a matrix , so copies each row's sum across that row.
- The softmax page's row result: for one row, with upstream , the Jacobian is symmetric (the softmax page, Problem 3), so , that is, .
- Every matrix below has rows as positions, so every formula is the row formula stacked.
- Self-attention: , , with , so . is the number of heads and indexes them.
Builds on: Jacobians and the chain rule, The softmax Jacobian
Problems
- ·
Give the shapes of , and , and the number of multiply-adds to form . What does that cost look like as the sequence length grows?
- ··
Compute .
- ··
Compute .
- ···
Compute from and , using the softmax Jacobian row by row, and write it as one matrix expression with no Jacobians.
- ··
Compute .
- ··
Compute .
- ···
Self-attention: , , . Compute , , and .
- ··
A causal mask sets for before the softmax. What are the masked entries of , and what are the masked entries of ? Show it from Problem 4.
- ··
If the entries of and are independent with mean and variance , what is ? What does dividing by achieve, and where does the factor appear in the backward pass?
- ···
Multi-head: with and . Compute and .
Answers
- ; ; costs multiply-adds, quadratic in the sequence length
- ()
- ()
- ()
- ()
- ; ; ;
- masked , and since , masked exactly; the unmasked entries follow Problem 4 with the masked columns contributing nothing to the row sums
- ; dividing by makes the scores order 1 so the softmax does not saturate; the same multiplies and (Problems 5–6) and nothing else
- ; = columns to of
Worked solutions
Problem 1
Give the shapes of , and , and the number of multiply-adds to form . What does that cost look like as the sequence length grows?
- is , and so is . is , and dividing by the scalar keeps the shape.
- is .The softmax is applied to each row separately and keeps each row's length.
- is .The inner dimensions agree: each query position gets a weighted sum of the value rows.
- takes multiply-adds, and there are entries.Each score is the dot product of query row with key row , both of length .
- ; ; costs multiply-adds, quadratic in the sequence lengthWith the cost is : doubling the sequence length quadruples it. costs another , and itself has entries to store for the backward pass, so time and memory both grow quadratically.
Problem 2
Compute .
- .Index form of , which is linear in for fixed .
- appears in for every , with coefficient , and in no entry of another column.Column of is built from column of only, and every query row uses key row .
- . depends on only through ; by step 2 the chain rule sums over the rows of in column .
- , with shapes .The sum runs over 's row index, which a product can only contract if is transposed.
- ()The shape of . The pattern is general: for , the gradient with respect to the right factor is .
Problem 3
Compute .
- appears in for every , with coefficient , and in no other row of .Problem 2, step 1: row of builds only row of .
- . depends on only through , and step 1 limits the sum to row .
- , with shapes .The sum runs over 's column index, so is transposed.
- ()The shape of . The companion pattern to Problem 2: for , the gradient with respect to the left factor is .
Problem 4
Compute from and , using the softmax Jacobian row by row, and write it as one matrix expression with no Jacobians.
Write , and for row of , and , as columns.
- , and depends on no other row of .The softmax is taken row by row, so row of reaches only through row of .
- .The softmax page's row result with and upstream , row of ; by step 1, is the only path from to , so no other row's upstream enters.
- is entry of , .A dot product of two rows is the sum along that row of their elementwise product.
- Row of is , with shapes .The outer product with copies entry of the column across row , so each row subtracts its own scalar.
- , the shape of : step 2 stacked, since an elementwise product acts row by row. Each row's Jacobian would be , and the whole matrix's ; neither is built, and the cost is . In array code the rowsum is a sum along the last axis with the dimension kept, so it broadcasts.
Problem 5
Compute .
- .Index form of , which is linear in for fixed .
- appears in for every , with coefficient , and in no other row of .Query row is dotted with every key row, and builds only row of .
- . depends on only through ; step 2 limits the chain rule to row . The sum runs over 's row index, which is also the column index of , so no transpose is needed.
- (), the shape of . Problem 3's left-factor pattern with , and gives ; since , each enters through times , so .
Problem 6
Compute .
- appears in for every , with coefficient , and in no other column of .Problem 5, step 1: key row is dotted with every query row and builds only column of .
- . depends on only through , and step 1 limits the chain rule to column .
- , with shapes .The sum runs over the row index of , so it is transposed.
- ()The shape of . Equivalently, has as its left factor, and .
Problem 7
Self-attention: , , . Compute , , and .
Here , and . , and are Problems 5, 6 and 2, each taken with the other two inputs held fixed.
- , . reaches only through , where it is the right factor: Problem 2's pattern.
- and , and .The same argument for and ; each weight matrix feeds one projection only.
- The part of through is , .In , is the left factor: Problem 3's pattern.
- The parts through and are and , both .The same pattern; is because .
- ; ; ; feeds all three projections, and the chain rule adds the contributions of every path from a variable to . Each gradient has the shape of its variable.
Problem 8
A causal mask sets for before the softmax. What are the masked entries of , and what are the masked entries of ? Show it from Problem 4.
Here , so is square and the mask keeps the diagonal and everything below it.
- for every .Row-wise softmax; the masked terms of the denominator are , so only remain, and the term makes the sum positive.
- for , and the unmasked entries of each row sum to .The numerator is ; the softmax renormalises over the positions a query may see.
- .Entry of Problem 4's formula.
- For the bracket is finite and the factor is , so . is finite, so the bracket is too; zero times a finite number is exactly zero, not a small number.
- The row sum has zero terms at the masked .Step 2: those are .
- masked , and since , masked exactly; the unmasked entries follow Problem 4 with the masked columns contributing nothing to the row sumsSo no gradient reaches or through a masked score: in Problems 5 and 6 those entries of multiply by zero. In code the is often a large negative number such as ; its exponential underflows to in floating point, so the same holds.
Problem 9
If the entries of and are independent with mean and variance , what is ? What does dividing by achieve, and where does the factor appear in the backward pass?
Here are one query row and one key row as columns, so is a score before scaling.
- . and are independent, so the expectation of the product factors, and each mean is .
- .Independence again; , and likewise for .
- .The terms for different are functions of disjoint sets of independent entries, so they are independent and their variances add.
- .Scaling a random variable by scales its variance by , here .
- Without the scaling the scores have standard deviation , at , so the softmax of a row is concentrated on a few keys, often one.The softmax weights of two scores differ by the factor , where is the gap between them. At this spread the top two scores of a row are typically to units apart, and already gives .
- The more concentrated a row , the smaller , so shrinks and less gradient reaches and .At a one-hot row, a standard basis vector , , and the entries and go to as approaches .
- In the backward pass appears only where is differentiated with respect to and ., and are computed from , and (Problems 2 to 4), and none of their formulas contains .
- ; dividing by makes the scores order 1 so the softmax does not saturate; the same multiplies and (Problems 5–6) and nothing elseThe projection gradients of Problem 7 inherit it through and 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: with and . Compute and .
In this problem is the output of the multi-head layer, , has the same shape, and , for , is head 's output from the earlier problems. Write , .
- , . with the right factor: Problem 2's pattern.
- , . is the left factor: Problem 3's pattern.
- Column of is column of , for .Concatenation places in columns to , in columns to , and so on; each entry of is exactly one entry of .
- ; = columns to of , the shape of . 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 add, as the three paths did in Problem 7.
Where this goes wrong
1. Gradient to V without the transpose
puts on the left, and it is tempting to keep it there on the way back.
- and Right so far: Problem 2, step 1.
- “ is multiplied by , so its gradient is times the gradient coming back.”The shortcut that causes the mistake: the scalar rule, where the derivative of is , applied to matrices without asking which index the sum runs over.
- The shapes do not multiply unless , and even then it is wrong. The index form, , puts the sum on 's row index, so the product is . In self-attention , 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 for the column of length that stacks the rows of an matrix one after another (row-major).
- for Right so far: the softmax page, Problem 3, for one row.
- “Flatten and to vectors and use the same Jacobian.”The shortcut that causes the mistake: treating the row-wise softmax as one softmax over all scores.
- with and The softmax is per row: entries in different rows do not interact, so the true Jacobian is block-diagonal, with the row Jacobians as blocks, and the row formula of Problem 4 is all there is. The flattened version subtracts , where 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.
- and from Problem 4Right so far.
- “ 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 of and . A constant needs no gradient of its own, but it still multiplies the gradient of everything it scales.
- contains the factor, so its derivative with respect to does too: (Problem 5). The error makes (and , if the factor is dropped there too) times too large, times at , which under plain SGD acts like a larger learning rate for (and ) 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.
- is the gradient reaching through Right so far: Problem 7, step 3.
- “ 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 and the keys and values from another sequence, so the query path really is the only one.
- In self-attention enters through , and , and the chain rule sums over every path: the terms and 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 for the matrix with for and otherwise.
- The causal mask must make for Right so far: Problem 8, step 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.
- after the softmaxRows no longer sum to , so the forward pass is wrong before the backward pass starts. Worse, each denominator still contains for the future positions , so the kept weights depend on later keys and gradient flows into them: the model can see the future. Mask the scores with and the softmax renormalises.
Print this set: attention-backward.pdf (problems, answers, and worked solutions on separate pages).