Training

Backprop

the chain rule, from the loss back to every weight

model
GPT-2 small 124M gradients
sentence
5 predictions loss 4.263
drawn
dims 1–8 of 3,072 × 768

All steps

  1. 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 pass

  2. A 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 positions

  3. Passing 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 down

  4. The 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 ≤ i

  5. Through 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 alive

  6. Check 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

Go deeper