多元微积分 多元微积分把导数和积分推广到多变量的函数,这至关重要,因为机器学习模型动辄有数百万个参数。本文件涵盖偏导数、梯度、雅可比矩阵、海森矩阵,以及让反向传播成为可能的多变量链式法则。 到目前为止,我们的函数都是接受单个输入 $x$、产生单个输出 $f(x)$。但在机器学习里,我们几乎从不止处理一个变量。 考虑一个二元函数,比如 $f(x, y) = x^2 + y^2$。它在三维空间里定义了一张曲面——一个碗的形状。我们想知道:如果在保持 $y$ 不变的同时,把 $x$ 稍微动一点点,$f$ 会怎么变?这就是偏导数(partial derivative)。
多元微积分把导数和积分推广到多变量的函数,这至关重要,因为机器学习模型动辄有数百万个参数。本文件涵盖偏导数、梯度、雅可比矩阵、海森矩阵,以及让反向传播成为可能的多变量链式法则。
到目前为止,我们的函数都是接受单个输入 x、产生单个输出 f(x)。但在机器学习里,我们几乎从不止处理一个变量。
考虑一个二元函数,比如 f(x, y) = x^2 + y^2。它在三维空间里定义了一张曲面——一个碗的形状。我们想知道:如果在保持 y 不变的同时,把 x 稍微动一点点,f 会怎么变?这就是偏导数(partial derivative)。
f 对 x 的偏导数记作 \frac{\partial f}{\partial x},做法是把其他所有变量都当成常数,然后照常对 x 求导。
对于 f(x, y) = x^2y + 3x - 2y:
在算 \frac{\partial f}{\partial x} 时,我们把 y 当成常数,所以 x^2y 求导得 2xy,3x 得 3,-2y 得 0。
在算 \frac{\partial f}{\partial y} 时,我们把 x 当成常数,所以 x^2y 求导得 x^2,3x 得 0,-2y 得 -2。
从几何上看,对 x 求偏导就像是拿一个平行于 xz 平面(固定某个 y 值)的平面去切这张三维曲面,然后求出切出来的那条曲线的斜率。
对于 f(x, y) = x^2 + y^2:\nabla f(x, y) = (2x, 2y)。在点 (1, 2) 处:\nabla f(1, 2) = (2, 4)。
梯度有两个关键性质:
方向:它指向函数增长最快的方向。想象一个登山者站在山上,他所处位置的梯度就指向正上方——沿最陡的那条路。
大小:\|\nabla f\| 给出在那个最陡方向上的增长率。梯度大说明地形陡峭,梯度小说明接近平坦。
既然梯度指向上坡,那么朝相反方向(-\nabla f)走就是下坡,走向更小的值。这个朴素的想法正是**梯度下降(gradient descent)**的基础——我们会在后续章节详细展开这种优化方法。目前只需要记住:梯度告诉你哪边是「上」,以及坡有多陡。
**方向导数(directional derivative)**推广了偏导数。它不再问「沿 x 轴方向 f 怎么变?」,而是问「沿任意方向 \mathbf{u},f 怎么变?」它的计算方法是梯度与一个单位向量的点积:
对于 f(x, y) = x^2 + y^2,在点 (1, 2) 处,沿方向 \mathbf{v} = (3, 4):先归一化得到 \mathbf{u} = (3/5, 4/5),再算 D_{\mathbf{u}} f = (2, 4) \cdot (3/5, 4/5) = 6/5 + 16/5 = 22/5。
偏导数其实是方向导数的特例——方向恰好沿某条坐标轴。如果某个方向上的方向导数为零,就说明函数在该点沿那个方向是平坦的。
等高线(contour lines)(也叫等值线)把函数值相同的点连起来。对于 f(x, y) = x^2 + y^2,等高线是以原点为圆心的同心圆:对应不同 c 值的 x^2 + y^2 = c。
等高线绝不会彼此相交(同一点不可能有两个不同的函数值)。
梯度始终垂直于等高线,并从低值指向高值。
等高线密集说明地形陡峭;稀疏说明坡度平缓。
到目前为止,我们的函数都只产生一个输出。但很多函数会产生多个输出。一个函数 \mathbf{F}: \mathbb{R}^n \to \mathbb{R}^m 接受 n 个输入、产生 m 个输出。**雅可比矩阵(Jacobian matrix)**把这种向量值函数的所有偏导数整齐地组织起来:
J = \begin{bmatrix} \frac{\partial f_1}{\partial x_1} & \cdots & \frac{\partial f_1}{\partial x_n} \\ \vdots & \ddots & \vdots \\ \frac{\partial f_m}{\partial x_1} & \cdots & \frac{\partial f_m}{\partial x_n} \end{bmatrix}
雅可比的每一行就是某一个输出分量的梯度。对于一个有 3 个输入、2 个输出的函数,雅可比是一个 2 \times 3 的矩阵。
雅可比把导数推广到了向量值函数。
正如标量函数的导数告诉你「输入每变一个单位,输出变多少」,雅可比告诉你「每一个输出相对于每一个输入怎么变」。
**雅可比行列式(determinant of the Jacobian)**衡量一个变换在局部把空间拉伸或压缩了多少。
如果行列式是 2,那么局部小区域的面积会翻倍。如果是 0,说明这个变换把空间压扁到了更低维度(回想一下矩阵那一章:行列式为零意味着这是一个奇异的、不可逆的变换)。
当多个变换被复合起来(一个接一个地喂入下一个),整体映射的雅可比就等于各个雅可比的乘积。我们在后续章节会看到这个想法变得至关重要。
如果说梯度捕捉的是一阶信息(斜率),那么**海森矩阵(Hessian matrix)**捕捉的就是二阶信息(曲率)。
对于标量函数 f(x_1, \ldots, x_n),海森矩阵是所有二阶偏导数组成的 n \times n 矩阵:
H = \begin{bmatrix} \frac{\partial^2 f}{\partial x_1^2} & \frac{\partial^2 f}{\partial x_1 \partial x_2} & \cdots \\ \frac{\partial^2 f}{\partial x_2 \partial x_1} & \frac{\partial^2 f}{\partial x_2^2} & \cdots \\ \vdots & \vdots & \ddots \end{bmatrix}
H = \begin{bmatrix} 6x & 4y \\ 4y & 4x - 6y \end{bmatrix}
对角元(6x 和 4x - 6y)告诉你:沿 x 方向走时,x 方向的斜率怎么变;对 y 也类似。
非对角元(4y)则告诉你:沿一个方向走时,另一个方向的斜率怎么变。
**克莱罗定理(Clairaut's theorem)**保证:只要函数的二阶导数连续,混合偏导数就相等:\frac{\partial^2 f}{\partial x \partial y} = \frac{\partial^2 f}{\partial y \partial x}。
这意味着海森矩阵是对称的,而(正如我们在矩阵那一章看到的)对称矩阵必定有实特征值和正交的特征向量。
海森矩阵告诉我们函数在某个临界点(梯度为零处)附近的形状:
**多元链式法则(multivariate chain rule)**把链式法则推广到了多变量函数。如果 z = f(x, y),其中 x = g(t)、y = h(t),那么:
从 t 到 z 的每一条路径都贡献一项:沿该路径的偏导数,乘以中间变量对 t 的导数。
例如,若 z = x^2 y + 3x - y^2,x = \cos(t),y = \sin(t):
除了手动求导,我们还有三种方式:
jax.grad 计算 f(x, y) = x^2 y + 3x - 2y 在点 (1, 2) 处的梯度。因为 f 接受的是向量输入,所以要配合 argnums 使用 jax.grad。import jax import jax.numpy as jnp def f(x, y): return x**2 * y + 3*x - 2*y df_dx = jax.grad(f, argnums=0) df_dy = jax.grad(f, argnums=1) x, y = 1.0, 2.0 print(f"∂f/∂x = {df_dx(x, y):.4f} (expected: {2*x*y + 3:.4f})") print(f"∂f/∂y = {df_dy(x, y):.4f} (expected: {x**2 - 2:.4f})")
jax.jacobian 计算一个向量值函数的雅可比,并与手算结果对比。import jax import jax.numpy as jnp def F(x): return jnp.array([x[0]**2 + x[1], x[0] * x[1]**2]) J = jax.jacobian(F) x = jnp.array([1.0, 2.0]) print(f"Jacobian at (1,2):\n{J(x)}") # 预期结果: [[2*x[0], 1], [x[1]**2, 2*x[0]*x[1]]] = [[2, 1], [4, 4]]
jax.hessian 计算 f(x, y) = x^3 + 2xy^2 - y^3 的海森矩阵,并验证它是对称的。import jax import jax.numpy as jnp def f(xy): x, y = xy[0], xy[1] return x**3 + 2*x*y**2 - y**3 H = jax.hessian(f) point = jnp.array([1.0, 2.0]) hess = H(point) print(f"Hessian:\n{hess}") print(f"Symmetric: {jnp.allclose(hess, hess.T)}") # 预期结果: [[6x, 4y], [4y, 4x-6y]] = [[6, 8], [8, -8]]
Var 都记录自己的值,以及如何沿链式法则把梯度向后传。class Var: def __init__(self, val, children=(), backward_fn=None): self.val = val self.grad = 0.0 self.children = children self.backward_fn = backward_fn def __add__(self, other): out = Var(self.val + other.val, children=(self, other)) def _backward(): self.grad += out.grad # d(a+b)/da = 1 other.grad += out.grad # d(a+b)/db = 1 out.backward_fn = _backward return out def __mul__(self, other): out = Var(self.val * other.val, children=(self, other)) def _backward(): self.grad += other.val * out.grad # d(a*b)/da = b other.grad += self.val * out.grad # d(a*b)/db = a out.backward_fn = _backward return out def backward(self): # 拓扑排序后再传播梯度 # 我们会在数据结构与算法那一章详细讲 order, visited = [], set() def topo(v): if v not in visited: visited.add(v) for c in v.children: topo(c) order.append(v) topo(self) self.grad = 1.0 for v in reversed(order): if v.backward_fn: v.backward_fn() # f(x, y) = x*x*y + x 在 (3, 2) 处 x = Var(3.0) y = Var(2.0) f = x * x * y + x # = 3*3*2 + 3 = 21 f.backward() print(f"f = {f.val}") # 21.0 print(f"df/dx = {x.grad}") # 2*x*y + 1 = 13.0 print(f"df/dy = {y.grad}") # x*x = 9.0