Chain Rule -- The Decomposition Technique for Differentiating Composite Functions

The chain rule is the mathematical soul of the backpropagation algorithm.Outer derivative × inner derivative, propagated layer by layer.


Concept Analysis

Univariate Chain Rule

If y = g(x), z = f(y), then:

\[ \frac{dz}{dx} = \frac{dz}{dy} \cdot \frac{dy}{dx} \]

Backpropagation = Chain rule applied in reverse along the computational graph

x
y=g(x)
∂y/∂x
z=f(y)
∂z/∂y
L
∂L/∂z

During backpropagation, the gradientpropagates layer by layer from output to input: each layer's "local gradient" × "upstream gradient" = the gradient of that layer's parameters with respect to the loss.


Real-life Example

Gas pedal → Fuel consumption

Fuel consumption depends on speed, and speed depends on how deeply the gas pedal is pressed.

"If you press the gas pedal a little more, how much does fuel consumption increase?" = (rate of change of fuel consumption with respect to speed) × (rate of change of speed with respect to gas pedal).

This is the chain rule:Propagate changes layer by layer, multiply to get the total effect.。


Python Hands-on Practice

Example

import sympy as sp

x = sp.Symbol('x')
# f(g(x)): sin(x^2), outer sin, inner x^2
d_direct = sp.diff(sp.sin(x**2), x)
d_chain = sp.cos(x**2) * 2*x
print(f"direct differentiation: {d_direct}")
print(f"chain rule: cos(x^2)·2x = {d_chain}")
print(f"consistent: {sp.simplify(d_direct-d_chain)==0}")

# Backpropagation for a two-layer network
w, xi, yi = sp.symbols('w x_i y_i')
a = sp.Symbol('a')            # Define intermediate variables separately as symbols
a_expr = w * xi               # The specific expression of a with respect to w
loss = (a - yi)**2            # loss is defined as a function of the symbolic a

dL_da = sp.diff(loss, a)                      # ∂L/∂a
da_dw = sp.diff(a_expr, w)                    # ∂a/∂w
dL_dw_chain = dL_da.subs(a, a_expr) * da_dw   # Chain rule: Substitute a=w*x_i and multiply

loss_direct = loss.subs(a, a_expr)            # Fully expand loss as a function of w
dL_dw_direct = sp.diff(loss_direct, w)

print(f"\n"Chain rule ∂L/∂w: {sp.simplify(dL_dw_chain)}")
print(fDirect ∂L/∂w: {sp.simplify(dL_dw_direct)})
print(fConsistent: {sp.simplify(dL_dw_chain-dL_dw_direct)==0})

Output:

直接求导: 2*x*cos(x**2)
链式法则: cos(x^2)·2x = 2*x*cos(x**2)
一致: True

链式 ∂L/∂w: 2*x_i*(w*x_i - y_i)
直接 ∂L/∂w: 2*x_i*(w*x_i - y_i)
一致: True

Application Scenarios in AI

Backpropagation = Systematic application of the chain rule on computational graphs

PyTorch's autograd engine automatically records the computation graph during every tensor operation. When loss.backward() is called, it starts from the loss node and traverses the graph in reverse, applying the chain rule to each node: local gradient × upstream gradient = downstream gradient. Developers don't need to write any derivation code by hand.

Gradient Check

If you write a custom layer or loss function by hand, how do you verify that the backpropagation gradient computation is correct? Use the numerical gradient \( (f(x+h)-f(x-h))/2h \) as the "ground truth" and compare it with the gradient computed by autograd. The relative error should be less than 1e-5. This is the principle of torch.autograd.gradcheck in PyTorch.

Computational Graph Optimization

Deep learning compilers (such as TensorRT, XLA) optimize computation graphs: merge adjacent operations (operator fusion) to reduce memory reads and writes. But it must be ensured that gradient computation under the chain rule remains correct—this is an important constraint of graph optimization.


Other extensions