接手过几个深度学习项目之后,我越来越觉得"会用PyTorch"和"理解深度学习框架"是两码事。反向传播为什么能自动算梯度?计算图和变量之间的关系到底是什么?为了把这些疑问彻底弄清楚,我决定照着DeZero的思路从零手写一个可用的深度学习框架。整个体系最底层的两块积木就是变量(Variable)和函数(Function),这两样东西加起来不到50行代码,却几乎已经决定了后面所有自动微分能力的地基结构。
如果你也想搞明白框架内部到底发生了什么,推荐你跟着这篇文章一步步动手敲一遍。它不要求你有多深的Python功底,也不需要你懂复杂的数学推导,只需要你熟悉最基础的NumPy操作。我会把为什么要这样设计、为什么这个函数要放在类里而不是直接用一个Python函数代替、以及中间踩过的坑都摊开讲清楚。
1. 为什么要从零写框架:PyTorch最值得追问的一句话
1.1 反向传播不是魔法,而是一张记录计算过程的账本
大多数人第一次接触深度学习框架时做的事情都一样:model(x)传进一个张量,调loss.backward(),然后梯度就变出来了。这个过程太顺滑了,顺滑到很少有人停下来问一句:backward()是怎么知道该往哪个方向传、传多少数值的?
答案是它记住了一笔账。反向传播之所以能工作,是因为在前向计算的时候,框架已经悄悄把"输入是谁、输出是谁、经过了哪个函数"全部记录下来。这些信息组成了计算图(computational graph),梯度就是沿着这张图反向流动的。
而你要从零实现一个框架,第一步要做的并不是直接去写反向传播,而是先把"账本"上最基本的元素准备好:一个能装数据的变量,一个能做计算的函数。只要这两个东西设计得好,后面往上拼任何神经网络层都会非常顺手。
1.2 DeZero的哲学:最小可运行,然后逐步长出肌肉
DeZero不是那种一上来就甩给你几百个API的框架。它的设计思路非常克制:每一章只增加一个概念,每个概念都能跑,跑完都能验证。这种方式对学习者极其友好,因为每一步你都能看见自己改了什么、影响是什么。
很多自学资料的问题在于"起步太高",一上来就让你写自动微分模块、写卷积层,结果概念之间互相依赖,最后什么都没吃透。DeZero的方式是先把一个盒子做出来——这个盒子能装数据;再做一个工具——这个工具能对盒子里的数据做数学变换。随着你慢慢往盒子里加属性、往工具里加回调逻辑,一个完整的深度学习框架就自己长出来了。
1.3 本文的代码环境和准备
在动手之前,先把环境准备好:
- Python 3.8以上即可,不需要任何深度学习库
- 只用NumPy处理数值运算,它是唯一的外部依赖
- 一个趁手的编辑器,建议打开行号显示,因为你后面会反复回看错误堆栈
import numpy as np本文所有代码都基于这个文件。后面的小节中,我会逐段解释代码背后的设计考虑,而不是简单堆一段能跑的东西就完事。
2. Variable类:先让数据在计算图里有个正式户籍
2.1 最朴素的Variable实现
先看第一版代码,总共只有三行逻辑,但它是整个框架的起点。
class Variable: def __init__(self, data): self.data = data你没看错,第一版就是这么多。构造时接收一个NumPy数组,把它存成实例属性self.data。
这个类看起来不起眼,却解决了深度学习框架里的一个基础问题:如何区分"框架里的数据"和"框架外的Python浮点数、列表"。
在之后实现函数类的时候,你就能感受到这样一个统一包装带来的好处——所有函数的输入输出都约定为一个Variable实例,内部数据统一用variable.data来拿。这个约定让函数之间可以无缝串联,不需要关心数据类型的分支判断。
2.2 data的三条使用纪律
使用Variable时有几条纪律值得从现在就开始养成,等后续写到反向传播会受益非常多。
第一条:self.data必须是NumPy数组。不要塞Python原生浮点数,更不要塞列表。NumPy数组所有的广播、切片、数学函数都齐全,之后的梯度计算也大量依赖ndarray的特性。
第二条:不要在外部直接随意修改variable.data。现在看起来没什么风险,但一旦我们有了自动微分,data会和其他属性(比如梯度grad、创建者creator)产生联动。你直接塞一个新数组进去,就可能绕过了框架需要记录的"属性同步关系"。
第三条:构造变量后,养成用variable.data访问内部数组的习惯,而不是自己私下再保存一份引用。因为在后面的章节里,函数类会频繁地从输入变量中取数据、构造新变量、再返回出去,统一的访问方式能减少很多怪问题。
2.3 为反向传播预留的坑位:grad与creator
现阶段我们不需要在Variable里加任何额外属性。但为了不返工,如果在阅读后面章节时你需要扩展,通常会在__init__中补上两个字段:
class Variable: def __init__(self, data): self.data = data self.grad = None # 反向传播时累积的梯度 self.creator = None # 生成这个变量的函数creator是计算图里"谁创建了我"的记录,它让数据可以反向追踪到生成它的函数;grad则是该变量对应的梯度,在反向传播时会被填充。现在先知道这两个坑位,不需要急着实现。
3. Function类:把一次数学计算变成可组合的积木
3.1 __call__与forward之间的分工逻辑
有了能装数据的Variable,下一步就需要一个能加工它的Function类。这里有一个非常经典的设计决定:为什么不用一个普通的Python函数,而要用一个类,还要搞得这么绕?
直接写一个def square(x): return x ** 2不就好了吗?
关键在于:只做计算是不够的,你还要为未来的自动微分做铺垫。反向传播需要知道三件事:输入是什么、输出是什么、这个函数如何把输出端的梯度传给输入端。
所以Function类需要有一个固定的"入口流程"和"实际计算"的分离。入口流程负责统一接住Variable、提取数据、调用实际计算、再把结果打包成新的Variable返回。实际计算则交给你重写。
代码如下:
class Function: def __call__(self, input_variable): x = input_variable.data y = self.forward(x) output_variable = Variable(y) return output_variable def forward(self, x): raise NotImplementedError()__call__方法是Python里非常方便的一种设计,它让你的实例可以像函数一样被调用:f = Function(); result = f(x)。所有自定义的计算类只需继承Function类并改写forward,后续反向传播需要的钩子都可以集中放在__call__里,而不用每个子类都重复写一遍"取数据、打包、返回"这套模板逻辑。
用生活类比来说:__call__像是门店前台,它负责接待客人、登记信息;forward像是后厨,只需要专心把菜做好。前台和后厨分开后,任何新菜品都只需要调整后厨,前台服务流程完全不变。
3.2 实现Square:第一个真正的Function子类
有了父类模板,定义一个平方运算只需要三行:
class Square(Function): def forward(self, x): return x ** 2然后就可以串联使用了:
x = Variable(np.array(3.0)) f = Square() y = f(x) print(type(y)) # <class '__main__.Variable'> print(y.data) # 9.0这里有个关键细节:y的类型仍然是Variable,而且是Square这个函数调用后产生的新变量。这意味着,你可以继续把y传给下一个函数,形成一个链条。比如:
g = Square() z = g(y) # 等价于 (3.0^2)^2 = 81.0数据不断在变量和函数之间交替传递,计算图的概念已经在这两行代码里悄然成型。
3.3 为什么这种写法能复用到任意算子
如果你现在想实现三角函数、指数函数、加法,流程完全一致——继承Function,在forward里写清NumPy运算:
class Exp(Function): def forward(self, x): return np.exp(x)因为__call__负责了所有"包装逻辑",每个新算子只要关心单变量的数学变换即可。后续章节中实现神经网络层、损失函数的时候,你也只需要照着这个模式写forward,框架的骨架完全不需要再改动。
我在实践中发现一个非常重要的习惯:任何自定义算子都要保持forward返回的是NumPy数组,不要返回Variable,更不要返回多个值。如果你破坏了这条约定,__call__里Variable(y)这层的包装就会失效,报错会变得极其隐蔽。
4. 数值微分先行:在自动求导之前,我们先需要一把尺子
4.1 中心差分比单侧差分好在哪里
现在变量和函数已经能跑,接下来要做的事情非常重要——哪怕还没有写任何反向传播的代码,也要先准备一个"标准答案",用来验证以后自动求导算得对不对。这个标准答案就是数值微分。
所谓数值微分,就是用"极限定义"近似求导:
df(x)/dx ≈ (f(x + h) - f(x - h)) / (2h)这叫做中心差分(central difference)。你可能在教科书上看过单侧差分(f(x + h) - f(x)) / h,它同样能近似导数,但精度不如中心差分。原因是中心差分的误差项是O(h^2)量级,而单侧差分只有O(h)。用大白话说,同样的步长,中心差分的结果至少比单侧差分更接近真实值。
代码实现:
def numerical_diff(f, x, eps=1e-4): x0 = Variable(x.data - eps) x1 = Variable(x.data + eps) y0 = f(x0) y1 = f(x1) return (y1.data - y0.data) / (2 * eps)注意这里我把f当成了"可调用对象",而不是直接传NumPy数组。这正是第3.1节里那个设计的优势:数值微分不关心你是哪种函数,它只要求入参是Variable、返回的是Variable(或被调用后能取到.data),这个通用性让后续所有算子都能共享同一个测试函数。
4.2 用sin(x)验证Variable+Function的数据通路
先实现一个Sin函数类:
class Sin(Function): def forward(self, x): return np.sin(x)注意:np.sin接收的是NumPy数组,返回的也是NumPy数组。我们的forward里不应该主动去动变量本身的包装层,包装的事交给__call__。
然后定义一个"函数对象"并跑数值微分:
x = Variable(np.array(0.0)) f = Sin() result = numerical_diff(f, x) print(result) # 约 0.999999999995...,非常接近 cos(0) = 14.3 精度观察:h到底取多少合适
我在实际测试中发现,eps取1e-4通常能给出足够好的结果,取1e-6可能因为浮点数舍入误差而出现微妙抖动。原因在于中心差分公式里有个2 * eps的分母,如果eps太小,f(x+eps)和f(x-eps)减去之后的有效数字会大量损耗,导致误差反而变大。
经验法则:先用1e-4跑一遍,如果和真实导数差异超过1e-4级别,再逐步缩小到1e-5、1e-6。不要一上来就迷信极小步长,那是新手最容易踩的坑。
在这一步,你会体会到框架的价值:你先实现了干净的"变量-函数"抽象,然后测试代码可以写得非常简洁。后续无论新增多少算子,数值微分测试都只有几行。
5. 跑通第一个完整实验:sin(x)的导数与梯形积分
5.1 实现Sin Function并验证
上面的验证已经做了一半,现在我们把整个链路串成一个完整的脚本,验证Sin的导数:
import numpy as np class Variable: def __init__(self, data): self.data = data class Function: def __call__(self, input_variable): x = input_variable.data y = self.forward(x) output_variable = Variable(y) return output_variable def forward(self, x): raise NotImplementedError() class Sin(Function): def forward(self, x): return np.sin(x) def numerical_diff(f, x, eps=1e-4): x0 = Variable(x.data - eps) x1 = Variable(x.data + eps) y0 = f(x0) y1 = f(x1) return (y1.data - y0.data) / (2 * eps) x = Variable(np.array(np.pi / 4)) f = Sin() deriv = numerical_diff(f, x) print(deriv, np.cos(np.pi / 4)) # 两者应非常接近输出结果大概在0.70710678左右,和cos(pi/4) = 0.70710678对比,差值在1e-6量级。到这一步,我们亲手搭建的变量和函数已经具备真实计算能力。
5.2 梯形法求面积:顺手复习框架的数据流
除了求导,还可以直接用这套框架做一点微积分练习:用梯形法计算sin(x)在0到π之间的积分,理论值是2。
思路很简单:把积分区间等分成N份,每一份用小梯形逼近面积,然后求和。这次我们把"函数"传入累加循环:
def trapezoid_integrate(f, a=0.0, b=np.pi, n=1000): h = (b - a) / n total = 0.0 for i in range(n): x0 = Variable(np.array(a + i * h)) x1 = Variable(np.array(a + (i + 1) * h)) y0 = f(x0).data y1 = f(x1).data total += (y0 + y1) * h / 2.0 return total f = Sin() print(trapezoid_integrate(f)) # 接近 2.0这里同样能看到抽象一致性:f(x0)返回的是Variable,必须通过.data取出数值参与累加。如果你在写循环时漏掉.data,Python就会因为拿一个对象做加法而报出莫名其妙的类型错误。
5.3 从数值微分到反向传播的下一步跳跃
这个例子让我们彻底理解了框架的数据流。为什么说它是"第一步跳跃"?因为接下来你完全可以把numerical_diff做成一个自动化的测试工具,用来校验自己写的反向传播公式对不对。
通常的做法是:给每个自定义的Function实现一个backward方法,然后拿数值微分的结果和backward的结果对比,误差控制在1e-5以内就算通过。这套"数值微分校验法"会一直贯穿整个框架开发流程,无论后面你写多少次新层,都是这么验证的。
当你理解了变量和函数之间的关系后,反向传播就不是什么玄学:你只需要在__call__里把"输出的函数是谁"记录下来(即给output_variable.creator = self),然后在反向传播时让每个函数知道怎么把梯度回传给输入。
6. 变量与函数联动时常见的三个坑
6.1 原地修改data导致的诡异结果
新手最容易犯的一个错是:拿到Variable之后,直接原地操作它的data属性。
x = Variable(np.array(3.0)) x.data -= 1 # 原地修改 f = Square() y = f(x) print(y.data) # 变成4.0:因为x已经变成2.0了看起来符合预期,但一旦后续引入反向传播,这个问题就会"阴间"起来:如果这个变量曾经参与过一次前向计算,它的creator和grad信息可能已经记录在案。你原地修改数据,会让已有的计算图失效,而报错位置却往往远离修改点。我的建议是:永远用新的Variable承载新数据,不要原地改。
6.2 __call__里不要顺手覆盖入参
有的同学可能会在__call__里为了某种"内部整洁"去做这样的事:
def __call__(self, input_variable): input_variable.data = input_variable.data * 2 # 越权操作 x = input_variable.data y = self.forward(x) output_variable = Variable(y) return output_variable这会在不知不觉中改变外部变量本身的值。前面说过,同样的数据可能在计算图中被多个函数引用,你一改入参,所有引用它的节点都会一起变化,结果就是所有梯度校验全部失败,而且极难排查。记住:__call__只负责"读入参、调forward、包新变量",不要做任何多余的操作。
6.3 忘记转换类型导致的mysterious行为
数值微分的代码中,Variable(x.data - eps)这一步要求x.data是NumPy的数组或标量。如果你使用过程中把x.data换成了普通Python浮点数,有时候NumPy会以奇怪的方式广播(或者直接报错),让你花大量时间在无关的细节上。
我的建议是:在Variable.__init__里加一个极简的类型断言,即使只写一行注释也好:
class Variable: def __init__(self, data): if not isinstance(data, np.ndarray): raise TypeError('Variable.data must be a numpy.ndarray') self.data = data这个断言会让调试周期缩短很多,因为它能把问题暴露在源头,而不是等到数值微分或者积分循环里才爆出看不懂的错误。
从变量与函数这两个最基础的构建块出发,我已经完成了数值微分验证和积分求解两个实验。回头来看,整个系统的可组合性来自于两个设计决策:Variable统一承载数据,Function统一封装计算入口,并把__call__和forward分离。这种分工让以后每新增一个算子都变得非常机械,也为自动微分留下了天然的扩展点。
我个人在写这一步时的体会是:不要急着一步到位写上梯度相关的代码,先把变量和函数的数据流摸透、把数值差分测试跑通,后面的路会顺很多。当你亲手把这三个类拼到一起并看到输出与理论值吻合时,你对深度学习框架的内核就已经建立起第一层直觉了。下一步就可以引入creator属性,开始实现真正的反向传播。