框架与自动微分
机器学习框架定义模型代码与硬件之间的可执行契约。这份契约涵盖张量值、算子语义、导数规则、设备放置、分布式布局和执行模式。一个程序即使在数学上成立,也可能因为其中一部分未定义、不受支持,或被另一个后端作出不同解释而失败。
本章沿着这份契约,从一个标量损失一路追踪到设备实际执行的工作。我们先解释自动微分,再说明反向模式在内存与正确性上的要求,最后介绍框架如何捕获、编译、分派和分布式执行得到的程序。下层的硬件限制见 第 62 章;编译器与内核流水线将在 第 64 章 继续展开。
获取导数的三种方法
有限差分通过扰动输入来估计导数。对在 处求值的标量函数 ,坐标 上的前向差分估计为
这里, 是输入, 是第 个分量为一的基向量, 是非零步长, 是估计值。 项表示截断误差, 项表示与尺度有关的浮点舍入和抵消误差。减小步长会降低一种误差,却可能放大另一种误差。若要算出每个坐标的导数,前向差分对每个输入都需要一次扰动,中心差分则需要两次。因此,它不适合为十亿参数的训练生成梯度,不过方向有限差分仍然很适合检查梯度。
符号微分把一个数学表达式变换成另一个表达式。它可以得到有用的闭式结果,但朴素展开可能重复共享子表达式,造成表达式膨胀。控制流、修改操作、外部库调用和大型张量程序,也需要比普通计算机代数表达式更丰富的程序表示。
自动微分(automatic differentiation) 采用另一条路径。它对实际执行的数值程序中每个原语组合局部导数规则,同时保留程序的中间结构 (Baydin et al. 2018)。具体实现可以记录即时磁带、追踪计算图、变换中间表示,或改写源代码。在理想算术中,它对选定的可微路径精确应用链式法则;在机器上,它以机器精度计算该导数,仍然会受到舍入、溢出和下溢影响。它不会对离散的分支选择本身求导。遇到不可微点时,框架必须选择一种约定、返回未定义值,或抛出错误。
这段历史早于神经网络库。Wengert 在 1964 年描述了基本运算求值的线性列表 (Wengert 1964)。Linnainmaa 为计算机程序提出了反向累积,后来这一方法又以反向传播之名在神经网络中普及 (Linnainmaa 1976; Rumelhart et al. 1986)。现代框架把这套思想从标量算术推广到了张量算子。
线性化是核心接口
考虑一个可微函数
这里, 是输入维度, 是输出维度, 是函数的求值点。雅可比矩阵 包含所有局部敏感度,其中第 项等于 。框架很少需要把整块矩阵实体化,它们需要的是这个矩阵作用在向量上的结果。
前向模式计算雅可比向量积(JVP):
其中, 是输入切向量。若取基切向量 ,结果会返回雅可比矩阵的一列;任意 则给出一个方向导数。前向模式把切向量与每个中间值一起向前传播,不需要保留反向磁带。
反向模式计算向量雅可比积(VJP):
其中, 是输出余切向量, 表示转置。若取基余切向量,结果会返回雅可比矩阵的一行的转置。若 是标量损失,那么 ,把输出种子设为 ,就能在一次反向扫描中得到完整梯度:
这种几何关系决定了模式选择。要构造完整雅可比矩阵,通常需要 个 JVP 列或 个 VJP 行,除非批处理或结构能够减少工作。神经网络训练有许多输入和一个标量损失,因此反向模式很合适。前向模式仍适用于方向导数、输入少而输出多的函数,以及 Hessian 向量积等组合 (JAX Authors 2026)。
廉价梯度结论属于算术模型,不是墙钟时间的服务级目标。Baur 与 Strassen 证明,在特定的操作计数约定下,有理直线程序及其全部一阶偏导数具有常数因子上界 (Baur and Strassen 1983)。更一般地说,如果原语的 VJP 规则相对成本有界,其中 表示标量函数的求值成本,那么反向算术成本为 。这个上界不包括已保存张量的流量、内核启动、通信、同步,也不包括存储和写入 个梯度值。实际反向传播时间不必是前向传播时间的固定倍数。
反向累积如何工作
即时执行引擎会记录实际运行过的张量运算所形成的有向无环图(DAG)。每个节点保存一个输出值,以及一条把传入的输出伴随量映射为各父节点的贡献的反向规则。用标量记号表示,若节点 被后续节点使用,其累积伴随量为
这里, 是损失对节点 的敏感度, 表示它的直接后继,求和会从每一次使用中收集一项贡献。每个节点都累积这些贡献。标量输出的伴随量以一为种子,随后按反向拓扑顺序处理节点,确保节点向父节点传递梯度之前,所有下游贡献都已经到达。对向量值算子,反向规则应用局部雅可比矩阵的转置,而不是标量导数。
前向:
执行每个原语,并记录其父节点和局部 VJP 规则
反向:
把所有伴随量清零,并把标量输出的伴随量设为一
每个节点只按反向拓扑顺序访问一次
把每个节点的 VJP 贡献加到各个父节点的伴随量上
对 、 和 ,反向扫描得到 、、,以及 。对 的两项贡献必须相加。
下面的可运行示例直接实现了这个调度不变量。第一种情况刻意复用一个非叶子节点:,。若递归反向程序沿每条路径传播已经累积的伴随量,这张图会得到错误答案。拓扑调度只处理每个节点一次,在 时得到 。
import math
class Var:
def __init__(self, val, parents=()):
self.val = float(val)
self.parents = tuple(parents) # (parent, local derivative)
self.grad = 0.0
def __add__(self, other):
return Var(self.val + other.val, ((self, 1.0), (other, 1.0)))
def __mul__(self, other):
return Var(self.val * other.val, ((self, other.val), (other, self.val)))
def backward(self, seed=1.0):
order = topo(self)
for node in order:
node.grad = 0.0
self.grad = seed
for node in reversed(order):
for parent, local in node.parents:
parent.grad += node.grad * local
def sin(value):
return Var(math.sin(value.val), ((value, math.cos(value.val)),))
def topo(root):
order, seen = [], set()
def visit(node):
if node in seen:
return
seen.add(node)
for parent, _ in node.parents:
visit(parent)
order.append(node)
visit(root)
return order
def finite_difference(fn, x, h=1e-6):
return (fn(x + h) - fn(x - h)) / (2 * h)
x = Var(2.0)
a = x * x
z = a + a
z.backward()
print("shared analytic", x.grad, "expected", 8.0)
x = Var(2.0)
z = x * x + sin(x)
z.backward()
numeric = finite_difference(lambda t: t * t + math.sin(t), 2.0)
print("finite difference", round(numeric, 6), "analytic", round(x.grad, 6))
真实引擎保存的是向量 VJP 函数,而且只保存这些函数需要的张量,并不会在每条边上保存一个标量局部导数。它们还会区分两种累积。扇出产生的贡献在同一张反向图内相加;与此同时,参数 .grad 缓冲区往往跨多次反向调用保留,因此优化器循环会在训练步之间将其清零 (PyTorch Contributors 2026)。
可微性是算子契约的一部分
自动微分是否正确,取决于它组合的原语规则是否正确。下面几种情况都需要明确决策:
- 不可微操作可以采用文档明确的次梯度、返回零、传播
NaN,或拒绝请求。max的并列值、裁剪边界、索引、排序和整数输出都应该通过测试确认,不能依赖猜测。 - detach 或 stop-gradient 操作会有意切断一条路径。通过某个丢弃历史记录的 API 意外复制张量,也可能造成同样结果。
- 原地修改或错误的别名声明可能覆盖为反向传播保存的值。框架的版本计数器能捕获许多情况,但自定义算子仍必须如实声明修改和别名关系。
- 自定义导数可以把外部代码接入自动微分,但它的反向规则、JVP、批处理、自动混合精度、追踪和高阶行为是彼此独立的契约。如果一个操作能用内置张量算子表达,组合这些算子通常能自动保留更多契约 (PyTorch Contributors 2026)。
- 混合精度会同时改变数值和梯度的数值行为。FP16 训练可能需要损失缩放;裁剪或检查梯度之前先取消缩放,
NaN或无穷值检测也属于优化器协议。
梯度检查会在双精度下的中心有限差分与解析方向导数之间做比较,并且避开不连续点。对高维输入,选择方向 ,比较 与 。这里, 是测试方向, 是有限差分步长。这样无需逐坐标构造数值梯度,就能沿一个方向检查完整梯度。自定义规则还应覆盖零尺寸、广播、非连续和极端输入;如果承诺支持高阶导数,还要接受二阶梯度检查 (PyTorch Contributors 2026)。
反向模式以时间换内存
前向传播会保存后续 VJP 规则所需的张量。它不一定保留所有中间值,而且已保存张量可以在反向传播消费后释放。即便如此,这些激活值占用的空间仍可能远大于动态图元数据。训练总内存还包括参数、梯度、优化器状态、临时工作区、通信缓冲区和分配器预留。
激活检查点会省略部分需要保存的张量,并在反向传播时重新执行一段前向计算。对一条包含 个阶段、按每段 个阶段划分的均匀链,同时保存的激活状态数量可以用下面的简单模型估算:
这里, 是保存状态数量的峰值, 是链的长度, 是分段长度, 表示向上取整。取 时,,而且每一段最多重算一次 (Chen et al. 2016)。递归调度可以达到其他时间与内存折中。对任意 DAG、不等大小的激活值、跳跃连接或通信密集型计算图,这些上界不能原样套用。
重计算必须复现原始前向传播的数值。随机数生成器状态、可变全局量、设备迁移或其他副作用一旦改变,重算路径就可能与原路径不同,造成无声的梯度错误。框架的检查点工具会保留部分随机数生成器状态,但具体覆盖哪些设备和状态,属于它们对外说明的契约 (PyTorch Contributors 2026)。应在真实工作负载上剖析结果:检查点可能降低峰值内存,却增加算术量、暴露更多通信,或延长训练步时间。
即时执行、追踪与编译是不同维度
早期张量系统直接暴露优化所需的程序表示。Theano 构建符号计算图,TensorFlow 1 通过 session 执行静态数据流图 (Theano Development Team 2016; Abadi et al. 2016)。Chainer 和 PyTorch 记录每次具体运行的操作,让即时执行成为实用方案;PyTorch 的设计保留普通 Python 控制流,同时把张量热路径移到 Python 之外 (Paszke et al. 2019)。JAX 采用另一条路径:通过追踪来变换纯函数,把微分、向量化和编译组合成程序变换 (Frostig et al. 2018)。
这段历史并不是一场赢家通吃的框架战争。当前系统把即时执行与一种或多种分阶段表示结合起来:
| 模式 | 会发生什么 | 可以复用什么 | 主要失败模式 |
|---|---|---|---|
| 即时执行 | 宿主代码分派算子时,算子立即运行;系统仍可记录自动微分 DAG。 | 内核缓存与分配器缓存 | 宿主开销、融合不足、意外修改 |
| 追踪或分阶段执行 | 函数使用抽象值或符号值运行,生成中间计算图。 | 针对某种输入签名特化的计算图 | 追踪阶段的副作用、不支持由数据决定的控制流、重复追踪 |
| 带守卫的计算图捕获 | 系统捕获动态宿主程序中的部分区域,同时记录成立所需的假设。 | 只要守卫仍成立,就能复用缓存的可执行程序 | 计算图中断、守卫失败、重新编译、编译延迟 |
| 提前导出 | 在部署之前,按有界程序与输入契约完成降级。 | 面向该契约的序列化产物 | 不支持的动态行为、运行时假设不匹配 |
在 PyTorch 即时模式下,每次被记录的前向传播都会重新构建一张 Function DAG。torch.compile 是另一条捕获路径:TorchDynamo 提取 FX 区域,记录类型、形状、dtype、设备和步幅等假设的守卫,再把捕获的计算图交给 AOTAutograd 与后端。采用部分捕获时,不受支持的代码会造成计算图中断,随后可以在另一个捕获区域恢复执行。守卫失败时,系统可以选择另一个缓存的可执行程序,也可能触发重新编译;若要求完整计算图,中断会直接变成错误 (Ansel et al. 2024; PyTorch Contributors 2026)。
JAX 的 grad 和 vmap 是程序变换,jit 则会分阶段处理并编译一个特化版本。兼容的调用会复用缓存的可执行程序。普通 Python 副作用发生在追踪阶段,而不是每次设备执行时。若 jitted 区域内存在由运行时数据决定的分支,就需要使用结构化控制流原语,或改换分阶段边界 (JAX Authors 2026)。TensorFlow 2 默认即时执行;GradientTape 记录即时执行的 TensorFlow 操作,tf.function 则追踪并缓存特化的 TensorFlow 计算图。XLA 编译是额外选择,不是 tf.function 的同义词 (TensorFlow Authors 2026; TensorFlow Authors 2026)。
因此,真正有用的比较并不是“计算图还是磁带”,而是追踪阶段能看见哪些语义,哪些信息进入缓存键,哪些副作用仍可观察,失败时如何回退,以及冷编译成本是否由热执行摊销。
框架必须明确什么
现代框架是一组跨层契约,并不是一摞可以随意互换的整齐层次:
- 张量表示定义形状、dtype、设备、布局、步幅、别名、梯度状态,有时还包括分布式放置。
- 算子模式定义接受的输入、广播、dtype 提升、输出元数据、修改和别名。
- 分派器不仅按后端选择实现,也要处理自动微分、批处理、自动混合精度、函数化和张量子类等变换与模式。
- 自动微分变换提供 VJP 规则,并在受支持时提供 JVP 规则和已保存张量契约。
- 编译器需要抽象张量或 fake tensor 行为、计算图降级、形状推理和后端实现。
- 设备运行时负责流、事件、同步和内存分配器。异步错误与分配器碎片会经由这一层暴露,而不是经由模型方程暴露。
- 分布式张量系统把逻辑布局附着到设备网格,并定义算子如何传播或改变这种布局。
自定义算子会让遗漏暴露出来。注册一个自定义算子可能需要算子模式、后端内核、供追踪使用的 fake 或 meta 实现、自动微分规则、批处理规则、自动混合精度策略和分布式分区规则。注册检查器可以验证这些组件是否接好,却无法证明自定义导数在数学上正确。后者仍需要梯度检查,以及即时执行与编译执行的代表性对照测试。
分布式布局属于张量语义
一个分布式张量在设备网格上表示一个逻辑值。它的放置方式可以是 Shard,表示每个网格坐标只存储某个张量维度的一部分;可以是 Replicate,表示每个坐标都保存完整值;也可以是 Partial,表示每个坐标保存一项等待归约的贡献。Partial 并不是完整的逻辑结果。
改变放置方式同时具有通信语义。Shard 转为 Replicate 通常需要全收集。Partial 转为 Replicate 需要全归约,Partial 转为 Shard 可以使用归约散布。改变分片维度可能需要全交换。框架可以插入这些集合通信,却无法让传输的字节凭空消失 (PyTorch Contributors 2026)。第 62 章 中的物理拓扑与可持续链路速率,仍然约束着程序。
GSPMD 展示了编译器如何把少量分片注解传播到计算图中,并生成单程序多数据执行 (Xu et al. 2021)。PyTorch 的分布式张量与 FSDP 把类似的全局张量语义整合进分派器、自动微分和分配器 (Zhao et al. 2023)。手写逐设备程序能提供更多控制,却也把集合通信顺序、全局与局部形状以及损失归一化交给程序员负责。网格或集合通信顺序不一致可能让所有 rank 挂起,即使每个本地张量操作都有效。建立在这些接口上的并行算法见 第 10 章。
运行检查清单
正确性先于任何加速结论。对一条新的模型路径、算子或框架后端,应完成以下检查:
- 在有代表性的输入上,比较即时执行与编译执行的输出、梯度、状态更新和随机数消耗。
- 在双精度下、避开不可微点,运行方向有限差分和框架自带的梯度检查。
- 覆盖形状、dtype、设备、布局、非连续视图、广播、零尺寸张量、别名和修改,并且有意测试
NaN与无穷值行为。 - 分别测量冷编译和热执行。记录缓存键或输入签名、计算图中断位置、守卫失败和重新编译次数。
- 剖析内核时间、主机空隙、已保存张量的字节数、峰值内存、分配器预留和通信量。更快的内核仍可能让整个训练步变慢。
- 在分布式测试中,断言全局与局部形状、放置变换、归约分母、网格身份、集合通信顺序、检查点恢复,以及 rank 或链路失效后的行为。
- 明确可复现性目标:容差内一致、统计等价、重启等价,或在固定平台上逐位一致。纯函数和固定种子会有所帮助,但编译路径、归约顺序、world size、硬件和版本仍可能改变精确比特。
由编译缓存、分配器行为和后端覆盖不足引起的服务事故,将在 第 31 章 继续讨论。
框架无法抹去它下面的机器。支持的 dtype、内存布局、分块形状、集合通信带宽和编译器覆盖范围,决定哪些算子契约可以高效实现。反过来,如果一种架构的关键操作缺少正确的导数、降级或分布式规则,即使它在数学上很有希望,评估成本也可能高得无法接受。第 64 章 会沿着这项约束,从捕获的计算图继续追踪到设备代码。
尚无定论的问题,是持久的抽象边界应该放在哪里。一种观点倾向功能丰富的即时框架,由编译器捕获相容区域;另一种倾向纯粹的分阶段程序,让所有变换显式可见;第三种则把更多语义放进可移植的算子与编译器接口。这些方案会交换调试自由度、编译范围、自定义内核控制和后端可移植性。它们都不能免除明确导数、副作用、布局和失败行为的责任。只有当这些契约在目标系统上真正可用的实现越多,框架的可移植性才越强。
延伸阅读
- Baydin et al., “Automatic Differentiation in Machine Learning: a Survey,” 2018. jmlr.org这篇综述区分自动微分、符号微分与数值微分,并系统说明机器学习程序中的前向与反向累积。
- Wengert, “A simple automatic derivative evaluation program,” 1964. doi.orgWengert 将函数分解为带中间变量的基本步骤,奠定了基于求值列表进行自动微分的思路。
- Baur & Strassen, “The complexity of partial derivatives,” 1983. web.vu.ltBaur 与 Strassen 证明:在算术电路模型中,有理函数及其全部一阶偏导可在常数倍运算量内共同计算。
- Chen et al., “Training Deep Nets with Sublinear Memory Cost,” 2016. arXiv:1604.06174本文分析激活重计算调度,以额外的前向计算换取次线性的激活保存内存。
- Paszke et al., “PyTorch: An Imperative Style, High-Performance Deep Learning Library,” 2019. arXiv:1912.01703PyTorch 设计论文说明其即时张量接口、动态自动微分图、分发器、分配器与 C++ 执行路径。
- Frostig et al., “Compiling Machine Learning Programs via High-Level Tracing,” 2018. mlsys.org本文提出对纯数组程序进行高层追踪,并将微分、向量化与编译组织为可组合变换。
- Ansel et al., “PyTorch 2: Faster Machine Learning Through Dynamic Python Bytecode Transformation and Graph Compilation” (给即时执行框架追加装配的编译器), 2024. docs.pytorch.org本文说明 PyTorch 2 编译路径中的带守卫 Python 字节码捕获、图中断、AOTAutograd 与 TorchInductor。
- Xu et al., “GSPMD: General and Scalable Parallelization for ML Computation Graphs” (作为 JAX pjit 底层机制的 XLA 与 TPU 分片标注), 2021. arXiv:2105.04663GSPMD 在计算图中传播张量分片标注,并生成分区后的单程序多数据程序。
- Zhao et al., “PyTorch FSDP: Experiences on Scaling Fully Sharded Data Parallel” (VLDB 论文), 2023. arXiv:2304.11277FSDP 论文说明参数分片如何与 PyTorch 自动微分、内存分配、通信和状态管理相互作用。
- PyTorch Contributors, “Autograd Mechanics” (持续更新的官方文档), 2026. docs.pytorch.org该官方说明记录 PyTorch 的动态自动微分图、保存张量、不可微约定与原地操作正确性检查。
- JAX Authors, “The Autodiff Cookbook” (持续更新的官方文档), 2026. docs.jax.dev该手册推导雅可比向量积与向量雅可比积,并解释输入输出维度如何决定高效的微分模式。
- PyTorch Contributors, “torch.compile Programming Model” (持续更新的官方文档), 2026. docs.pytorch.org该编程模型指南定义图捕获、图中断、守卫、重新编译以及部分图与完整图契约。
- TensorFlow Authors, “Better Performance with tf.function” (持续更新的官方文档), 2026. tensorflow.org该指南解释追踪、ConcreteFunction 缓存、输入特化、重新追踪与追踪期 Python 副作用。
- PyTorch Contributors, “PyTorch DTensor: Distributed Tensor” (持续更新的官方文档), 2026. docs.pytorch.orgDTensor 契约以设备网格上的 Shard、Replicate 与 Partial 放置定义逻辑张量及其重分布语义。
评论
登录后评论