news 2026/9/23 14:17:20

5分钟吃透复合函数求导法则,附Python源码解析

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
5分钟吃透复合函数求导法则,附Python源码解析

5分钟吃透复合函数求导法则,附Python源码解析

报错一堆看不懂 StackTrace?别急,先深呼吸。很多刚接触自动微分或数值计算的朋友,看到满屏的 TracebackAssertionError 就头疼,觉得这是天书。其实,这背后往往不是代码逻辑写崩了,而是对底层数学原理的理解出现了断层。今天咱们不整虚的,直接通过一个 Python 实战项目,把【复合函数求导法则】拆碎了揉碎了讲给你听,配合详细的【源码解析】,让你彻底搞懂链式法则在代码里到底是怎么跑的。

项目目标

咱们这个项目很简单,就是手写一个极简版的自动微分引擎。市面上像 PyTorch 或 TensorFlow 这样的框架,底层全是 C++ 和 CUDA 写的,普通人根本摸不到核心逻辑。我们要做的,是用纯 Python 实现一个类 Tensor,支持加法、乘法和非线性函数(如 sin, exp)的前向计算和反向求导。

目标只有一个:当你执行 y = f(g(x)) 时,代码能自动算出 dy/dx,而且精度要和数学推导一致。这不仅仅是为了炫技,更是为了让你明白,那些高大上的深度学习框架,在反向传播阶段,到底是在遍历什么样的计算图。通过这个项目,你会对“计算图”、“梯度累积”、“叶子节点”这些概念有肌肉记忆般的理解,以后再遇到 gradNone 或者梯度爆炸的问题,你至少知道该去查哪个环节。

目录结构

为了保持代码的可复现性和工程化,我们采用标准的项目结构。虽然代码不多,但规范不能少,这是职场人的基本素养。

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 函数,它捕获了外部的 selfother。这就是 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 out

4. 反向传播引擎

有了上面的节点,我们需要一个引擎来触发整个反向传播过程。这就是 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, f"Gradient mismatch: {x.grad} vs {expected}"print(f"Success: 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 列表,看看遍历顺序对不对。

优化扩展

基础版能跑通后,我们可以考虑一些工程化的优化。

  1. 支持标量混合:实际使用中,经常有 Tensor + float 的情况。我们需要重载 __add____mul__,判断 other 是否是 Tensor,如果是标量,就不需要记录父节点,梯度传递时标量的梯度为 0。
  2. 内存管理:目前的实现中,每个 Tensor 对象都保存在计算图中,直到 backward 结束。在生产环境中,我们需要支持 zero_grad()retain_grad(),以便在反向传播后释放内存,或者保留中间节点的梯度用于调试。
  3. 数值稳定性:对于 exp 函数,如果输入很大,math.exp 会溢出。在生产级框架中,通常会使用对数空间(log-space)或者截断(clipping)来处理。虽然本项目追求简洁,但你在阅读 PyTorch 源码时,会发现它们对每一个算子都做了大量的边界条件检查。

小结

通过这个项目,我们从零搭建了一个支持复合函数求导的微型引擎。核心在于理解了【复合函数求导法则】在代码中的映射:前向计算存值,反向计算存梯度,拓扑排序保顺序

很多人觉得数学难,其实是因为没有把它具象化。代码就是数学最好的翻译。当你看着 self.grad += ... 这一行行代码,你就真正懂了链式法则。

这个项目虽然简单,但麻雀虽小五脏俱全。你可以在此基础上,尝试添加 ReLU 激活函数,或者构建一个两层的感知机。GitHub 上有大量的类似开源项目,比如 microgradtinygrad,推荐大家去 Star 并 Fork 下来跑一跑,对比一下我们的实现,你会发现工程细节上的巨大差异,这正是从“会做题”到“会做工程”的跨越。

在反向传播的过程中,你有没有遇到过梯度消失或者梯度爆炸的问题?或者对拓扑排序的递归实现有性能上的顾虑?还有什么不懂的?评论区留言挨个回。

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/23 14:17:09

步尚雪源码解析:3个环境坑让你少熬2夜

步尚雪源码解析:3个环境坑让你少熬2夜 配置环境就卡半天,这简直是每个刚接触步尚雪的新人噩梦。我当年为了跑通一个示例项目,把电脑重启了五次,差点把键盘敲烂。别笑,这真不是个例,很多人盯着报错日志发呆,其实问题就出在最基础的依赖加载逻辑上。今天咱们不整虚的,直接扒开步尚雪的源码解析,看看那些官方教程里…

作者头像 李华
网站建设 2026/9/23 14:17:08

手机号码归属地查询软件下载源码解析实战指南

手机号码归属地查询软件下载源码解析实战指南 看了一堆教程还是不会写项目?别慌,这不是你的错,是大部分教程只教你“怎么下”,不教你“怎么改”。 很多人以为 手机号码归属地查询软件下载 就是去某个官网点一下“下载”,或者在 PyPI 上 pip install…

作者头像 李华
网站建设 2026/9/23 14:16:58

5步搞定走遍美国视频下载避坑指南

5步搞定走遍美国视频下载避坑指南 配置环境就卡半天,是不是你也经历过这种崩溃时刻?明明照着教程敲代码,结果报错一片红,最后发现是库版本不匹配或者依赖冲突。别急,这篇避坑指南专门为你整理,基于我过去三年处理数百个爬虫项目的实战经验,帮你一次性搞定环境搭建。…

作者头像 李华
网站建设 2026/9/23 14:16:51

面试必问:3个关于亚洲精品国产免费精情侣的源码坑

面试必问:3个关于亚洲精品国产免费精情侣的源码坑 刚毕业那会儿,我盯着屏幕上的代码,感觉脑子像被浆糊糊住。看了一堆教程还是不会写项目,这是多少新人的噩梦?别慌,今天咱们不聊虚的,直接拆解【亚洲精品国产免费精情侣】这类复杂业务场景下的典型源码问题。这可是 面试必问…

作者头像 李华
网站建设 2026/9/23 14:16:37

吉他新手必看:低弦距的重要性与选购指南

1. 为什么低弦距对新手如此重要&#xff1f;作为一名教过上百名吉他初学者的老师&#xff0c;我见过太多人因为选错吉他而放弃。其中最致命的错误&#xff0c;就是忽视了弦距这个关键指标。你可能不知道&#xff0c;一把弦距合适的吉他&#xff0c;能让你的学习效率提升30%以上…

作者头像 李华