The Chain Rule and Backpropagation
The chain rule says rates of change multiply along nested functions; backpropagation is the chain rule executed backwards over a computation graph, one sweep for all gradients. This pair is THE algorithm that made deep learning possible — here it is from first principles, with code you can write from memory.
Backprop Is the Chain Rule in Reverse
Blue: forward pass computes values. Red: backward pass multiplies local derivatives from the output back to every input — each edge contributes its local slope, accumulated along every path.
01.The Problem: The Loss Sits at the End of a Long Chain
Here's the setup. Topic 9 taught you to measure a slope: nudge the input, watch the output. Topic 10 taught you to do it one knob at a time.
But a neural network isn't one function. It's thousands of functions stacked.
Say you have:
- a multiply:
a = x · w - then an error:
e = a − y - then a square:
L = e²
The loss L depends on w... but only through a. You can't just ask "slope of L against w" directly — L doesn't even see w. It sees a.
So the question becomes
If I nudge w a tiny bit, how does the effect travel through the middle steps and arrive at L?
And the scaling question:
A network has billions of weights. How do we answer this for every weight without rebooting the whole calculation billions of times?
The first question is answered by the chain rule. The second by backpropagation. Together they are the reason deep learning exists.
02.The Idea in Plain Words: Rates Multiply Along Chains
The chain rule is simply
Effect size in = effect size out, multiplied step by step.
Formally: if h(x) = f(g(x)), then
h'(x) = f'(g(x)) · g'(x)
Derivative of the outer function, evaluated at the inner output, times the derivative of the inner.
Example: h(x) = sin(x²). Outer: sin, whose slope at input u is cos(u). Inner: x², slope 2x. So h'(x) = cos(x²) · 2x.
Deeper nesting just keeps multiplying: y = f₃(f₂(f₁(x))) gives dy/dx = f₃' · f₂' · f₁' (each evaluated at the right point).
The one-line intuition: rates of change multiply along chains. A 3× amplifier followed by a 0.1× amplifier yields a 0.3× attenuator. Nobody has to test the pair — you just multiply the gains.
With several variables (Topics 10–11), the same rule runs on vectors: if L depends on z, and z depends on the vector θ, then
∇θL = (∂L/∂z) · (∂z/∂θ)
— a scalar times a Jacobian (or its transpose, by layout convention). And when a variable flows into several places, the total effect is the sum over all paths through the graph. Multiply along each path, add across paths. Tattoo that.
# Graph: a = x * w ; L = (a - y)**2 ; x=2, w=3, y=4
x, w, y = 2.0, 3.0, 4.0
a = x * w # forward
L = (a - y) ** 2
# backward: local derivatives
dL_da = 2 * (a - y) # 4.0 (outer: square)
da_dw = x # 2.0 (inner: product rule)
da_dx = w # 3.0
dL_dw = dL_da * da_dw # 8.0 <- chain rule!
dL_dx = dL_da * da_dx # 12.0
print(dL_dw, dL_dx)
w -= 0.1 * dL_dw # one SGD step: 3.0 -> 2.203.A Simple Worked Example: Three Numbers, Two Ops
Walk the code block above with your eyes. Inputs: x = 2, w = 3, y = 4.
Forward pass — just compute values:
a = x · w = 2 · 3 = 6L = (a − y)² = (6 − 4)² = 4
The model predicts 6, truth is 4, loss is 4. Done.
Backward pass — local slopes, multiplied along the wire:
Each gate knows only its own little derivative. That's all it ever needs:
- Square gate:
dL/da = 2(a − y) = 2·2 = 4. "Push a by 1, L moves by about 4." - Multiply gate:
da/dw = x = 2andda/dx = w = 3. "Push w by 1, a moves by 2. Push x by 1, a moves by 3."
Chain rule — multiply along the path:
dL/dw = dL/da · da/dw = 4 · 2 = 8
dL/dx = dL/da · da/dx = 4 · 3 = 12
Read the verdict: at this setting, nudging w by +0.01 raises the loss by about +0.08. So shrink w. One SGD step with lr = 0.1: w ← 3 − 0.1·8 = 2.2. The linear model predicted a loss drop of 8 · 0.8 = 6.4 — even below zero, which is impossible for a squared loss. The real loss went from 4 to L = (2.2·2 − 4)² = 0.16: still a huge drop, but the tangent line oversold it because the parabola bends. Locally right, globally sloppy — exactly Topic 9's lesson.
That's all backprop is — on a network, the same 3-line pattern (local slope → multiply along wire → sum across branches) executed for every edge, once.
04.Visual Intuition: Values Go Right, Slopes Go Left
Draw the example as a computation graph: nodes are operations, wires carry numbers.
codeFORWARD PASS: values flow → BACKWARD PASS: slopes flow ← x=2 ──┐ seed: dL/dL = 1 ├─( × )── a=6 ──┐ │ w=3 ──┘ ( × ) ├─( − )─ e=2 ─( ² )─ L=4 ▼ │ │ │ local dL/da = 4 y=4 ───────┘ │─────┘ slope ┌────┴────┐ da/dw=2 │ da/dx=3 ▼ ▼ (swap!) x=2 w=3 ×da/dw=2 ×da/dx=3 ↓ ↓ dL/dw=8 dL/dx=12 └── w update, x update
Three patterns to internalize from the picture:
- Forward pass caches the values (a = 6, e = 2). The backward pass needs them —
dL/da = 2(a−y)couldn't be computed without knowing a and y. This is why training uses much more memory than inference, and why the multiply-gate can hand back "the other operand" as its local slope. - Backward runs in reverse order of forward. Every arrow flips.
- Each node does a tiny, local job: take whatever arrived from downstream, multiply by your own local slope, pass upstream. No node sees the whole graph.
The mermaid diagram above and the framework you use every day implement exactly this flip-and-multiply loop.
05.The Analogy: A Stack of Amplifier Pedals
Carry one analogy through the rest of the topic: an electric guitar wired through a stack of pedal amplifiers.
Guitar → pedal 1 → pedal 2 → pedal 3 → amp speaker.
Each pedal has a gain knob: ×3, ×0.5, ×10. The signal reaching the speaker has been multiplied by the gains in order.
Now you're the sound engineer, standing at the speaker, and the show is too loud. You want to know how much each pedal contributes to the volume:
Walk the chain backwards. At each pedal, note its local gain, and multiply.
- Pedal nearest the speaker is ×4, the one before it ×2: turning that one affects the speaker by 4·2 = 8. That's
dL/dw = 8from the worked example. The backward walk is the chain rule. - Twenty attenuator pedals in a row, each ×0.25, and the last pedal's tweak arrives at the speaker multiplied by
0.25²⁰ ≈ 10⁻¹². The guitar might as well not exist. This is the vanishing gradient problem, one multiplication away from obvious. - Twenty ×10 pedals and a whisper becomes a deafening explosion — the exploding-gradient case.
- The fix everyone discovered: run a direct bypass wire around each pedal that adds the untouched signal back (gain ×1 in parallel). Then turning the guitar knob always reaches the speaker at full strength — no product can shrink below the bypass. That wire is the residual connection (Topic 2), and it is why we can train 100-layer nets.
- A tweak to one pedal that feeds two output paths? Add the two effects — sum over paths.
Every deep-learning horror story ("early layers never learn", "loss went NaN", "why do residual connections work?") is a pedal-stack story. Section 7 replays them all.
06.Why AI Cares: Backpropagation = Reverse-Mode Autodiff
A neural network is one gigantic composition: L = loss(W₃ ∘ σ ∘ W₂ ∘ σ ∘ W₁ ∘ x). How do you get the slope against all p weights?
The naive ways all collapse:
- Finite differences (Topic 9's code): nudge one parameter, re-run the whole forward pass, watch the loss. Billions of parameters → billions of evaluations per step. Dead end.
- Forward-mode chain rule: propagate d(everything)/d(one input) — one full sweep per input. Also a dead end when p ≫ number of outputs.
- Reverse mode (backprop): exploit the fact that the loss is a scalar. Seed
dL/dL = 1at the output and sweep backwards, accumulatingupstream_grad × local_op_derivativeat each node. One sweep yields all p partial derivatives at once. Cost ≈ 2× the forward pass, no matter how many parameters.
Every framework automates precisely this: PyTorch's loss.backward(), TensorFlow's GradientTape, JAX's jax.grad (literally a composition of vector-Jacobian products over the traced program). Each op ships a VJP rule — its local derivative in matrix form. For a matmul Y = XW, the product rule gives dX = dY·Wᵀ and dW = Xᵀ·dY (Topic 5) — the chain rule with matrices, nothing more.
07.What the Chain Reveals: Vanishing and Exploding Gradients, and the Fixes
Because slopes multiply along the chain (pedal stack), a depth-L network's gradient is a product of L local Jacobians. The arithmetic decides your fate:
- Sigmoid layers (slope ≤ 0.25, Topic 9): the gradient shrinks like
0.25ᴸ→ early layers never learn. This killed 10-layer nets in the 1990s. - Large weights / ReLU stacks with bad initialization: product > 1 → gradients explode, updates rocket past the optimum, loss goes NaN.
- The remedies are all "fix the product" tricks, and now you can see exactly why each one works:
- Residual connections add a +1 (identity) term to each layer's Jacobian, so the chain always contains a bypass path whose product is 1 (Topics 2, 12) — the direct wire around the pedals.
- Normalization layers keep per-layer gains near 1, so products neither inflate nor decay.
- Careful init (Xavier/He) targets unit-variance products at step zero.
- Clipping bounds the total norm when the exploded case slips through anyway.
- Gradient checkpointing — recompute forward slices during the backward pass — trades compute for memory in every 2024–2026 LLM training run. It exists because the chain rule's memory bill (Section 4: backward needs cached forward values) got itemized: billions of activations would otherwise not fit in HBM.
Attention adds its own chain subtlety: the softmax Jacobian is cheap, but backprop through QKᵀ is where FlashAttention fuses tiles — to avoid materializing the O(T²) intermediates twice, once forward and once backward.
08.In Practice: Write Backprop by Hand Once (Do It)
The canonical exercise (karpathy-style "micrograd"): build scalar nodes with a forward method plus a per-op backward, then topologically sort the graph and call backward from the loss. Verify against finite differences — always. The four local rules to internalize:
- add: gradient passes through unchanged (distribute to all parents).
d(a+b)/da = 1. - mul: swap-and-multiply (each parent gets upstream × the other operand). That's
da/dw = xin the example. - max/select: route gradient to the winner only (this is why ReLU and attention's max-like ops behave — the losing branch gets 0).
- matmul: upstream times the other factor, transposed (Topic 5):
dX = dY·Wᵀ,dW = Xᵀ·dY.
Once you see that loss.backward() is just this loop over billions of nodes in a DAG — cached forward values × local derivatives × upstream accumulation, in reverse topological order — "magic autodiff" becomes a bookkeeping algorithm. And debugging NaN losses becomes a skill: walk the graph backwards and find which local slope first blew up.
One habit separates people who've done the exercise from people who've read about it: they know gradients sum across paths when a variable feeds multiple branches (shared weight matrices, Siamese nets — Section 2's rule).
Architectural Trade-offs & Production Realities
Architectural Advantages
- All p partials in one reverse sweep — O(#ops), independent of parameter count; makes deep training computable.
- Modular: each op only needs its local VJP rule; new layers integrate without touching the engine.
- Exact (to float precision), unlike finite differences; composes cleanly with matmul kernels on GPUs.
Trade-offs & Constraints
- Must cache forward activations → memory scales with depth × batch × width (hence activation checkpointing).
- Multiplicative chains amplify pathologies: vanishing/exploding gradients are arithmetic destiny.
- Control flow, discrete ops, and custom kernels complicate graph tracing; gradients of nondifferentiable steps need surrogates.
A GPT training step traces thousands of ops per layer; the backward engine executes cached VJP rules in reverse order while all-reducing gradients across devices. JAX composes transforms so jax.grad = VJP pullbacks, and gradient checkpointing recomputes forward slices to fit 100K-token contexts in HBM — the chain rule engineered for clusters.
Staff+ Engineering Takeaways
- The chain rule: rates of change multiply along nested functions; total derivative = sum over all paths through the graph.
- Backpropagation = reverse-mode autodiff: seed dL/dL = 1 and sweep backward multiplying local derivatives — all partials for ~2× forward cost.
- Scalar losses make reverse mode unbeatable; each op contributes a hand-written VJP rule (matmul: dX = dY·Wᵀ, dW = Xᵀ·dY).
- Vanishing/exploding gradients are multiplicative chain pathology; residuals, normalization, and good init repair the product.
- Autodiff memory (cached activations) scales with graph size — checkpointing is the practical trade of compute for memory.
Topic Knowledge Check
Exercise 1 of 3 • Test your architectural comprehension.
With a = x·w, L = (a − y)², x=2, w=3, y=4, what is dL/dw?
How clear and actionable was this distributed systems breakdown?