Level U6 · Theory · Under the hood · runs in your browser

Backward through attention

How does the gradient get back through attention, and what must training keep in GPU memory for it?

Side trip · best after level 17 · The whole Transformer

This level uses the backward rules of U3 and the number formats of U2. If you skipped them, read them quickly first.

The full model trains with one line, loss.backward(). In U3 you built that line yourself for single numbers: every operation knows its own backward rule, and the chain rule connects the rules. Attention needs two rules that U3 didn’t have: one for softmax, and one for a matrix product.

This level writes them, takes one attention head backward by hand, and checks the result against numeric slopes. Then it counts what the backward pass costs: every value a backward rule needs must be kept in GPU memory from the forward pass until the backward pass uses it.

The level’s demo.py prints every number on this page. To run it, download it into a folder of its own and run python demo.py there (setup).

1. Softmax backward

Softmax turns scores zz into probabilities pp:

pi=ezi∑jezjp_i = \frac{e^{z_i}}{\sum_j e^{z_j}}

Backward receives dpdp, the gradient of the loss with respect to each pip_i, and must return dzdz. The new part: every score moves every probability, because all of them share the same denominator (the sum under the fraction line). If z1z_1 grows, p1p_1 grows. p0p_0 and p2p_2 shrink, so the total stays 1.

So the slopes form a table. Row ii, column jj holds ∂pi/∂zj\partial p_i / \partial z_j: how fast pip_i moves when zjz_j moves. This table of all slopes is the Jacobian JJ. For softmax it has a short formula:

Both lines together: J=diag(p)−p⊤pJ = \text{diag}(p) - p^\top p. Here diag(p)\text{diag}(p) is the table with pp on the diagonal and 0 elsewhere, and p⊤pp^\top p (np.outer(p, p)) is the table of every product pipjp_i p_j. Vectors are rows in this course, so p⊤pp^\top p is (n, 1) @ (1, n) → (n, n).

Take z=[ln⁡2,0,0]z = [\ln 2, 0, 0]. Then ez=[2,1,1]e^z = [2, 1, 1], the sum is 4, and p=[0.5,0.25,0.25]p = [0.5, 0.25, 0.25].

Softmax backward: one Jacobian, or one short formula

Type new scores z or a new incoming gradient dp. Both ways of computing dz always give the same numbers.

scores z
p = softmax(z)
0.500.250.25
dp (from the layer after)
J = diag(p) − pᵀp
z0z1z2
p00.2500??
p1??-0.0625
p2?-0.0625?
row sum
?
?
?
dz = dp @ J
???
dz = p ⊙ (dp − dp·p)
???
dp · p = ?, so each dzᵢ is pᵢ × (dpᵢ − dp · p).
row i of J: how pᵢ moves when each score moves ? answer a question below to see this number
Number Softmax gave p = [0.5, 0.25, 0.25]. Its Jacobian is J = diag(p) − pᵀp, so J[i][i] = pᵢ(1 − pᵢ). What is J[1][1]?
Go deeper Where the two formulas come from

Write s=∑jezjs = \sum_j e^{z_j}, so pi=ezi/sp_i = e^{z_i} / s. When ziz_i moves, both ezie^{z_i} and ss move, at the same rate ezie^{z_i}. The rule for a fraction gives

∂pi∂zi=ezis−eziezis2=pi−pi2=pi(1−pi).\frac{\partial p_i}{\partial z_i} = \frac{e^{z_i} s - e^{z_i} e^{z_i}}{s^2} = p_i - p_i^2 = p_i (1 - p_i).

When another score zjz_j moves, only ss moves: ∂pi/∂zj=−eziezj/s2=−pi pj\partial p_i / \partial z_j = -e^{z_i} e^{z_j} / s^2 = -p_i \, p_j.

🔒 Answer the question above to unlock

Now an entry off the diagonal. Its sign says which way p0p_0 goes when another score grows.

Number Softmax gave p = [0.5, 0.25, 0.25]. Off the diagonal, the Jacobian J = diag(p) − pᵀp has J[i][j] = −pᵢ·pⱼ. What is J[0][1]?
🔒 Answer the question above to unlock

One more fact about softmax: adding the same number cc to every score changes nothing. ezi+c=ec ezie^{z_i + c} = e^c \, e^{z_i}, and the factor ece^c appears in the numerator and in every term of the denominator, so it cancels.

