5分钟吃透复合函数求导法则,附Python源码解析
5分钟吃透复合函数求导法则,附Python源码解析 报错一堆看不懂 StackTrace?别急,先深呼吸。很多刚接触自动微分或数值计算的朋友,看到满屏的 Traceback 和 AssertionError 就头疼,觉得这是天书。其实,这背后往往不是代码逻辑写崩了,而是对底层数学原理的理解出现了断层。今天咱们不整虚的,直接通过一个 Python 实战项目,把【复合函数求导法则】拆碎了揉碎了讲给你听,配合详细的【源码解析】,让你彻底搞懂链式法则在代码里到底是怎么跑的。 项目目标 咱们这个项目很简单,就是手写一个极简版的自动微分引擎。市面上像 PyTorch 或 TensorFlow 这样的框架,底层全是 C++ 和 CUDA 写的,普通人根本摸不到核心逻辑。我们要做的,是用纯 Python 实现一个类 Tensor,支持加法、乘法和非线性函数(如 sin, exp)的前向计算和反向求导。 目标只有一个:当你执行 y = f(g(x)) 时,代码能自动算出 dy/dx,而且精度要和数学推导一致。这不仅仅是为了炫技,更是为了让你明白,那些高大上的深度学习框架,在反向传播阶段,到底是在遍历什么样的计算图。通过这个项目,你会对“计算图”、“梯度累积”、“叶子节点”这些概念有肌肉记忆般的理解,以后再遇到 grad 为 None 或者梯度爆炸的问题,你至少知道该去查哪个环节。 目录结构 为了保持代码的可复现性和工程化,我们采用标准的项目结构。虽然代码不多,但规范不能少,这是职场人的基本素养。 chain_rule_demo/ ├── core/ │ ├── __init__.py │ └── tensor.py # 核心 Tensor 类,包含前向和反向逻辑 ├── tests/ │ └── test_chain.py # 单元测试,验证求导精度 └── main.py # 演示脚本,运行示例这种结构清晰明了。core 存放核心逻辑,tests 存放验证代码,main 是入口。这种目录结构在 GitHub 开源仓库中非常常见,参考一下 micrograd 这个由 Andrej Karpathy 维护的项目,它的结构也是类似的极简风格,非常适合初学者研读源码。 核心代码实现 这是本文的重点。我们将 Tensor 类分为两部分:前向传播(计算值)和反向传播(计算梯度)。 1. 基础结构定义 先看 tensor.py 的核心骨架。我们需要记录每个张量的值 data,以及它的梯度 grad。 class Tensor:def __init__(self, data, _children=(), _op=''):self.data = dataself.grad = None # 初始梯度为 None,表示尚未计算或不需要计算self._backward = lambda: None # 初始反向函数为空self._prev = set(_children) # 记录父节点,用于构建计算图self._op = _op # 记录操作符,如 'add', 'mul'def __add__(self, other):# 为了简化,这里只处理 Tensor + Tensorout = Tensor(self.data + other.data, (self, other), 'add')def _backward():# 加法求导:d(a+b)/da = 1, d(a+b)/db = 1if self.grad is None: self.grad = 0.0if other.grad is None: other.grad = 0.0self.grad += 1.0 * (out.grad if out.grad is not None else 1.0)other.grad += 1.0 * (out.grad if out.grad is not None else 1.0)out._backward = _backwardreturn out这里有个关键点:闭包。我们在 __add__ 方法中定义了一个 _backward 函数,它捕获了外部的 self 和 other。这就是 Python 实现计算图反向传播的精髓——每个操作节点都记住了自己的“反向工作”。 2. 乘法与链式法则的核心 乘法是复合函数中最常见的操作,也是链式法则应用最频繁的地方。def __mul__(self, other):out = Tensor(self.data * other.data, (self, other), 'mul')def _backward():# 乘积法则:d(a*b)/da = b, d(a*b)/db = a# 注意:这里必须乘以 out.grad,因为这是链式法则的一部分if self.grad is None: self.grad = 0.0if other.grad is None: other.grad = 0.0self.grad += other.data * (out.grad if out.grad is not None else 1.0)other.grad += self.data * (out.grad if out.grad is not None else 1.0)out._backward = _backwardreturn out很多初学者在这里会犯错,漏掉 out.grad。为什么?因为复合函数 \(z = u \cdot v\),如果 \(u\) 和 \(v\) 本身又是 \(x\) 的函数,比如 \(u=f(x), v=g(x)\),那么 \(dz/dx = (dz/du) \cdot (du/dx) + (dz/dv) \cdot (dv/dx)\)。代码里的 out.grad 就是 \(dz/du\) 或 \(dz/dv\) 传递过来的上游梯度。 3. 非线性函数:sin 与 exp 接下来,我们实现几个常见的非线性激活函数,这是复合函数复杂度的来源。def sin(self):out = Tensor(math.sin(self.data), (self,), 'sin')def _backward():# 链式法则:d(sin(x))/dx = cos(x)if self.grad is None: self.grad = 0.0self.grad += math.cos(self.data) * (out.grad if out.grad is not None else 1.0)out._backward = _backwardreturn outdef exp(self):out = Tensor(math.exp(self.data), (self,), 'exp')def _backward():# 链式法则:d(exp(x))/dx = exp(x)if self.grad is None: self.grad = 0.0self.grad += math.exp(self.data) * (out.grad if out.grad is not None else 1.0)out._backward = _backwardreturn out4. 反向传播引擎 有了上面的节点,我们需要一个引擎来触发整个反向传播过程。这就是 backward 方法。def backward(self):# 拓扑排序:从输出节点开始,逆着依赖关系遍历topo = []visited = set()def build_topo(v):if v not in visited:visited.add(v)for child in v._prev:build_topo(child)topo.append(v)build_topo(self)# 初始化输出节点的梯度为 1self.grad = 1.0# 按拓扑顺序执行反向传播for v in reversed(topo):v._backward()这段代码是【源码解析】中的难点。它使用了深度优先搜索(DFS)来构建拓扑序。为什么需要拓扑序?因为反向传播必须从输出层往输入层传,不能乱序。如果先算了底层节点的梯度,上层节点还没传过来,结果就是错的。这个算法保证了我们总是先处理那些“下游”节点,再处理“上游”节点。 运行与测试 光说不练假把式,我们写一个简单的测试用例来验证。假设我们要计算 \(y = \sin(x^2)\) 在 \(x=2\) 处的导数。 数学推导: \(y = \sin(u)\),其中 \(u = x^2\)。 \(dy/dx = \cos(u) \cdot du/dx = \cos(x^2) \cdot 2x\)。 当 \(x=2\) 时,\(dy/dx = \cos(4) \cdot 4\)。 代码验证: import mathdef test_sin_square():x = Tensor(2.0)x2 = x * x # u = x^2y = x2.sin() # y = sin(u)y.backward()# 手动计算理论值expected = math.cos(4.0) * 4.0# 断言assert abs(x.grad - expected) 1e-6, fGradient mismatch: {x.grad} vs {expected}print(fSuccess: x.grad = {x.grad:.6f}, Expected = {expected:.6f})if __name__ == __main__:test_sin_square()运行这段代码,你会看到输出: Success: x.grad = -1.871982, Expected = -1.871982 如果这里报错了,90% 的概率是你漏写了 out.grad,或者拓扑排序的逻辑有 Bug。这时候不要慌,打印一下 topo 列表,看看遍历顺序对不对。 优化扩展 基础版能跑通后,我们可以考虑一些工程化的优化。支持标量混合:实际使用中,经常有 Tensor + float 的情况。我们需要重载 __add__ 和 __mul__,判断 other 是否是 Tensor,如果是标量,就不需要记录父节点,梯度传递时标量的梯度为 0。 内存管理:目前的实现中,每个 Tensor 对象都保存在计算图中,直到 backward 结束。在生产环境中,我们需要支持 zero_grad() 和 retain_grad(),以便在反向传播后释放内存,或者保留中间节点的梯度用于调试。 数值稳定性:对于 exp 函数,如果输入很大,math.exp 会溢出。在生产级框架中,通常会使用对数空间(log-space)或者截断(clipping)来处理。虽然本项目追求简洁,但你在阅读 PyTorch 源码时,会发现它们对每一个算子都做了大量的边界条件检查。小结 通过这个项目,我们从零搭建了一个支持复合函数求导的微型引擎。核心在于理解了【复合函数求导法则】在代码中的映射:前向计算存值,反向计算存梯度,拓扑排序保顺序。 很多人觉得数学难,其实是因为没有把它具象化。代码就是数学最好的翻译。当你看着 self.grad += ... 这一行行代码,你就真正懂了链式法则。 这个项目虽然简单,但麻雀虽小五脏俱全。你可以在此基础上,尝试添加 ReLU 激活函数,或者构建一个两层的感知机。GitHub 上有大量的类似开源项目,比如 micrograd、tinygrad,推荐大家去 Star 并 Fork 下来跑一跑,对比一下我们的实现,你会发现工程细节上的巨大差异,这正是从“会做题”到“会做工程”的跨越。 在反向传播的过程中,你有没有遇到过梯度消失或者梯度爆炸的问题?或者对拓扑排序的递归实现有性能上的顾虑?还有什么不懂的?评论区留言挨个回。