Multivariable chain rule

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

What is it?

When a variable influences the output through several paths, add the contributions of every path, each the product of the local derivatives along it: ∂z∂x=∑i∂z∂ui∂ui∂x\frac{\partial z}{\partial x} = \sum_i \frac{\partial z}{\partial u_i}\frac{\partial u_i}{\partial x}. In matrix form, Jacobians multiply. This is exactly what backpropagation computes on a network's graph.

Why does it exist?

Real computations branch and merge: a weight affects many neurons, which all affect the loss. The one-variable chain rule handles a single path; the multivariable version handles any directed acyclic graph of operations — any program without loops, or any unrolled one.

Intuition

Draw the computation as a graph: inputs at the bottom, the output at the top, each edge labelled with a local partial derivative. ∂(output)/∂(input)\partial(\text{output})/\partial(\text{input}) is the sum, over all paths, of the products of labels along the path. Backpropagation computes all these sums at once by sweeping the graph from the top, accumulating at each node "how much the output cares about me" (the adjoint).

Formal definition

If g:ℝn→ℝmg : \R^n \to \R^m is differentiable at xx and f:ℝm→ℝpf : \R^m \to \R^p at g(x)g(x), then

Jf∘g(x)=Jf(g(x)) Jg(x).J_{f\circ g}(x) = J_f\big(g(x)\big)\,J_g(x).

For scalar z=f(u1,…,um)z = f(u_1, \dots, u_m) with ui=ui(x)u_i = u_i(x): ∂z∂xj=∑i=1m∂z∂ui∂ui∂xj\frac{\partial z}{\partial x_j} = \sum_{i=1}^{m}\frac{\partial z}{\partial u_i}\frac{\partial u_i}{\partial x_j}.

Formulas

∂z∂x=∑i∂z∂ui ∂ui∂x\frac{\partial z}{\partial x} = \sum_{i}\frac{\partial z}{\partial u_i}\,\frac{\partial u_i}{\partial x}
JfL∘⋯∘f1=JfL⋯Jf2 Jf1J_{f_L\circ\cdots\circ f_1} = J_{f_L}\cdots J_{f_2}\,J_{f_1}
a deep network: a product of Jacobians
∂L∂W(ℓ)=δ(ℓ) (a(ℓ−1))𝖳,δ(ℓ)=((W(ℓ+1))𝖳δ(ℓ+1))⊙σ′(z(ℓ))\frac{\partial L}{\partial W^{(\ell)}} = \delta^{(\ell)}\,\big(a^{(\ell-1)}\big)^{\mathsf T}, \qquad \delta^{(\ell)} = \Big(\big(W^{(\ell+1)}\big)^{\mathsf T}\delta^{(\ell+1)}\Big)\odot\sigma'\big(z^{(\ell)}\big)
the backpropagation equations

How is it computed?

The product JL⋯J1J_L\cdots J_1 can be evaluated right-to-left (forward mode: cost ∝ number of inputs) or left-to-right (reverse mode: cost ∝ number of outputs). A loss has one output and millions of inputs, so reverse mode wins by a factor of millions — that choice is backpropagation.

Example

z=xy+sin⁡xz = xy + \sin x with x=t2x = t^2, y=ety = e^t. Two paths from tt to zz (through xx and through yy): dzdt=(y+cos⁡x)⋅2t+x⋅et=(et+cos⁡t2) 2t+t2et\frac{\dd z}{\dd t} = (y + \cos x)\cdot 2t + x\cdot e^t = (e^t + \cos t^2)\,2t + t^2e^t.

Why does it matter?

It is the mathematical content of backpropagation and of every autodiff library (PyTorch, JAX, TensorFlow). It also gives robot velocities through the kinematic chain and the sensitivities used in engineering design (adjoint methods in CFD and weather forecasting are the same reverse-mode idea).

Where it shows up in computing

  • Weather and climate modelling★★★★★advancedPhysics and simulation

    4D-Var data assimilation uses adjoint (reverse-mode) models to fit the initial state of forecasts.

Where it shows up in AI

  • Backpropagation★★★★★fundamentalAI and machine learning

    Backprop = the multivariable chain rule evaluated in reverse order on the network's graph.

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

    Forward and reverse mode are the two natural orders of multiplying the Jacobian chain.

Where is it used?

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

What depends on it

Exercises

1Computation

If w=f(x,y)w = f(x, y) with x=rcos⁡θx = r\cos\theta, y=rsin⁡θy = r\sin\theta, express ∂w/∂r\partial w/\partial r and ∂w/∂θ\partial w/\partial\theta.

Solution

wr=fxcos⁡θ+fysin⁡θw_r = f_x\cos\theta + f_y\sin\theta, wθ=−fxrsin⁡θ+fyrcos⁡θw_\theta = -f_x r\sin\theta + f_y r\cos\theta.

2AI

A network computes L=ℓ(W2 σ(W1x))L = \ell(W_2\,\sigma(W_1x)). Write ∂L/∂W1\partial L/\partial W_1 using the chain rule and say which factors are reused from ∂L/∂W2\partial L/\partial W_2.

Solution

Let z1=W1xz_1 = W_1x, a=σ(z1)a = \sigma(z_1), z2=W2az_2 = W_2a. With g=∂ℓ/∂z2g = \partial\ell/\partial z_2: ∂L/∂W2=g a𝖳\partial L/\partial W_2 = g\,a^{\mathsf T} and ∂L/∂W1=((W2𝖳g)⊙σ′(z1))x𝖳\partial L/\partial W_1 = \big((W_2^{\mathsf T}g)\odot\sigma'(z_1)\big)x^{\mathsf T}. The upstream gradient gg is computed once and reused: that sharing is what makes backprop cheap.

↑ ↓ to navigate · ↵ · Esc