ChooseFor p = [0.5, 0.25, 0.25], J = diag(p) − pᵀp, and J[i][j] is how fast pᵢ moves when zⱼ moves. Adding the same number c to every score leaves p unchanged. What is the sum of each row of J?
🔒 Answer the question above to unlock

From dp to dz

Each score zjz_j moves every pip_i, so dzjdz_j collects one part from each of them: dpi×Jijdp_i \times J_{ij}. This is the += of U3: a value used in several places gets the sum of all its parts.

dzj=∑idpi Jij,that is,dz=dp  Jdz_j = \sum_i dp_i \, J_{ij}, \quad \text{that is,} \quad dz = dp \; J

When you put the formula for JJ into the sum, the sum becomes short. dp  diag(p)dp \; \text{diag}(p) is just dpj pjdp_j \, p_j, and dp  p⊤pdp \; p^\top p is (dp⋅p) pj(dp \cdot p) \, p_j, where dp⋅pdp \cdot p is one number for the whole row. So

dz=p⊙(dp−dp⋅p)dz = p \odot (dp - dp \cdot p)

where ⊙\odot means entry by entry (* in NumPy). Subtract the one number dp⋅pdp \cdot p from every entry of dpdp, then multiply by pp.

Check both ways on numbers, with p=[0.5,0.25,0.25]p = [0.5, 0.25, 0.25] and dp=[1,1,0]dp = [1, 1, 0]:

  • the long way: column 0 of JJ is [0.25,−0.125,−0.125][0.25, -0.125, -0.125], so dz0=1×0.25+1×(−0.125)+0×(−0.125)=0.125dz_0 = 1 \times 0.25 + 1 \times (-0.125) + 0 \times (-0.125) = 0.125;
  • the short way: dp⋅p=0.5+0.25+0=0.75dp \cdot p = 0.5 + 0.25 + 0 = 0.75, so dz0=p0(dp0−0.75)=0.5×0.25=0.125dz_0 = p_0 (dp_0 - 0.75) = 0.5 \times 0.25 = 0.125.
Number Softmax gave p = [0.5, 0.25, 0.25], and the gradient arriving at p is dp = [1, 0, 2]. Use the short form for dz from above, without building J. What is dz[2]?
I got stuck here Why not just build J and multiply?

It gives the same numbers, but it is much more work. With nn scores, JJ has n×nn \times n entries, and the short form needs about 2n2n multiplications. In attention, each row of scores has one entry per token: with 1000 tokens, JJ for one row has a million entries, and there are 1000 rows per head. The short form never builds JJ.

🔒 Answer the question above to unlock

In NumPy, and checked

In attention, softmax runs along each row of a table, so dp·p is needed once per row: a sum along the last axis. A sum along an axis normally removes that axis: for a (2, 3) table, .sum(axis=-1) gives shape (2,). The argument keepdims=True keeps the axis with size 1, so the result is (2, 1), one number per row. Broadcasting can then subtract it from every entry of its own row.

How do you know a backward pass is right? Compare it with numeric slopes (L(z+h)−L(z−h))/2h(L(z + h) - L(z - h)) / 2h, as in level 6 and U5. A loss that makes the check easy is L=∑dp⊙softmax(z)L = \sum dp \odot \text{softmax}(z): its gradient at pp is exactly dpdp. The template does this check for you.

CodeWrite the backward pass of a softmax that ran along each row: return dz = p ⊙ (dp − dp·p), with dp·p computed row by row.

Enter keeps the indent · Tab indents · Esc then Tab leaves the editor · ⌘/Ctrl + Enter runs

Try it

In the lab, set the scores to z=[10,0,0]z = [10, 0, 0]. Then p≈[1,0,0]p \approx [1, 0, 0], and almost every entry of JJ is 0. Change dpdp to anything: dzdz stays near 0. A softmax that is almost certain passes almost no gradient back. That is one reason level 14 divides the scores by dk\sqrt{d_k}: it keeps the scores small, so softmax is not this certain at the start of training.

🔒 Answer the question above to unlock

2. One attention head, backward

The forward pass of one head (level 14) has three lines. Here are their shapes, with LL tokens:

ForwardShape
S = Q @ K.T / √dₖ(L, dₖ) @ (dₖ, L) → (L, L)
A = softmax(S), each row(L, L)
O = A @ V(L, L) @ (L, dᵥ) → (L, dᵥ)

Backward handles the lines in reverse order, starting with the last one. It starts with dO, which has the shape of O. Two rules do all the work:

  • A matrix product Y=XWY = X W (level 6): dX = dY @ W.T and dW = X.T @ dY.
  • Softmax, each row: section 1’s formula, with AA in place of pp.

