Chain rule

Level FundamentalDifficulty ★★★★★Concept⌖ Open in the map

What is it?

The derivative of a composition is the product of the derivatives: (g∘f)′(x)=g′(f(x)) f′(x)(g\circ f)'(x) = g'(f(x))\,f'(x). Rates of change multiply along a chain — and backpropagation is this rule applied, very efficiently, to a neural network.

Why does it exist?

Almost every function we meet is built by composing simpler ones. Without the chain rule we would have to go back to limits for each new formula; with it, knowing the derivatives of a few primitives is enough to differentiate everything built from them — including a program.

Intuition

Gears. If uu turns 3 times as fast as xx, and yy turns 2 times as fast as uu, then yy turns 2×3=62 \times 3 = 6 times as fast as xx. In Leibniz notation the rule looks like cancelling fractions: dydx=dydu⋅dudx\frac{\dd y}{\dd x} = \frac{\dd y}{\dd u}\cdot\frac{\dd u}{\dd x}. Along a chain of nn functions you multiply nn local rates — and if they are all smaller than 1, the product vanishes (that is the vanishing-gradient problem).

Formal definition

If ff is differentiable at xx and gg is differentiable at f(x)f(x), then g∘fg \circ f is differentiable at xx and

(g∘f)′(x)=g′(f(x)) f′(x).(g \circ f)'(x) = g'\big(f(x)\big)\,f'(x).

Idea of the proof

Write g(f(x+h))−g(f(x))=(g′(f(x))+η(k)) kg(f(x+h)) - g(f(x)) = \big(g'(f(x)) + \eta(k)\big)\,k with k=f(x+h)−f(x)k = f(x+h) - f(x) and η(k)→0\eta(k) \to 0; divide by hh and let h→0h \to 0. (Dividing and multiplying by kk directly fails when k=0k = 0, which is why the auxiliary function η\eta is needed.)

Formulas

(g∘f)′(x)=g′(f(x)) f′(x),dydx=dydu dudx(g \circ f)'(x) = g'\big(f(x)\big)\,f'(x), \qquad \frac{\dd y}{\dd x} = \frac{\dd y}{\dd u}\,\frac{\dd u}{\dd x}
ddx fn(⋯f2(f1(x)))=fn′(un−1)⋯f2′(u1) f1′(x)\frac{\dd}{\dd x}\,f_n(\cdots f_2(f_1(x))) = f_n'(u_{n-1})\cdots f_2'(u_1)\,f_1'(x)
a chain of nn functions: a product of nn local derivatives

How is it computed?

Name the intermediate values: u1=f1(x)u_1 = f_1(x), u2=f2(u1)u_2 = f_2(u_1), … Compute the local derivatives fk′(uk−1)f_k'(u_{k-1}) and multiply. The order in which you multiply does not change the result, but it changes the cost when the functions have many inputs and outputs: left-to-right is forward mode, right-to-left is reverse mode (backpropagation).

Example

y=σ(wx+b)y = \sigma(wx + b) with σ\sigma the sigmoid. Set z=wx+bz = wx + b. Then ∂y∂w=σ′(z)⋅x=σ(z)(1−σ(z)) x\frac{\partial y}{\partial w} = \sigma'(z)\cdot x = \sigma(z)(1 - \sigma(z))\,x. This is literally the gradient a neural network computes for one weight of one neuron.

Why does it matter?

Backpropagation (Rumelhart, Hinton and Williams, 1986; Linnainmaa's reverse-mode AD, 1970) is the chain rule organised so that the gradient with respect to all weights costs about as much as one evaluation of the network. Without that trick, training modern models would be impossible.

Where it shows up in computing

  • Robot Jacobian (velocity kinematics)★★★★★frequentRobotics and control

    End-effector velocity is the chain rule through the arm: p˙=∂p∂q q˙\dot p = \frac{\partial p}{\partial q}\,\dot q.

Where it shows up in AI

  • Backpropagation★★★★★fundamentalAI and machine learning

    Backprop is the chain rule evaluated from the loss backwards, reusing every intermediate product.

  • Automatic differentiation★★★★★fundamentalAI and machine learning

    Forward and reverse mode are two orders of multiplying the same chain of local derivatives.

Where is it used?

Computing topics reachable from here, through the chain of ideas that leads to them:

What depends on it

Exercises

1Computation

Differentiate h(x)=ln⁡(1+e3x)h(x) = \ln\big(1 + e^{3x}\big).

Solution

h′(x)=11+e3x⋅3e3x=3 σ(3x)h'(x) = \frac{1}{1 + e^{3x}}\cdot 3e^{3x} = 3\,\sigma(3x) — the derivative of softplus is a sigmoid.

2AI

A 20-layer network uses sigmoid activations. Bound the factor by which the gradient can shrink when it goes through the 20 activations, ignoring the weights.

Solution

Each σ′≤14\sigma' \le \frac14, so the product is at most 4−20≈9⋅10−134^{-20} \approx 9 \cdot 10^{-13}: vanishing gradients. ReLU (derivative 1\text{derivative } 1 when active) and residual connections avoid it.

↑ ↓ to navigate · ↵ · Esc