Training
- model
- GPT-2 small 124M gradients
- sentence
- 5 predictions loss 4.263
- drawn
- dims 1–8 of 3,072 × 768
All steps
Blame flows backwards
GPT-2 small reads “The cat sat on the floor” and scores its 5 next-token guesses: mean loss 4.263. Training needs, for each of its 124M weights, how the loss would change if that weight moved. Backpropagation gets all of them in one pass from the loss back to the input.
loss 4.263 · 124M gradients in one backward passA linear layer: dW = Xᵀ · dY
Take block 12’s last matrix, y = x · W. Its gradient is dW = Xᵀ · dY: each weight’s blame is its input times the gradient arriving at its output, summed over the 5 positions. Hover a cell for its sum.
dW = Xᵀ · dY · summed over 5 positionsPassing it down: dX = dY · Wᵀ
The same layer also passes blame to its input, dX = dY · Wᵀ, and so on down: through GELU (times its slope), the up-projection, LayerNorm, and attention. Every weight matrix gets its dW on the way. The residual add copies the gradient straight past each sub-layer.
each layer: keep dW, pass dX downThe causal mask, backwards
Position i’s loss can only reach tokens at or before i, because attention never looked ahead. Each row is one position’s loss; each column is how strongly it reaches that token’s embedding.
loss at i → embeddings at j ≤ iThrough twelve blocks
The gradient reaches every block. Its size on the residual stream barely shrinks from block 12 to block 1, because each block adds to the stream and the gradient flows back through the adds unchanged. Only the embeddings see it grow, as LayerNorm divides their small vectors less.
residual adds keep the gradient aliveCheck it by nudging
To check, nudge one weight by ±0.0001, run the model forwards twice, and see how the loss moves. The slope matches backpropagation to six digits. That takes two passes per weight; backpropagation gets all 124M at the cost of about two.
(L(w + ε) − L(w − ε)) / 2ε = ∂L/∂w
Code
logits = model(ids[:, :-1])
loss = F.cross_entropy(logits.flatten(0, 1), ids[:, 1:].flatten())
loss.backward() # fills .grad of every parameter, last layer first
# what it does for each linear layer y = x @ W + b, given dy = ∂L/∂y:
W.grad += x.T @ dy # summed over positions (and the batch)
b.grad += dy.sum(0)
dx = dy @ W.T # passed down to the layer below
# check one weight numerically
w[i] += eps; up = loss_fn(); w[i] -= 2 * eps; down = loss_fn(); w[i] += eps
assert abs((up - down) / (2 * eps) - w.grad[i]) < 1e-6