The figure shows the order. Every gradient has the shape of its value, which makes the transposes easy to check:

GradientComes fromShape
dVO = A @ V(L, dᵥ)
dAO = A @ V(L, L)
dSA = softmax(S)(L, L)
dQS = Q @ K.T / √dₖ(L, dₖ)
dKS = Q @ K.T / √dₖ(L, dₖ)
dQdKdSdVdAdOQKSQ @ K.T/ √dₖAsoftmaxOA @ VV
Solid arrows: the forward pass. Dashed arrows: the backward pass. It starts with dO on the right, and each arrow is labelled with the gradient it gives. First O = A @ V gives dV and dA, then the softmax gives dS, then the scores give dQ and dK.

An example by hand

Two tokens, dk=4d_k = 4 (so dk=2\sqrt{d_k} = 2), dv=2d_v = 2. The queries and keys are chosen so that every score is 0:

Q=[11000011]K=[1−100001−1]V=[2011]Q = \begin{bmatrix} 1 & 1 & 0 & 0 \\ 0 & 0 & 1 & 1 \end{bmatrix} \quad K = \begin{bmatrix} 1 & -1 & 0 & 0 \\ 0 & 0 & 1 & -1 \end{bmatrix} \quad V = \begin{bmatrix} 2 & 0 \\ 1 & 1 \end{bmatrix}

Every dot product of a row of QQ with a row of KK is 0, so SS is all zeros and every weight in AA is 0.5. Each output row is the average of the rows of VV: O=[[1.5,0.5],[1.5,0.5]]O = [[1.5, 0.5], [1.5, 0.5]]. The gradient that arrives from the layers above is dO=[[1,0],[0,2]]dO = [[1, 0], [0, 2]].

Number One attention head with 2 tokens: the weights are A = [[0.5, 0.5], [0.5, 0.5]] and the gradient arriving at the output O = A @ V is dO = [[1, 0], [0, 2]]. Use the matrix-product rule from the list above. What is dV[0][1]?
🔒 Answer the question above to unlock

The product O=AVO = A V also gives dAdA: dA=dO  V⊤=[[2,1],[0,2]]dA = dO \; V^\top = [[2, 1], [0, 2]]. Row 0 of dAdA says how the loss changes when token 0 gives more weight to each token. Now the softmax step, one row at a time.

Number Two tokens. The weights are A = [[0.5, 0.5], [0.5, 0.5]] and the gradient arriving at them is dA = [[2, 1], [0, 2]]. Softmax backward, row by row: dS = A ⊙ (dA − (dA·A for that row)). What is dS[1][0]?
🔒 Answer the question above to unlock

The full score gradient is dS=[[0.25,−0.25],[−0.5,0.5]]dS = [[0.25, -0.25], [-0.5, 0.5]]. Each row adds up to 0, like the rows of the Jacobian. The last step is the backward pass of S=QK⊤/dkS = Q K^\top / \sqrt{d_k}. The forward pass divided the scores by dk\sqrt{d_k}, so both dQdQ and dKdK are divided by dk\sqrt{d_k} too.

Number Two tokens, dₖ = 4. The score gradient is dS = [[0.25, −0.25], [−0.5, 0.5]] and the keys are K = [[1, −1, 0, 0], [0, 0, 1, −1]]. The scores were S = Q @ K.T / √dₖ. What is dQ[1][0]?
🔒 Answer the question above to unlock

Now the keys. SS has one row per query and one column per key: S[1][0]S[1][0] is query 1 against key 0. Look at where key jj appears in S=QK⊤/dkS = Q K^\top / \sqrt{d_k}.

Number The same example: Q = [[1, 1, 0, 0], [0, 0, 1, 1]], dₖ = 4, and the score gradient is dS = [[0.25, −0.25], [−0.5, 0.5]]. What is dK[0][2], the gradient at entry 2 of key 0?
🔒 Answer the question above to unlock
I got stuck here Why does dK use dS.T, and dQ use dS?

Sij=qi⋅kj/dkS_{ij} = q_i \cdot k_j / \sqrt{d_k}. Query ii appears only in row ii of SS, so dqidq_i collects row ii of dSdS, each entry times its key: that is dS @ K / √dₖ. Key jj appears only in column jj of SS, so dkjdk_j collects column jj of dSdS. dS.T turns the columns into rows, and dS.T @ Q / √dₖ does the same sum for the keys.

