The chain rule
When one function feeds another, their slopes multiply. Learn the chain rule, use it to find a model’s gradient, and see that backpropagation, the method that trains every neural network, is the chain rule walked backwards.
Step 1 A function inside a function
A function can be built in stages, each feeding the next. y = (3x + 1)² works in two: first work out u = 3x + 1, then square it, y = u². The first stage is called the inner function and the second the outer function.
| x | u = 3x + 1 | y = u² |
|---|---|---|
| 0 | 3 × 0 + 1 = 1 | 1² = 1 |
| 1 | 3 × 1 + 1 = 4 | 4² = 16 |
| 2 | 3 × 2 + 1 = 7 | 7² = 49 |
Machine learning models are built the same way, from many simple stages: multiply by a weight, add, square, and so on.
Try it yourself
Set the data point and the weight. The chain runs forwards to the loss and backwards to the slope at every stage. Take steps and watch the loss fall as the weight moves.
L against w the slope dL/dw
forward:
p = w × x = 1 × 2 = 2
e = p − y = 2 − 3 = −1
L = e² = (−1)² = 1
backward:
dL/dL = 1
dL/de = 1 × 2e = 1 × (−2) = −2
dL/dp = dL/de × 1 = −2
dL/dw = dL/dp × x = (−2) × 2 = −4
step 1:
new w = 1 − 0.1 × (−4) = 1.4
dL/dw is negative: making w larger lowers the loss, so the step makes w larger.
Check by nudging w by 0.001: L goes from 1 to 0.996004, and (0.996004 − 1) ÷ 0.001 = −3.996, close to −4.
Programming exercise
Write the forward and backward walks for a one-weight model in plain Python, then train it. Save chain.py and test_chain.py in the same folder, fill in each function in chain.py, and run the tests:
python test_chain.py
"""The chain rule: programming exercise. Use the chain rule on a two-stage function, then write the forward andbackward walks for a model with one weight and train it, in plain Python.Run the tests from this folder: python test_chain.py The model predicts p = w * x. For one data point (x, y), the error ise = p - y and the loss is L = e squared.""" def chain_slope(a, b, x): """Return the slope of y = (a*x + b) squared at x, by the chain rule. The inner stage is u = a*x + b and the outer stage is y = u squared. """ raise NotImplementedError def forward(w, x, y): """Walk forwards and return (p, e, L): the prediction, the error and the loss.""" raise NotImplementedError def backward(w, x, y): """Walk backwards from dL/dL = 1 and return (dL_de, dL_dp, dL_dw). Each is the slope of the loss with respect to that value. """ raise NotImplementedError def measured_slope(w, x, y, h=1e-6): """Return dL/dw measured by nudging: (L at w + h, minus L at w - h) / (2h).""" raise NotImplementedError def fit(w, x, y, rate, steps): """Run gradient descent on w: take `steps` steps of w = w - rate * dL/dw, and return the final w.""" raise NotImplementedError Stuck? Do the forward walk first and keep every value; the backward walk multiplies by one stage’s slope at a time, starting from 1. The Solution tab has one way to write it.
In practice: backpropagation
Backpropagation is the chain rule. PyTorch records every operation as the forward walk runs, building the chain behind the scenes, and loss.backward() walks it back, multiplying slopes stage by stage exactly as in step 5. Module 2 builds this machinery from scratch.
Vanishing slopes. Many slopes below 1 multiplied together shrink towards 0, so the first layers of a deep network get almost no signal. The sigmoid function, once common in neural networks, has a slope of at most 0.25, so ten sigmoid stages can multiply the slope by as little as 0.25¹⁰ ≈ 0.00000095. Modern networks use functions such as ReLU, whose slope is exactly 1 for any positive input, and shortcut connections that let the signal skip stages.
Exploding slopes. Slopes above 1 multiply up instead, and one huge step can wreck training. A common fix is gradient clipping: if the gradient is longer than a set limit, scale it down to that length before taking the step.
Test your knowledge
01y = (2x − 1)³. What is dy/dx at x = 1?Show answer
Let u = 2x − 1, so y = u³. Then dy/du = 3u² and du/dx = 2. At x = 1, u = 1, so dy/dx = 3 × 1² × 2 = 6.
02Why are the slopes of the stages multiplied, not added?Show answer
Each stage scales the change it receives. If u moves 3 for each 1 that x moves, and y moves 8 for each 1 that u moves, then one step in x moves u by 3, and each of those 3 moves y by 8: 3 × 8 = 24.
03In backpropagation, what number does the backward walk start from, and why?Show answer
1. It is dL/dL, how fast L changes as L itself changes, and L always changes by exactly as much as itself. Every slope after it is that 1 multiplied by the slopes of the stages passed on the way back.
04A network has 20 stages, each with slope 0.5. By how much is the slope at the first stage multiplied?Show answer
0.5²⁰ ≈ 0.00000095. The first stages get almost no signal, which is called a vanishing gradient.