ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

DeZero框架最小实现:函数连续调用与数值微分

DeZero框架最小实现:函数连续调用与数值微分 如果你也看过《深度学习入门2自制框架》一定对DeZero这个名字不陌生。这个项目用极少代码演示了一个深度学习框架最核心的骨架。今天这篇博文我们只围绕两个最基础也最关键的点展开函数连续调用与数值微分。看完之后你会发现所谓“深度学习框架”在最小实现里不过就是几个Python类。适合谁看想手工复现框架、准备AI面试、或者读原书卡在数值微分这一步的同学都可以从中拿到一套能直接运行的最小实现。我会把每一步的代码、计算过程和踩坑记录都摆出来尽量还原真实从零构建的过程。1. 项目动机与整体设计1.1 从“一个函数不够用”说起深度学习中几乎不存在“输入到输出只经过一次变换”的模型。以最简单的多层感知机为例输入要经过线性变换、激活函数再来一次线性变换最后经过softmax。这一整串操作本质上就是多个函数连续作用在一个变量上。但“连续调用”四个字落到代码层面远比口头描述麻烦如果每个函数都直接修改原始数组中间结果会全部丢光如果每个函数都只接收ndarray、返回ndarray以后想挂“梯度”“计算节点”这类附加信息时就要大规模改动现有代码。所以第一步不是写具体函数而是先搭一套“变量即对象、函数即对象”的最小骨架。1.2 设计目标最小但完整的一条可运行链路这次复现的目标非常小只有两个能力Variable一个包装numpy数组的类目前只负责存放数据Function所有计算节点的基类输入Variable输出Variablenumerical_gradient给定一个目标函数和一个输入点输出该输入点处的数值梯度。为了把“函数连续调用”和“数值微分”这两者打通我们暂时不做反向传播不设计计算图也不考虑算子注册。为什么这么克制因为我见过太多人一开始就想把自动微分、动态图、GPU支持全塞进去结果连一个x**2的梯度都调试不出来。先把最小链路跑通之后再往上添加能力反而是最快的路径。整个项目的核心链路可以概括为Variable - Function - Variable - Function - ... - Variable最后对末端Variable求数值梯度。1.3 为什么变量对象是“箱子”而不是“数据本身”生活化类比函数就是工厂流水线上的工位每个工位只做一件事——接上半成品做完传给下一个工位。Variable对象就是装着半成品的箱子箱子本身不参与运算但它保证了每个工位都按同样的“收货标准”工作。如果直接用numpy数组就相当于让半成品裸奔每个工位都得自己处理包装、运输和异常。DeZero不管设计多简单都坚持让变量成为对象。这不只是风格问题更是一个稳定的扩展边界以后只要在Variable上增加grad、creator等属性就能平滑过渡到反向传播。2. 基础架构构建先让Variable和Function能跑起来2.1 Variable类数据转换必须做在源头先写第一个类。这里我建议直接在初始化时把数据转成np.float64的ndarray而不是原样保存。这个选择的理由在写数值微分时会体现得特别明显如果用户传入一个Python标量或者整数数组后面做扰动加法时会因为类型问题直接出错或者出现“加了h但取整成原值”的诡异现象。import numpy as np class Variable: def __init__(self, data): self.data np.asarray(data, dtypenp.float64)np.asarray和np.array的区别在于如果传入的已经是对应类型的ndarraynp.asarray不会拷贝原数组。这个特性在数值微分中很重要因为我们希望通过修改x.data内部的值来完成“扰动”而不是每次重新创建一个新数组。当然如果你的代码中不允许原数组被大规模修改也可以显式用np.array(data, dtypenp.float64)拷贝一份但那样后续内存开销会变大。这里有一个容易忽略的细节dtypenp.float64保证了data是可浮点运算的数组。如果用户传入np.array([1, 2, 3])默认是int64后面x.data[idx] tmp_val h时浮点数会被截断成整数梯度必然出错。把类型转换放在Variable构造器里相当于所有问题在最源头就被拦截掉了。2.2 Function类__call__和forward为什么要分开接下来是Function基类。很多初学者不理解为什么不能只写一个forward非要再加一层__call__。我们看代码class Function: def __call__(self, input): x input.data y self.forward(x) output Variable(y) return output def forward(self, x): raise NotImplementedError()__call__负责三件事从Variable里取出数据、调用真正的计算逻辑forward、把计算结果重新包装成Variable。这意味着使用方永远面对的是Variable对象而不是直接操作numpy数组。先不说这会让代码更整洁更重要的是为将来留了后路——如果后面要扩展反向传播我们只需要在__call__里补上“保存输入到input属性”和“把输出和输入关联起来”这两步接口本身完全不用改。forward里只写具体计算比如平方、指数、sin。这样做还有一个好处以后做自动微分时每个算子只需要额外再写一个backward方法__call__就能在完成前向的同时把反向所需的信息准备好。可以说这个接口划分是整个框架后续所有扩展的地基。2.3 先实现三个基础函数Square、Exp、Sin有了基类我们立刻实现几个具体算子。为了贴近深度学习场景我选了平方、指数和sin这三个函数足够支撑连续调用和梯度验证。class Square(Function): def forward(self, x): return x ** 2 class Exp(Function): def forward(self, x): return np.exp(x) class Sin(Function): def forward(self, x): return np.sin(x)为了方便连续调用再包一层函数式接口def square(x): return Square()(x) def exp(x): return Exp()(x) def sin(x): return Sin()(x)这里每次调用都会创建新的Function实例但不用担心性能——算子本身是无状态的。函数式接口带来的好处是调用方不需要关心类实例化细节写出来的表达式更接近数学公式square(x)、exp(x)、sin(x)。这一步没有魔法本质就是在语法层面做了一层简化。简单测一下x Variable(np.array(2.0)) y square(x) print(y.data) # 4.0到这里我们已经有了Variable和Function的最小闭环。目前这个框架只能做单一变换还看不出“连续调用”的威力。下一步我们把它们串起来。3. 函数连续调用的实现把多个函数串成一条计算链3.1 顺序串联让数据流经多个工位函数连续调用最直观的写法是嵌套比如z square(exp(x))。但更贴近真实框架工作方式的写法是分步串联x Variable(np.array(0.5)) a exp(x) # a e^0.5 b square(a) # b (e^0.5)^2 e^1 y sin(b) # y sin(e^1)每一步都返回一个新的Variable对象。这样做给调试带来的好处是立竿见影的你可以在任何一步打印a.data、b.data判断计算是否符合预期。如果直接写一个大嵌套表达式一旦结果不对你只能从头到尾重新心算一遍很难定位是哪一层出错。从设计上看“连续调用”能力是被两个约定天然支持的每个Function接收Variable、返回Variable每个具体函数在forward里只关心从x.data取出numpy数组做运算。这两个约定保证了任意两个函数都能拼接不会因为入参类型不匹配而中断。3.2 为什么中间变量值得保留有同学会问数值微分计算梯度时并不需要中间变量为什么还要这么麻烦地分步保存这里要分两层来看。第一层从代码可维护性来看保留中间变量是“可观测性”。比如调试时你能查看每一层的输出是否在合理区间。如果第3层出现了NaN你可以直接定位到具体是哪一层引入了非法值而不是靠猜。第二层从框架演进来看后续实现反向传播时每个节点都需要知道“我的输入是什么”“我的输出是什么”“我的局部梯度是什么”。如果你在连续调用时把所有中间结果都丢掉将来自动微分根本无从下手。可能有人觉得分步写代码不如嵌套写法优雅但你要记住我们不是在写一次性脚本而是在搭框架。框架的第一原则是给未来留出空间而不是让当前这一行代码最简洁。所以哪怕DeZero后续可以用嵌套写法它的内部依然会创建指向中间结果的Variable对象。3.3 不固定函数数量时用循环维护调用链有时候调用链的长度不是写代码时固定的而是由配置或数据决定的。比如函数列表来自一个列表变量这时用循环串联更通用x Variable(np.array(0.5)) funcs [Exp(), Square(), Sin()] y x for f in funcs: y f(y) print(y.data)这段代码和前面分步串联完全等价但逻辑上更适合动态场景。你只要保证列表中每个元素都是Function实例循环就能不断更新y。这里有一个细节点y被反复赋值前一个Variable对象并没有被销毁只是不再被y引用所以中间结果仍然保留在内存中。如果某个时段内存压力大可以手动del不需要的中间变量但在小型Demo里完全没必要。实际上这种“用循环串接算子”的模式就是很多深度学习框架最底层迭代的雏形。框架内部会维护一个更结构化的计算图但外部表现同样是函数连续调用。理解了这一步后面读框架源码时会轻松很多。3.4 连续调用中容易踩的三个坑第一个坑在forward里修改输入数组本身。比如写x * 2会直接改写input.data导致同一个Variable被多次调用时结果不可复现。正确的做法是让forward返回一个新数组比如return x ** 2。第二个坑复用同一个Function实例时如果实例内部记录了最近一次调用的输入输出就可能出现数据串线。现阶段我们写的Function都是无状态的所以可以安全复用但如果你在上面加“缓存”之类的优化就要非常小心。第三个坑忘掉输入类型检查。__call__里直接取input.data如果用户传了ndarray进来会直接AttributeError。更隐蔽的情况是传了Python列表虽然不会报错但.data属性完全不存在。我的建议是明确约定“这个框架的运算单位是Variable”并在外部文档里写清楚而不是在代码里加一堆防御性判断。加太多if会把主干逻辑淹没。4. 数值微分的实现用“小小的扰动”逼近真实梯度4.1 中心差分比单侧差分更聪明数值微分是最朴素、也最可靠的梯度近似方案。它不需要你手推导函数只需要基于导数定义给输入一个微小扰动观察输出变化。但直接使用单侧差分并不够好我们一般用中心差分。中心差分公式[ f(x) \approx \frac{f(xh) - f(x-h)}{2h} ]相比单侧差分(f(xh)-f(x))/h中心差分从左右两侧各取一个点几何上相当于把割线的斜率取在x正中央。误差分析上单侧差分截断误差是O(h)中心差分是O(h^2)。同样是h1e-4中心差分的截断误差会小很多。为什么不用解析求导因为手工推导复合函数偏导容易出错而且有些函数根本没有闭合表达式。数值微分虽然慢但它不依赖推导过程适合作为基准答案来验证其他微分实现是否正确。DeZero在早期引入自动微分之前就是先用数值微分验证梯度结果的。4.2 通用数值梯度函数逐元素扰动我们要写一个通用的numerical_gradient(f, x)其中f是接收Variable并返回Variable的标量函数x是待求梯度点。对于多维输入需要对每个元素分别扰动一次同时把其他元素保持为原始值。用np.nditer可以优雅地遍历任意维数组的每个位置def numerical_gradient(f, x): h 1e-4 grad np.zeros_like(x.data, dtypenp.float64) it np.nditer(x.data, flags[multi_index]) while not it.finished: idx it.multi_index tmp_val x.data[idx] x.data[idx] float(tmp_val) h fxh1 f(x).data.item() x.data[idx] float(tmp_val) - h fxh2 f(x).data.item() grad[idx] (fxh1 - fxh2) / (2 * h) x.data[idx] tmp_val it.iternext() return grad这段代码有几个关键点。第一np.nditer配合multi_index可以同时处理一维、二维甚至更高维数组不需要写多层for循环。第二每次扰动前都要保存原始值tmp_val扰动完立刻恢复否则会影响下一个位置的梯度。第三f(x).data.item()要求目标函数输出的是标量。如果f(x).data是多元素数组.item()会直接抛错这种设计反而是一种保护数值梯度通常只用于标量损失函数。这里我特意把grad也用dtypenp.float64初始化避免整数数组做除法后自动丢失小数位。4.3 作用于连续调用链的数值微分我们已经有连续调用能力现在把数值微分套上去。考虑一个简单复合函数y square(x)理论导数是2x。x Variable(np.array(3.0)) grad numerical_gradient(square, x) print(grad) # 约 6.0这里square是一个函数式接口接收Variable并返回Variable正好满足numerical_gradient的要求。在x3处数值梯度会非常接近6说明基本链路是通的。但更复杂的连续调用链也同样适用。例如def composite(x): t exp(x) s square(t) y sin(s) return y x Variable(np.array(0.5)) grad numerical_gradient(composite, x) print(grad)这段代码会先顺着调用链算出y sin((e^x)^2)然后在x0.5处用中心差分得到梯度。注意数值微分不关心这个函数是几步完成的它只关心“当x.data发生微小变化时最终输出变化多少”。因此哪怕你的嵌套层级再深只要f(x)能正常算出标量输出数值微分就能工作。5. 连续调用与数值微分合流一个完整实战验证5.1 单变量复合函数手动解析对比理论推导永远是最好的验证方式。设[ y \sin\left((e^x)^2\right) \sin(e^{2x}) ]对x求导[ \frac{dy}{dx} \cos(e^{2x}) \cdot 2e^{2x} ]我们用数值微分和解析解同时计算并观察两者是否一致x Variable(np.array(0.5)) grad_numeric numerical_gradient(composite, x) c np.exp(2 * 0.5) grad_manual np.cos(c) * 2 * c print(数值梯度:, grad_numeric) print(解析梯度:, grad_manual)在我的环境里输出大约是数值梯度: [-0.56838] 解析梯度: -0.56839误差在1e-4量级完全符合中心差分的预期。这个验证过程建议每个人都亲手跑一遍它不仅能确认框架没有写错还能让你直观体会“数值方法逼近解析解”是怎么回事。5.2 多维输入的梯度检查每个位置深度学习里的参数大多是向量或矩阵所以我们还要验证多维输入场景。以最简单的“平方和”为例def sphere(x): return (x ** 2).sum() x Variable(np.array([1.0, 2.0, 3.0])) grad numerical_gradient(sphere, x) print(grad)输出应该是[2. 4. 6.]。因为sum(x^2)对每个分量的偏导数是2*x_i。如果输出与预期不符最常见的原因是输入数组在Variable构造时被转成了整数类型导致扰动丢失。所以前面的dtypenp.float64非常关键。np.nditer在这里会自动遍历向量中的三个元素并为每个元素计算中心差分。要注意每计算一个元素的梯度都要把数组恢复成原始值再去扰动下一个元素。不然第一次扰动会污染后续结果。5.3 再复杂一点逐元素平方后求和我们可以把连续调用和多维输入结合起来模拟一层最简单的网络。假设输入x先逐元素平方再求和整个链路虽然短但已经有了“逐元素变换 汇总”的结构这很像一个没有线性层和偏置的全连接层雏形。def layer_like(x): h square(x) y h.sum() return y x Variable(np.array([1.0, 2.0, 3.0])) grad numerical_gradient(layer_like, x) print(grad) # [2. 4. 6.]理论上layer_like的梯度就是2*x和sphere相同因为平方操作没有额外的线性耦合。但你可以在这个基础上随意替换函数改成h exp(x)再求和梯度就会完全不一样。这种“函数连续调用数值微分”的组合已经足够你验证很多自己构造的计算链路了。5.4 一个提醒调试这类代码的观察顺序实战中我总结出一个小经验梯度结果不对时不要先怀疑数值微分函数而是先检查目标函数本身的计算过程。你可以打印每一步中间变量确认前向输出是否符合数学预期。如果前向输出都不对数值微分一定不可能对。其次检查输入x.data的类型和形状确保扰动操作真的生效。最后再看numerical_gradient中是否每次都正确恢复了原始值。按照这个顺序排查绝大多数问题能在三分钟内定位。6. 常见问题与排查技巧实录6.1 h为什么取1e-4而不是越小越好很多人直觉认为h越小越精确但浮点数计算会打破这个直觉。当h取到1e-8甚至1e-10时xh和x-h在浮点数表示里可能非常接近导致中心差分的分子f(xh)-f(x-h)丢失大量有效位舍入误差急剧上升。而h取太大时截断误差又会变大。实践经验里1e-4是个平衡点中心差分截断误差O(h^2)大约是1e-8量级舍入误差还远没有到失控的程度。如果你需要更精细的结果可以再用更高精度的数值微分公式但在这个阶段没必要。6.2 梯度算出来全零或明显偏大先查什么全零通常不是巧合而是输入被“整数化”了。比如x.data是np.array([1,2,3])当你给x.data[0]加上1e-4后numpy会按int64规则把结果截断成1前后没有任何改变梯度自然就是0。解决方法很简单Variable构造时统一转成float64。梯度明显偏大则更可能是函数本身的问题。比如你在forward里对数组做了原地更新导致两次扰动调用时输入状态不一致。另一个常见原因是目标函数输出不是标量但numerical_gradient没有正确报错而是做了隐式广播。我建议在函数入口加一行断言或检查明确要求f(x).data是0维数组写法和报错信息都清晰。6.3 为什么这篇文章不急着加运算符重载你可能已经在想象y (x ** 2) 3这样的写法这确实更好看。但运算符重载是“糖”不是“核心”。过早加糖容易掩盖数据流的方向。比如__add__、__mul__背后仍然要创建Function对象、保留中间结果如果这些机制还没扎实直接加运算符只会让调试变成“黑洞”。我见过不少Demo表达式写得很爽一旦梯度出错根本不知道是运算符重载的坑还是数值微分的坑。先把显式函数调用调通再考虑减负会更稳。6.4 下一步向反向传播平滑过渡数值微分虽然可靠但性能上有天然劣势每个参数维度都要跑两次完整前向百万参数模型几乎不可用。所以DeZero在验证梯度正确之后下一步就会引入计算图和链式法则实现反向传播。目前这个项目留了一个很好的扩展点Function.__call__里只需要记录输入以及在Variable上增加creator属性就能把计算图逐步组织起来。到时候你会看到反向传播的梯度结果和数值微分对比误差在1e-8量级才算合格。我在实际写这段过程中最深的体会是数值微分不是“备用轮子”而是一台可信赖的基准仪器。哪怕以后实现了自动微分我也强烈建议保留这套数值梯度代码凡是遇到反向传播结果可疑就用它做随机梯度检验。这个习惯能帮你抓住很多隐藏bug而且越早养成越好。
返回列表