The shapes give a quick check on any line you write. Try it with sizes that are all different.

Shape One attention head: L = 5 tokens, dₖ = 3, dᵥ = 6. Q and K are (5, 3), V is (5, 6). What shape is dK?
Go deeper With a batch, heads and a mask

A real model computes all of this for BB sequences and hh heads together. The tensors get two extra axes at the front: (B, h, L, L) for S and A. Every line stays the same. Only .T must become .swapaxes(-1, -2), so that the batch and head axes stay where they are.

The causal mask needs no extra backward code. A blocked score became −∞, so its weight in AA is 0. Every entry of dS = A * (…) is multiplied by its weight, so a blocked position gets exactly 0 gradient.

🔒 Answer the question above to unlock

The whole head

Now write the backward pass in NumPy: five lines, in the order of the second table, with the rules you just used by hand. The template runs it on the example above, then checks every gradient against numeric slopes, on random numbers with 3 tokens, dk=2d_k = 2 and dv=4d_v = 4.

CodeWrite the backward pass of one attention head: from dO, compute dV and dA, then dS (softmax backward, row by row), then dQ and dK. Several lines, all at the same indent.

Enter keeps the indent · Tab indents · Esc then Tab leaves the editor · ⌘/Ctrl + Enter runs

🔒 Answer the question above to unlock

3. What backward must keep in GPU memory

Look at the five backward lines again. Each one uses values from the forward pass: A, V, K, Q. Those values must still exist when the backward pass reaches them, so the forward pass keeps them. These kept values are called saved activations. In your code they were the tuple saved.

ChooseSoftmax backward computes dz = p ⊙ (dp − dp·p). To run it for A = softmax(S), which value from the forward pass must training keep?
🔒 Answer the question above to unlock

So one head keeps Q, K, V and A, and the other layers keep what their own backward rules need (the input of every matrix product, for example). Training uses GPU memory in two ways.

Per parameter: a fixed cost

Large models train in mixed precision (U2). The forward and backward passes use 16-bit numbers, bfloat16 (bf16 for short). The update is done on a 32-bit float32 copy (fp32 for short), so that small changes aren’t lost to rounding. With Adam (level 7), each parameter needs:

Kept for every parameterFormatBytes
the weight used in forward and backwardbf162
its gradientbf162
the master copy of the weightfp324
Adam’s mfp324
Adam’s vfp324

That is 16 bytes per parameter, whatever the batch or the length.

Number Training with mixed precision and Adam keeps, for every parameter: the weight in bf16 (2 bytes), its gradient in bf16 (2 bytes), an fp32 master copy (4 bytes), and Adam’s m and v in fp32 (4 bytes each). How many GB does a model with 7 billion parameters need for these alone? (1 GB = 10⁹ bytes)
🔒 Answer the question above to unlock

Per token: activations

Saved activations grow with the data in the batch. Most of them have one row per token, shape (B, L, d_model): with B=4B = 4, L=1000L = 1000, d_model = 1000 and 2 bytes per number, one such tensor is 8 MB. A layer keeps several of them, for example the inputs of its matrix products. The attention weights A are different: they have one number for every pair of tokens, in every head. Their shape is (B, h, L, L).

Number One attention layer keeps its weights A, of shape (B, h, L, L), for the backward pass. B = 4 sequences, h = 8 heads, L = 1000 tokens, 2 bytes per number. How many MB is that? (1 MB = 10⁶ bytes)
🔒 Answer the question above to unlock

How many (B, L, ·) tensors does one layer keep? Here is one pre-norm block (level 16), with an FFN of width 4 × d_model. The table lists what its backward rules need, step by step, in a plain implementation:

Kept for the backward passWidth
the input of the first LayerNormd_model
the input of the Q, K, V products (the LayerNorm’s output)d_model
Q, K and Vd_model each
the heads’ output, the input of the output productd_model
the input of the second LayerNormd_model
the input of the FFN’s first productd_model
the FFN’s hidden layer before ReLU4 × d_model
the FFN’s hidden layer after ReLU, the input of its second product4 × d_model
Number One pre-norm block keeps the tensors in the table above for its backward pass. Count them in units of one (B, L, d_model) tensor: a tensor of width 4 × d_model counts as 4. How many units?
🔒 Answer the question above to unlock

Now make the sequences longer and keep everything else.

ChooseFor every layer, training keeps activations of shape (B, L, width) and attention weights of shape (B, h, L, L). You double L from 1000 to 2000 and keep B, h and the width. The activations double. What happens to the saved attention weights?
🔒 Answer the question above to unlock

The lab adds up everything training keeps, for one model. Move LL and watch which part grows fastest.

What training keeps in GPU memory, and what grows with L

Move B and L, or switch recomputation on. The model stays the same: 24 layers, width 1000, 8 heads, 300 million parameters. Each mark on the scale is 10 times the one before.

Total 9.4 GB, within one 80 GB GPU. The attention weights are 16% of it.
Weights, gradients, Adam16 bytes × 300 million16 bytes × 300 million4.8 GB51%
Token activations16 × (B, L, d_model) × 24 layers16 × (B, L, d_model) × 24 layers3.1 GB33%
Attention weights(B, h, L, L) × 24 layers(B, h, L, L) × 24 layers1.5 GB16%

Token activations: the 16 tensors per layer that you counted above. On the scale, each boundary of the bar is the running total so far: the grey part always ends at 4.8 GB.

Go deeper Recomputation: more arithmetic, less GPU memory

The forward pass doesn’t have to keep everything. With recomputation, each layer keeps only its input. When the backward pass reaches that layer, it runs the layer’s forward pass again from the input, gets the values back, and uses them immediately.

In the example above, a layer then keeps one (4, 1000, 1000) input, 8 MB, instead of 64 MB for A alone plus everything else. The values of that one layer exist again while the backward pass is inside it, so the peak is the kept inputs plus one layer’s values. The cost is one more forward pass. A backward pass costs about twice a forward pass. So one training step goes from 1 + 2 = 3 units of work to 3 + 1 = 4: about 33% more arithmetic, and much less GPU memory.

A kernel here is one program that runs on the GPU. A fused kernel does several steps in one program. Fused attention kernels save even more GPU memory than recomputation. They keep Q, K, V, O and one number per row (from the softmax’s sum), but not A. They recompute A in small blocks during the backward pass, so nothing of size L×LL \times L is kept. A later part of the course builds this. The softmax step still needs dA⋅AdA \cdot A for each row, and it can come from O and dO alone: dA=dO V⊤dA = dO \, V^\top, so ∑jdAijAij=∑cdOic∑jAijVjc=∑cdOicOic\sum_j dA_{ij} A_{ij} = \sum_c dO_{ic} \sum_j A_{ij} V_{jc} = \sum_c dO_{ic} O_{ic}. In section 2’s example, row 0 gives dO⋅O=1×1.5+0×0.5=1.5dO \cdot O = 1 \times 1.5 + 0 \times 0.5 = 1.5, and dA⋅A=2×0.5+1×0.5=1.5dA \cdot A = 2 \times 0.5 + 1 \times 0.5 = 1.5.

Number Section 2’s example: the outputs are O = [[1.5, 0.5], [1.5, 0.5]] and the gradient arriving at them is dO = [[1, 0], [0, 2]]. A fused kernel needs dA·A for each row without keeping A, so it uses dO·O instead. What is dO·O for row 1?

Recap

a summary for when you finish the level

The key formulas and common mistakes appear here once you clear the level.

You can now

  • Write the Jacobian of softmax, and compute dz with the short form p ⊙ (dp − dp·p) without building it.
  • Take one attention head backward from dO to dQ, dK and dV, say the shape of every step, and check it against numeric slopes.
  • Estimate the GPU memory of training: 16 bytes per parameter, plus saved activations that grow with B × L × d<sub>model</sub> and B × h × L².

Keep in mind

  • , and the sum of every row of is 0
  • s = (dp * p).sum(-1, keepdims=True), then dz = p * (dp - s)
  • A matrix product goes backward as dX = dY @ W.T and dW = X.T @ dY; every gradient has the shape of its value
  • Query is in row of the scores and key in column : dK needs the transpose of dS, dQ does not, and both keep the
  • Mixed-precision Adam: 2 + 2 + 4 + 4 + 4 = 16 bytes per parameter
  • One pre-norm block keeps about 16 tensors of shape (B, L, d<sub>model</sub>) for its backward pass, besides A
  • Saved attention weights: B × h × L × L numbers per layer (in a plain implementation; a fused kernel keeps none)

Common mistakes

  • Forgetting the 1/√dₖ in dQ and dK, or confusing the rows and the columns of dS.
  • Counting only the bf16 weights (2 bytes each) for training: gradients, the master copy and Adam’s m and v make it 16.

Press ? for keyboard shortcuts

Reading mode · every part open, no stars