反向传播:把一个结果的责任高效分回全部参数
从计算图、局部导数和上游梯度开始,完整手算一次反向遍历,并理解分叉累加、向量—雅可比积、激活缓存和梯度检查。
- 逐个扰动参数的有限差分为什么不适合训练大模型?
- 计算图怎样把复杂函数拆成可复用的局部求导?
- 上游梯度、局部导数和参数梯度之间是什么关系?
- 变量被多条路径使用时,梯度为什么必须相加?
- 怎样手算完整前向与反向,并用有限差分检查实现?
- 前向按计算图得到预测与标量损失,并保存必要中间值。
- 损失节点以
∂L/∂L=1作为反向种子。 - 逆拓扑遍历让每个节点接收上游梯度。
- 节点用局部导数计算向量—雅可比积。
- 多条下游路径的贡献在共享变量处相加。
- 遍历结束后,每个可训练参数得到
∂L/∂θ。 - 有限差分和局部测试抽查梯度实现。
- 优化器随后才读取梯度并更新参数;目标是否正确仍由损失与评测负责。
1为什么不能把每个参数轻轻推一下试错动机
这一节在回答一个核心问题:既然有限差分可以近似导数,为什么训练神经网络还需要反向传播?
有限差分采用纯数值方法。对某个参数 θi,保持其他参数不变,分别把它向正、负方向移动一个很小的步长 ε,然后比较两次损失:
这个公式称为中心差分。它的输入是参数当前值和人为选择的扰动步长 ε,输出是损失沿第 i 个参数方向的近似斜率。
其计算过程可以分为三步:
- 将 θi 改为 θi+ε,执行一次完整前向计算,得到损失 L(θi+ε)。
- 将 θi 改为 θi-ε,再次执行完整前向计算,得到损失 L(θi-ε)。
- 用两个损失之差除以参数两次取值之间的距离 2ε,估计单位参数变化引起的损失变化。
中心差分使用扰动点两侧的信息,通常比只从一侧计算的前向差分更准确。但是,它仍然是一种近似方法,而不是从计算过程直接得到的解析梯度。
有限差分首先面临计算成本问题。假设模型有 P 个参数,检查每个参数的中心差分需要大约 2P 次完整前向计算。百万参数意味着约两百万次前向,十亿参数则需要约二十亿次前向。每次只改变一个参数,却要重新执行整个模型,绝大多数共享计算被反复重做,因此无法承担日常训练。
其次,步长 ε 存在数值上的两难选择。
当 ε 太大时,损失函数在这一段范围内可能已经明显弯曲。差商测得的是一段区间的平均斜率,而不是当前点的瞬时斜率,因此产生截断误差。
当 ε 太小时,两个损失值会极其接近。在有限精度浮点数中,对两个相近数做减法容易丢失有效数字;得到的微小误差再除以很小的 2ε,还可能被进一步放大。这就是舍入误差或消减误差。
因此,ε 并非越小越好。实际检查中常尝试多个数量级的步长,观察数值梯度是否在某个区间内保持稳定。通常随着 ε 减小,截断误差会先下降;当 ε 过小时,浮点误差又会上升。
这些限制决定了有限差分的合适用途:它适合做梯度检查,不适合直接训练模型。可以随机抽取少量参数,将有限差分结果与自动求导得到的梯度比较,以发现反向公式、路径累加或广播处理中的错误。
反向传播的效率来自不同的思路。神经网络的所有参数最终都影响同一个标量损失,而且它们到损失之间共享大量中间计算。有限差分把每个参数视为一次独立实验,因而不断重复这些计算;反向传播则保存前向计算的结构和必要中间值,再从损失出发,按照链式法则逆向遍历一次计算图。
在一次反向遍历中,共享节点对损失的敏感度只需计算一次,随后便可复用于所有上游参数。因此,反向传播能够以与少数几次前向计算相近的成本,得到全部参数的梯度。
两种方法的差别可以概括为:
- 有限差分逐个询问:“单独轻推这个参数,损失会怎样变化?”
- 反向传播统一计算:“损失的敏感度怎样沿共享计算路径传回所有参数?”
所以,有限差分的问题不是不能求梯度,而是不能高效、稳定地为海量参数反复求梯度。反向传播真正利用的是计算图中的共享结构,把原本需要逐参数重复进行的试错,转化成一次系统性的敏感度传播。
2计算图怎样把复杂函数拆成局部问题直觉
这一节解释反向传播的组织基础:计算图如何把一个复杂函数拆成许多简单、可复用的局部求导问题。
以
为例。虽然可以把损失直接写成:
并对整个复合公式求导,但真实神经网络包含大量矩阵乘法、偏置、归一化和非线性函数。若把整个模型展开为一个巨型公式,推导、实现和维护都会变得困难。
计算图采用不同的表示方法:把复合函数拆成连续的基本运算。例如:
这里的 u 就是预测值 ŷ,e 是预测误差。乘法、减法和半平方各自构成一个局部运算节点,节点之间的边表示数据依赖。由于一个运算必须等待其输入产生,普通前向计算形成一张有向无环图。
前向传播沿依赖方向计算数值:
反向传播则沿相反方向传递敏感度:
这里传递的不是前向数值,而是最终损失对当前节点输出的梯度,也就是“当前变量稍微变化时,损失会怎样变化”。
反向传播从:
开始。随后,每个节点只需把收到的上游梯度与自己的局部导数组合。
对于半平方节点:
局部导数为:
对于减法节点:
局部导数为:
对于乘法节点:
局部导数为:
按照链式法则拼接后,可以得到:
以及:
整个过程中,没有任何一个节点需要知道完整模型。乘法节点只负责乘法的反向规则,减法节点只负责减法的反向规则,平方节点只负责平方的反向规则。复杂网络的梯度就是这些局部规则按照图结构组合后的结果。
因此,计算图不只是公式的可视化,它还是一份可执行的求导计划。它至少需要记录两类信息。
第一类是运算之间的连接关系。例如 u 由 w 和 x 相乘得到,e 又依赖于 u 和 y。这些关系决定反向传播应按什么顺序经过哪些节点。
第二类是局部反向公式需要的前向值。例如:
- 计算乘法节点中 w 的梯度需要前向值 x;
- 计算 x 的梯度需要前向值 w;
- 计算半平方损失的反向需要误差 e=ŷ-y;
- 激活函数可能需要知道前向输入位于哪个区域。
这说明“图结构还存在”并不等于“反向一定能够执行”。如果框架仍然知道某个节点做过乘法,却已经丢失或覆盖了乘法反向所需的输入值,便无法正确计算局部梯度。
计算图与中间状态的关系也解释了训练显存高于推理显存的原因。推理完成一层后,若其激活不再参与后续前向,通常可以释放;训练却要保留反向所需的中间值,直到对应节点完成反向计算。
它也解释了原地修改为何危险。若某个张量在前向后被直接覆盖,而反向公式需要原来的值,反向读取到的就不再是产生当前损失时的状态。即使前向结果看起来正常,梯度也可能报错或失真。
当计算图存在分支时,同一个变量可能沿多条路径影响损失。各条路径会分别传回梯度贡献,并在共享节点处求和。这是多元链式法则在图结构中的表现,也是计算图能够处理残差连接、权重共享和循环结构的基础。
所以,计算图解决复杂求导问题的方法不是推导更大的整体公式,而是建立清晰的局部协作机制:
它把全局求导转化为许多小型局部问题,同时也划定了一条硬边界:反向传播的正确性依赖于前向计算图和必要中间值都保持完整。
3上游梯度乘局部导数是什么意思链式法则
这一节把反向传播缩小到一个运算节点,说明梯度在节点之间究竟如何传递。
设某个节点执行:
最终损失 L 又依赖于 z。反向传播经过该节点时使用链式法则:
这三个量各有明确含义:
- ∂ L/∂ z 是节点收到的上游梯度,表示 z 发生微小变化时,最终损失会怎样变化;
- ∂ z/∂ u 是节点的局部导数,表示 u 发生微小变化时,节点输出 z 会怎样变化;
- ∂ L/∂ u 是节点传给输入 u 的梯度,表示 u 的微小变化最终会怎样影响损失。
两种敏感度之所以相乘,可以从微小变化理解:
将第一式代入第二式:
因此:
第一项描述从 z 到损失的影响,第二项描述从 u 到 z 的影响。两段首尾相接后,就得到从 u 到损失的完整局部影响。
常用的横杠记号可以把公式写得更紧凑:
于是单个节点的反向规则写成:
也就是:先接收上游梯度,再用本节点的局部导数对它进行变换,最后把结果传给输入。
对于加法:
局部导数为:
所以:
两个输入都收到完整的上游梯度。这不是把梯度平均分配给两个输入,而是因为任意一个输入增加一点,输出都会增加同样的量。
对于乘法:
局部导数为:
因此:
计算 a 的梯度需要前向时的 b,计算 b 的梯度需要前向时的 a。这说明为什么反向传播必须保存某些前向中间值。
对于平方:
局部导数为:
所以:
上游梯度在通过平方节点时,会按照当前输入 u 的数值和符号被缩放。
对于 ReLU:
当 u>0 时,局部导数为 1;当 u<0 时,局部导数为 0。因此:
ReLU 可以被理解为一个梯度门:正区间让上游梯度原样通过,负区间将其阻断。
在 u=0 处,ReLU 的左右导数不同,经典导数不存在。自动求导框架必须约定一个值,通常选择某个次梯度。这个值并非由普通链式法则唯一决定,而是算子的实现约定。
对于矩阵乘法:
设上游梯度为:
则输入梯度为:
假设:
那么:
所得梯度分别与 W 和 X 的形状一致。检查梯度形状是判断矩阵反向公式中转置位置是否正确的一种基本方法。
在标量例子里,“上游梯度乘局部导数”看起来只是普通乘法;在向量和张量情形下,局部导数可能是雅可比矩阵。严格来说,节点执行的是上游梯度与局部雅可比的乘积。框架通常不会显式创建完整雅可比,而是直接实现这一乘积。
如果节点有多个输入,它需要分别为每个输入计算梯度。如果同一个输入通过多条路径影响损失,各条路径返回的贡献还必须相加。例如 u 同时影响 z1 和 z2,则:
因此,单条路径遵循“上游梯度乘局部导数”,多条路径在共享节点处还要执行“各路径贡献求和”。
反向传播没有创造新的微积分规则。它只是把链式法则组织成一种高效执行过程:每个节点只实现自己的局部反向规则,框架再按照计算图的逆序将这些规则组合起来。
这种机制也有明确前提:局部导数必须存在或有清晰约定,并且反向规则必须正确实现。如果自定义算子的前向计算正确、反向公式错误,模型仍能产生正常的前向数值,但错误梯度会从该节点继续传向所有上游变量,使整条训练链路受到污染。
| 局部操作 | 前向 | 反向给输入 |
|---|---|---|
| 加法 | z=a+b | 两边都接收上游梯度 |
| 乘法 | z=ab | ā=ż·b, b̄=ż·a |
| 平方 | z=u² | ū=ż·2u |
| ReLU | z=max(0,u) | u>0 时传递,否则为 0 |
| 矩阵乘 | Z=WX | 按转置乘法得到 W̄ 与 X̄ |
4完整手算:−4 怎样一步步出现数值例子
这一节通过完整手算说明:梯度 -4 并不是突然出现的,而是损失端的种子梯度经过多个局部节点逐步缩放后得到的结果。
给定:
模型和损失为:
它们共同表示复合函数:
计算时不直接对整个复合公式求导,而是先完成前向传播,保存中间值,再沿相反顺序执行反向传播。
首先计算前向。
乘法节点:
这里 u 是模型预测值。
减法节点:
这里 e=-2 表示预测值比目标值低 2。
半平方损失节点:
前向传播最终得到:
其中 u 和 e 不只是临时计算结果,也是反向传播所需的原料。平方节点的局部导数依赖 e,乘法节点的局部导数依赖 w 和 x。
反向传播从最终损失开始。其种子梯度是:
这是因为任何变量对自身的导数都等于 1。这个 1 不是学习率,也不是人为猜测的梯度大小,而是启动反向传播的一单位敏感度。
接下来经过半平方节点:
该节点对输入 e 的局部导数为:
把前向保存的 e=-2 代入,并乘以上游梯度 1:
这一步中的负号来自前向误差 e=-2。平方运算虽然让最终损失变成正数,但其导数仍保留误差的方向信息。
然后经过减法节点:
该节点关于 u 的局部导数为:
因此:
梯度经过通往 u 的减法路径时,大小和符号都没有改变。
最后经过乘法节点:
通往参数 w 的局部导数是:
因为前向时 x=2,所以:
因此,-4 的完整来源可以写成:
同一个乘法节点还可以把梯度传给输入 x。由于:
并且 w=2,所以:
若把目标值 y 也视为需要求导的变量,那么减法节点在通往 y 的路径上具有局部导数:
因此:
不过在普通监督学习中,x 和 y 通常是数据,优化器只更新可训练参数 w。
梯度:
表示在当前点 w=2 附近,如果让 w 增加一个很小的量 Δ w,损失的变化近似为:
因此,稍微增大 w 会使损失下降。这与当前预测偏低相符:因为 x=2>0,增大 w 会增大预测值 u=wx,使其从 4 向目标值 6 靠近。
但 -4 并不表示应该直接让 w 增加 4。若采用最简单的梯度下降,更新公式是:
代入当前梯度:
真正的更新幅度由学习率 η 和优化器规则决定。反向传播只提供当前点附近的方向和变化率,不负责决定实际步长。
最后,这些前向值和梯度只对:
这一组取值成立。参数更新后,预测值 u、误差 e、损失 L 和各局部梯度都会变化。因此,每次参数更新后都需要重新执行前向传播,再基于新的中间状态进行反向传播。
| 方向 | 节点 | 计算 | 值 |
|---|---|---|---|
| 前向 | u=wx | 2×2 | 4 |
| 前向 | e=u−y | 4−6 | −2 |
| 前向 | L=½e² | ½×4 | 2 |
| 反向 | 损失种子 | ∂L/∂L | 1 |
| 反向 | 经过平方 | ∂L/∂e=1×e | −2 |
| 反向 | 经过减法 | ∂L/∂u=−2×1 | −2 |
| 反向 | 经过乘法到 w | ∂L/∂w=−2×x | −4 |
| 反向 | 经过乘法到 x | ∂L/∂x=−2×w | −4 |
5分叉路径为什么要累加梯度多路径
这一节说明多元链式法则中的一个关键规则:当同一个变量通过多条计算路径影响最终结果时,它的总梯度等于所有路径贡献之和。
以:
为例。参数 w 在计算图中有两条通往 u 的路径:
以及:
第一条路径经过平方运算,第二条路径直接进入加法。只要 w 发生微小变化,这两条路径的输出就会同时变化,因此 u 的总变化必须包括两部分。
平方路径的导数贡献为:
直接路径的导数贡献为:
所以:
这里的加号不是代数展开后的偶然结果,而是计算图中分叉路径反向传播的一般规则:一个变量若通过多条下游路径影响输出,各条路径返回的梯度贡献必须在该变量处求和。
如果最终还有损失 L 依赖于 u,设从下游传到 u 的上游梯度为:
平方路径传回 w 的贡献为:
直接路径传回 w 的贡献为:
因此,w 的完整梯度是:
即:
这就是多元链式法则在计算图中的路径求和形式。
从微小变化的角度也可以得到同样结论。假设 w 变化了一个很小的量 Δ w。平方分支的变化近似为:
直接分支的变化为:
由于加法节点把两个分支的结果相加,u 的总变化为:
因此:
如果只沿平方路径反向传播,就会错误地得到:
这漏掉了直接路径的贡献 1。错误结果仍可能形状正确、数值有限,看起来像一个合理梯度,因此比直接报错更隐蔽。
漏掉路径也不一定总是让梯度变小。不同路径的贡献可能同号,也可能异号。例如一条路径贡献 5,另一条路径贡献 -3,正确总梯度是 2。若漏掉其中一条,结果可能变成 5 或 -3,不只改变大小,还可能改变符号。因此,漏掉路径可能使训练变慢,也可能使参数朝错误方向更新。
在自动求导的实现中,某个变量的梯度缓冲区不能在每次收到新贡献时直接覆盖:
而应执行累加:
如果使用赋值覆盖,后到达的路径会抹掉先到达的路径,最终只保留一部分导数。正确的反向遍历需要等待或汇总所有下游路径的贡献,得到完整总梯度。
这种路径累加在神经网络中十分常见。
残差连接通常写成:
输入 x 一条路径经过变换 F,另一条路径通过恒等连接直接到达加法节点。因此:
其中 I 是恒等路径的局部导数。反向传播必须同时保留变换分支和捷径分支的贡献。
权重共享也会产生路径累加。若同一个参数 w 在前向中被使用多次,例如:
那么:
虽然它们引用的是同一个参数对象,但每次使用都形成了一条独立的下游路径。所有使用位置产生的梯度贡献都必须累加到同一个参数上。
循环神经网络中,同一组权重在多个时间步反复使用。将循环沿时间展开后,每个时间步都形成一条使用共享权重的路径,因此该权重的最终梯度是所有时间步贡献的总和。这是通过时间反向传播中的基本机制。
还需要区分图内路径求和与跨 mini-batch 梯度累积。
图内路径求和发生在同一次前向和反向过程中。它来自函数本身的依赖结构,是计算正确总导数不可缺少的一部分。只要计算图中存在分叉、共享或重复使用,自动求导系统就必须执行这种累加。
跨 mini-batch 梯度累积则是一种训练策略。它在多次独立的前向—反向之后,暂时不清空参数的梯度缓冲区,从而将多个批次的梯度相加,用来模拟更大的有效批量。例如连续累积 K 个微批次后再更新参数。
这两类累加都使用加法,但原因不同:
- 图内分叉求和是链式法则的要求,不能省略;
- 跨批次梯度累积是训练配置,可以选择是否使用;
- 图内求和由计算图自动完成;
- 跨批次累积需要训练代码明确决定何时保留梯度、何时更新以及何时清零。
因此,理解分叉路径的关键是:一个变量的梯度不是任选一条路径得到的局部结果,而是该变量通过所有下游路径对最终损失造成的总影响。只有把每条路径的贡献完整相加,所得梯度才是真正的总导数。
6张量很多时为什么不显式构造巨大雅可比向量化
这一节解释反向模式自动微分为什么能够处理高维张量:它通常不构造完整雅可比矩阵,而是直接计算向量—雅可比积。
设某一层为:
其中:
这个函数的完整雅可比矩阵是:
矩阵中的元素为:
它记录每个输入分量对每个输出分量的局部影响。如果输入和输出都很大,显式存储 J 会非常昂贵。例如 m=n=10 000 时,雅可比包含一亿个元素;更大的张量会产生更加难以承受的存储和计算成本。
但神经网络训练真正需要的通常不是 J 本身。最终目标是一个标量损失:
当反向传播到当前层时,下游已经给出了损失对层输出的梯度:
根据链式法则,损失对层输入的梯度为:
这个运算称为向量—雅可比积,即 VJP。
如果统一使用列向量表示梯度,同一个关系也可以写成:
两种形式只采用了不同的行列向量约定,含义完全相同。
关键在于,最终结果只是一个与输入 x 同形的向量,包含 n 个元素。既然只需要 vTJ,就没有必要先创建含有 mn 个元素的完整雅可比,再执行乘法。框架可以利用每种运算的结构,直接从上游梯度计算输入梯度。
以线性变换为例:
损失对输出的梯度为:
则输入梯度可以直接计算为:
如果还需要参数 W 的梯度,则:
两个结果都可以通过矩阵乘法直接得到,无须创建输出对 x 或 W 中每个元素的完整偏导数组。
对于逐元素激活函数:
其雅可比在理论上是对角矩阵:
显式创建这个对角矩阵会存储大量无用的零。反向传播只需计算:
其中 ⊙ 表示逐元素乘法。整个过程只处理与输入同等规模的张量。
卷积也无须展开成庞大的矩阵。卷积的反向可以利用卷积或相关运算直接得到输入和卷积核的梯度。注意力、归一化和其他复杂算子同样会实现专门的反向规则,直接完成所需的 VJP。
因此,自动求导框架为算子实现的接口可以理解为:
完整雅可比只存在于数学描述中,不必作为真实张量出现在显存里。这种做法同时节省了存储、构造雅可比的时间以及随后进行大规模乘法的成本。
反向模式尤其适合神经网络训练,是因为训练问题通常具有大量输入和少量输出。若模型有 P 个参数:
而最终损失只有一个标量:
则完整导数是一个含 P 个元素的梯度:
反向模式从:
开始,一次逆向遍历便能得到损失对全部参数的梯度。它的成本主要随计算图规模增长,而不是为每个参数单独执行一次求导。
前向模式自动微分传播的是另一种乘积。给定输入方向 r,前向模式计算:
这称为雅可比—向量积,即 JVP。它回答的是:如果输入沿方向 r 发生微小变化,所有输出会怎样变化。
两种模式的区别可以概括为:
- 反向模式计算 VJP,适合大量输入、少量输出;
- 前向模式计算 JVP,适合少量输入、大量输出。
如果函数只有几个输入,却有大量输出,前向模式可以针对一个输入方向一次得到所有输出的变化。若用反向模式获得完整雅可比,则通常需要针对多个输出方向反复执行反向。
相反,神经网络训练有数百万甚至数十亿个参数,却通常只有一个标量损失。若采用前向模式计算每个参数方向的导数,就需要许多次传播;反向模式只需从标量损失反传一次,因此更为合适。
当被求导的输出不是标量时,反向模式必须提供一个与输出同形的种子向量 v。此时计算的是:
而不是整个 J。它等价于先将向量输出按照 v 加权组合成一个标量,再求这个标量对输入的梯度。标准训练中的损失已经是标量,因此初始种子自然是 1。
完整雅可比并非绝对不能获得。某些自动微分框架提供显式 Jacobian 接口,其实现通常通过多次 VJP 或 JVP 调用,将不同方向的结果组合成完整矩阵。但当输入和输出维度都很大时,这种计算的时间和空间成本仍然很高,因此不适合作为普通训练的默认方式。
所以,反向传播的高效之处不仅在于应用链式法则,还在于只计算训练真正需要的对象:沿标量损失给出的方向,把上游梯度直接变换为输入梯度,而不显式保存每个输出对每个输入的全部偏导。
| 模式 | 适合的输出/输入关系 | 直觉 |
|---|---|---|
| 反向模式 | 大量输入参数 → 少量标量输出 | 一次反向得到损失对所有参数梯度 |
| 前向模式 | 少量输入 → 大量输出 | 沿给定输入方向向前传播导数 |
7反向为什么比纯推理占更多内存工程
这一节解释训练显存高于纯前向推理的根本原因:反向传播需要使用前向阶段的中间状态,因此许多在推理中可以立即释放的张量,在训练中必须继续保留。
模型参数数量虽然没有改变,但训练和推理对中间结果生命周期的要求不同。
在纯前向推理中,一层的输出交给下一层后,如果上一层的激活不再参与后续计算,其内存通常就可以释放或复用。系统主要需要保留当前仍在前向路径上使用的张量,因此中间状态可以像流水线一样逐步产生、逐步丢弃。
训练除了前向计算,还要沿计算图逆序执行反向传播。许多局部反向公式都依赖前向时的输入、输出或统计量,不能只凭算子名称和上游梯度凭空计算。
例如,对于 ReLU:
反向规则为:
因此,反向阶段必须知道前向输入的哪些位置大于 0。框架可以保存输入、输出或压缩后的掩码,但总要保留足以恢复这个判断的信息。
对于矩阵层:
其反向公式为:
计算权重梯度需要前向输入激活 X,计算输入梯度则需要参数 W。因此,输入激活不能在前向完成后立即丢弃。
归一化层的反向可能需要前向计算中的输入、均值、方差或标准化结果。其他激活函数、注意力和卷积算子也都有各自必须保存的中间状态。
训练前向因此承担两项任务:
- 计算预测和损失;
- 为未来的反向传播保存必要材料。
这些等待反向使用的激活,是训练相较推理增加的第一大类内存占用。
训练时主要的内存项目通常包括:
- 模型参数;
- 前向激活和局部反向所需的中间状态;
- 参数梯度;
- 优化器状态;
- 算子执行期间使用的临时工作空间。
其中,参数是推理和训练都需要的基础占用。训练额外增加的主要是激活、梯度和优化器状态。
参数梯度通常与对应参数具有相同形状。对于每个可训练参数 θ,反向传播需要保存:
因此,仅参数和梯度两项就可能需要两份同等规模的存储,具体大小还取决于各自的数据类型。
优化器还可能为每个参数维护跨训练步骤持续存在的状态。例如动量类优化器会保存历史梯度统计量。这些状态不是一次反向结束后就能释放的临时数据,而是下一次更新仍要继续使用的训练状态。
激活内存则通常受以下因素影响:
- 批量大小;
- 序列长度;
- 特征维度;
- 网络深度;
- 算子需要保存的中间状态种类。
批量越大,每一层的激活包含的样本越多。序列越长,注意力和序列模型中的中间张量通常越大。网络越深,同时等待反向使用的层数也越多。因此,即使模型参数完全不变,增大批量大小或上下文长度也可能显著增加训练显存。
一次训练迭代中的激活生命周期大致为:
反向刚开始时,大量层的激活仍然同时保存在显存中,所以峰值占用通常远高于只做前向推理。
梯度检查点改变了“保存全部必要激活”的策略。它把网络划分成若干计算段,前向时只保存少数边界位置的激活,也就是检查点。段内其他中间值在前向使用后即可释放。
当反向传播到某一段时,系统从最近的检查点出发,重新执行该段前向,临时恢复反向所需的中间值,然后完成这段反向。
这种方法的核心权衡是:
保存的检查点越少,通常能释放的激活越多,但反向时需要重复计算的内容也越多。因此,梯度检查点用额外计算时间换取更低的显存占用。
在理想条件下,梯度检查点不会改变数学上的梯度。第一次前向与重算前向若得到完全相同的中间值,后续执行的仍是同一套局部导数和链式法则。
但这要求重算过程能够复现原前向结果。
例如 dropout 在前向中生成随机掩码。若第一次前向使用掩码 M1,重算时却使用另一张掩码 M2,那么恢复的中间值就不属于产生当前损失的那次计算。反向传播随后计算出的梯度也不再对应原损失。因此,检查点机制通常需要保存或恢复随机数状态,使重算得到相同掩码。
原地修改同样危险。如果检查点保存的输入或重算所依赖的张量后来被直接覆盖,那么重算得到的就不是原来的前向过程。
带副作用的算子也可能破坏可重复性。例如算子会修改外部状态、更新计数器、读取变化中的可变数据,或者依赖无法恢复的随机行为。第二次执行即使输入看似相同,也可能产生不同结果。
因此,使用梯度检查点时需要特别注意:
- 随机层能否恢复相同随机状态;
- 原地操作是否覆盖了重算所需的输入;
- 算子是否修改外部状态;
- 两次前向是否读取相同数据和模型状态;
- 实现中是否存在无法复现的非确定性行为。
还要区分“评估模式”和“关闭梯度记录”。将模型切换到评估模式,通常只会改变 dropout、归一化等层的行为,不一定停止计算图构建。如果程序仍要求梯度,框架仍可能保存中间状态。要获得纯推理的低内存特征,通常还需要明确禁用梯度记录。
因此,训练显存高并不是模型参数突然变多,而是系统同时承担了更多职责:保存前向历史、存放反向梯度,并维护优化器状态。梯度检查点则通过选择性遗忘和按需重算,重新安排了计算与存储之间的权衡。
8哪些代码操作会悄悄切断或污染计算图失败模式
这一节讨论自动求导的工程边界:即使数学公式完全正确,程序实际执行的张量操作仍可能切断或污染计算图,使梯度变成 None、全零、NaN、Inf,或者产生形状正确但数值错误的结果。
排查这类问题时,首先要区分不同症状可能表达的含义:
- 梯度为
None,通常说明参数与损失之间不存在可追踪的计算路径,或者该参数没有参与此次前向计算; - 梯度恒为零,可能是真实导数为零,也可能来自梯度阻断、激活饱和、不可微操作或低精度下溢;
- 梯度出现
NaN或Inf,通常与数值溢出、非法运算或不稳定的计算有关; - 梯度方向似乎合理但大小异常,可能来自漏掉路径、错误广播、旧梯度累积或缩放处理错误;
- 梯度随训练步骤不断增大,可能是忘记清零,也可能是模型本身发生梯度爆炸。
症状只能缩小排查范围,不能单独确定根因。可靠诊断必须回到代码实际执行的张量操作和独立的梯度检查结果。
第一类故障是计算图被切断。
设:
如果随后执行分离操作:
再计算:
虽然 b 的数值与 a 相同,但 b 已经不再保留 a 的计算历史。损失 L 无法沿:
这条路径返回到参数 w,因此 w 的梯度可能为 None。
把张量转换为普通数组、普通数值,或者交给不受自动求导系统追踪的外部程序,也可能产生相同效果。即使之后再把结果转换回张量,新张量通常也只是一个没有原计算历史的独立对象,原先的依赖关系不会自动恢复。
排查时应检查:
- 参数是否启用了梯度追踪;
- 损失是否确实依赖该参数;
- 中间是否使用了
detach或等价操作; - 是否把张量转成普通数组或标量后又重新构造张量;
- 相关计算是否发生在禁用梯度记录的作用域中;
- 参数是否只存在于未实际执行的条件分支;
- 模型中是否存在此次样本没有使用到的参数。
需要注意,梯度为 None 与梯度为零并不相同。None 通常表示没有形成可求导路径;零梯度则表示路径可能存在,但当前局部导数或总路径贡献为零。
第二类故障是原地修改保存的激活。
反向传播经常需要读取前向阶段保存的张量。若代码直接在原内存上修改该张量,反向阶段看到的可能不再是产生当前损失时的值。
例如某个节点的反向公式需要原始输入 u,但前向之后代码执行了:
如果这是原地覆盖,那么反向读取的 u 已经发生变化。局部导数可能无法计算,或者根据错误状态计算。
成熟框架通常会通过张量版本信息检测部分危险修改并报错,但没有报错并不保证所有情况都安全。共享存储、视图、自定义算子和绕过框架检查的操作仍可能造成隐蔽污染。
排查方法包括:
- 暂时把原地操作改成生成新张量的非原地形式;
- 启用自动求导异常检测;
- 检查多个张量是否共享底层存储;
- 查看报错指向的前向节点,而不只关注反向时报错的位置;
- 检查自定义算子是否保存后又修改了同一个张量。
第三类故障是忘记清零参数梯度。
许多自动求导框架默认把新梯度累加到参数已有的梯度缓冲区中:
这是为了支持同一计算图中的路径累加和跨 mini-batch 梯度累积。但如果训练代码希望每个批次独立更新,却没有在正确时机清零,那么后续步骤看到的就不是当前批次梯度,而是多个批次贡献之和。
典型训练顺序为:
如果有意进行梯度累积,则应明确:
- 累积多少个微批次;
- 损失或梯度是否需要按累积步数缩放;
- 何时执行参数更新;
- 更新后何时清零。
排查时可以记录每一步反向后的梯度值或梯度范数。如果梯度在没有数学原因的情况下持续包含历史贡献,应检查清零时机。
第四类故障来自混合精度中的下溢和溢出。
低精度浮点格式的可表示范围和有效数字有限。很小的梯度可能下溢为 0,很大的激活或梯度可能溢出为 Inf,后续运算又可能产生 NaN。
损失缩放的基本方法是先把损失乘以一个尺度 s:
于是:
原本过小的梯度被放大,降低下溢风险。参数更新前再把梯度除以 s,恢复原来的尺度。
诊断时应检查:
- 梯度和激活中有限值所占比例;
- 零值、
NaN和Inf的数量; - 问题是否只在低精度训练中出现;
- 临时切换到更高精度后问题是否消失;
- 损失缩放和反缩放顺序是否正确;
- 是否在反缩放之前错误地进行梯度裁剪或阈值判断。
如果动态损失缩放检测到非有限梯度,通常会跳过当前参数更新并调整缩放因子。这种跳步本身可能是正常保护行为,但若频繁发生,说明数值范围仍有问题。
第五类问题是不可微离散操作。
取整、硬阈值、离散采样、取类别索引和某些选择操作,可能使输出在输入的一大片区域内保持不变。此时局部导数通常为零;在跳变位置,经典导数可能不存在。
这并非计算图损坏,而是所选择的运算本身没有适合普通梯度下降的连续导数。解决方法取决于任务,可以采用:
- 平滑近似;
- 可微松弛;
- 代理梯度;
- 专门的梯度估计器;
- 重新设计损失或模型目标。
这些方法得到的可能是近似梯度、代理梯度或随机估计量,不应误认为原离散函数的精确普通导数。
第六类问题是错误广播。
广播允许不同形状的张量完成合法运算,但前向能够运行不等于轴语义正确。
例如:
如果 b 在批次维度上被重复使用,那么反向传播计算 b 的梯度时,必须把所有批次位置的贡献求和回 b 的原始形状:
如果张量轴排列错误,程序可能仍能广播成功,却在错误维度上重复参数。反向传播随后也会沿错误维度求和。最终梯度形状甚至可能仍然正确,但每个数值的实际含义已经错位。
因此,应检查:
- 每个张量的完整形状;
- 各轴分别代表批次、时间、通道还是特征;
- 广播实际发生在哪些轴;
- 反向归约是否对应被扩展的轴;
- 改成尺寸为 1、2、3 的小例子后能否手工验证。
一个实用的排查顺序是:
- 确认参数与损失之间是否存在完整的计算图路径。
- 检查是否使用了分离操作、普通数组转换或禁用梯度作用域。
- 检查反向所需的激活是否被原地修改。
- 明确参数梯度的清零和有意累积时机。
- 检查梯度及激活中的零、
NaN、Inf和其他异常值。 - 用更高精度进行对照,判断是否存在混合精度问题。
- 检查离散或不可微节点。
- 检查广播、归约和张量轴语义。
- 在小型确定性例子中,将自动梯度与手算或有限差分比较。
执行梯度检查时,应尽量关闭随机性、固定数据和模型状态、使用双精度,并一次只检查少量参数。若数值梯度稳定而自动梯度明显不符,应优先怀疑局部反向规则、图路径、分支累加和广播归约。
同一种症状可能由多个故障共同造成。例如,忘记清零可能与低精度下溢同时存在;修复清零问题后,零梯度现象仍可能继续。错误广播也可能与某个分支被分离同时出现。
因此,每修复一处,都应重新运行最小前向—反向测试和梯度检查。报错消失、梯度不再为零或损失开始下降,都只能证明某个症状发生了变化,不能单独证明整条计算图和所有梯度已经正确。
| 故障 | 机制 | 检查方法 |
|---|---|---|
detach / 转成普通数组 | 明确停止记录后续依赖 | 检查 requires-grad 和图边界 |
| 原地修改保存的激活 | 反向需要的前向值已被覆盖 | 启用异常检测,避免危险 in-place |
| 忘记清零梯度 | 多数框架默认累加到旧值 | 记录每步梯度并明确 zero-grad 时机 |
| 混合精度下溢/溢出 | 小梯度变 0 或大值变 Inf | 损失缩放、检查 finite 比例 |
| 不可微离散操作 | 局部没有可用连续导数 | 使用代理、估计器或改写目标 |
| 错误广播 | 形状合法但梯度在错误维度求和 | 检查张量形状和小例子手算 |
9运行示例:用有限差分给自动梯度做体检验证
这一节说明如何用有限差分检查自动求导:反向传播负责高效计算全部参数的梯度,有限差分则作为一条相对独立的数值路线,对少量参数进行抽查。
仍使用贯穿示例:
模型和损失为:
反向传播给出的解析梯度是:
在 w=2 处:
现在不使用这条反向公式,而是只运行前向损失,从参数两侧估计数值梯度。中心差分为:
取:
向正方向扰动:
对应的预测为:
损失为:
向负方向扰动:
对应的预测为:
损失为:
代入中心差分:
因此:
在这个例子中,损失关于 w 是二次函数。中心差分在理想算术下能消去二次项造成的一阶近似误差,所以结果特别准确;实际程序中仍会存在有限的浮点误差。
数值梯度可以作为自动梯度的独立参照,是因为它只依赖:
这两次前向计算的损失值。它不使用被检查节点的反向公式。如果自定义反向规则写错,而前向计算正确,自动梯度与数值梯度通常会出现差异。
比较二者时,可以使用相对误差:
其中 δ 是一个很小的正数,用于防止分母为零。
相对误差考虑了梯度自身的尺度。例如,当梯度大小为 103 时,10-5 的绝对差通常可以忽略;当梯度大小为 10-6 时,同样的绝对差就可能很严重。
不过,当两个梯度都接近零时,相对误差可能因分母很小而显得很大。因此,实际检查中应同时观察绝对误差:
只有结合绝对误差、相对误差和梯度尺度,才能合理判断结果。
有限差分检查需要满足若干条件,否则即使反向实现正确,也可能出现不一致。
第一,应尽量避开不可导点。以 ReLU 为例:
在 u=0 处,左右导数不同。中心差分会同时观察折点两侧,而自动求导框架会按照预先约定返回某个次梯度。二者可能不同,但这不一定说明反向实现有错。
对于分段函数,应分别选择每个可导区域内的测试点,并对不可导边界按照算子约定单独处理。
第二,前向计算必须具有确定性。若两次有限差分前向使用了不同的 dropout 掩码,那么:
同时包含参数扰动和随机噪声,差商就无法可靠表示参数方向的导数。
因此需要:
- 关闭 dropout 等随机行为,或严格恢复相同随机状态;
- 固定输入数据;
- 固定模型状态;
- 避免两次前向之间更新运行统计量;
- 避免带有不可控副作用的算子。
第三,ε 必须处于合适范围。
当 ε 太大时,差商测量的是较宽区间上的平均斜率,不能充分代表当前点的局部导数,会产生截断误差。
当 ε 太小时,两个损失值非常接近,相减时容易丢失有效数字。得到的舍入误差再除以很小的 2ε,可能被进一步放大。
因此,不应只尝试一个 ε,而应测试多个数量级。典型现象是:
- 较大的 ε 受截断误差影响;
- 中间某个范围内,数值梯度相对稳定;
- 极小的 ε 开始受到浮点舍入误差影响。
检查应优先使用这个稳定区间中的结果。
第四,梯度检查最好使用双精度。中心差分需要对两个非常接近的损失做减法,单精度更容易丢失有效数字。梯度检查的目标是验证求导逻辑,因此不必拘泥于正式训练所使用的低精度设置。
第五,应把网络和参数规模缩小。有限差分每检查一个参数元素,都需要额外执行前向计算。检查完整大模型既昂贵,也不利于定位问题。
对于张量参数,可以随机选取少量索引。每次只扰动其中一个元素,其他元素保持不变,再将数值梯度与自动梯度中同一位置的值进行比较。
第六,应按照从局部到整体的顺序检查:
- 先检查单个自定义算子的各输入梯度;
- 再检查一层或少量算子的组合;
- 最后检查小型端到端网络;
- 对不同参数位置和不同输入区域重复抽查。
如果直接检查完整网络并发现不一致,问题可能来自任何一个节点、广播维度或分叉路径,很难定位。局部测试能更快判断是哪条反向规则出现错误。
测试用例还应覆盖算子的不同分支。例如检查 ReLU 时,应分别选择正输入和负输入;检查广播时,应使用能暴露归约轴的小张量;检查共享参数时,应确保同一参数确实沿多条路径参与损失。
若梯度检查失败,可以根据现象进一步定位:
- 多个 ε 下数值梯度都稳定,但与自动梯度不一致:优先检查反向公式、路径累加和广播归约;
- 数值梯度随 ε 剧烈变化:检查浮点精度、随机性、不可导点和损失尺度;
- 自动梯度为
None:检查参数是否参与损失以及计算图是否被切断; - 自动梯度为零而数值梯度非零:检查分离操作、错误局部导数、不可微操作和低精度下溢;
- 单层检查通过但端到端检查失败:检查层间组合、共享变量、分叉求和和状态变化。
即使所有抽查都通过,也不能证明整个反向实现对所有输入和所有参数都正确。它只能说明:在所选择的参数元素、输入数据和取值点上,自动梯度与有限差分给出了相近结果。
因此,梯度检查属于有针对性的抽样验证。它应当覆盖代表性的参数、数值范围和计算分支,并结合小模型、双精度、确定性前向、中心差分和多个步长,尽量让这份独立证据可靠。
10把整条因果链连起来综合
这一节把一次训练迭代中的完整因果链连接起来,并划清反向传播与损失函数、优化器和评测之间的职责边界。
整个过程可以概括为:
反向传播位于损失计算和参数更新之间。它负责回答“当前损失对每个参数有多敏感”,但不负责定义训练目标,也不负责决定参数实际更新多少。
第一步是使用当前参数执行前向传播。
给定输入 x、目标 y 和模型参数 θ,模型产生预测:
损失函数再计算:
训练中通常把损失归约成一个标量。前向传播还会记录计算图的依赖关系,并保存反向公式需要的中间值,例如输入激活、误差、激活掩码和归一化统计量。
此时得到的只是当前参数在当前数据上的预测和损失,参数本身尚未改变。
第二步是从标量损失启动反向传播。
最终节点是损失自身,因此反向种子为:
这个 1 表示损失对自身的一单位敏感度。它不是学习率,也不是人为设定的优化信号,而是反向传播的数学起点。
如果待求导输出不是标量,则需要额外提供一个与输出同形的种子梯度。标准训练通常已经将输出误差归约为标量损失,所以初始种子自然是 1。
第三步是按照逆拓扑顺序遍历计算图。
前向传播必须先得到节点输入,才能计算节点输出。反向传播则必须先知道损失对节点输出的梯度,才能求损失对节点输入的梯度。因此,反向遍历顺序与前向计算顺序相反。
设某个节点执行:
节点从下游收到上游梯度:
再结合自身的局部导数,计算传给输入的梯度:
每个算子只负责自己的局部规则。乘法节点处理乘法的导数,激活节点处理激活函数的导数,矩阵乘法节点处理对应的转置乘法。它们不需要知道完整模型的结构。
第四步是直接计算向量—雅可比积。
对于张量节点:
完整局部导数是雅可比矩阵:
但反向传播不需要显式构造 J。如果下游梯度为:
当前节点只需计算:
这就是向量—雅可比积。算子利用自身结构直接完成这一计算,避免生成巨大的完整雅可比矩阵。
第五步是在分叉和共享变量处累加梯度贡献。
如果同一个变量通过多条路径影响损失,每条路径都会返回一份梯度。该变量的总梯度必须是所有路径贡献之和:
残差连接、权重共享和循环网络中的参数复用都依赖这个规则。只保留某一条路径,或让后到达的贡献覆盖先前贡献,都会漏掉总导数的一部分。
第六步是得到所有可训练参数的梯度。
逆拓扑遍历完成后,每个参与当前损失计算且要求求导的参数 θ 都会获得:
这个梯度描述当前参数点附近,参数发生微小变化时损失的一阶变化:
到这里,反向传播的主要职责已经完成。它交付的是梯度,不是新参数。
第七步可以使用有限差分和局部测试验证梯度实现。
有限差分通过参数两侧的前向损失估计数值梯度:
将它与反向传播得到的:
比较,可以抽查局部反向规则、广播归约、路径累加或自定义算子是否正确。
这类检查不是每次训练迭代的必要组成部分,也不能证明所有参数在所有取值上都正确。它是一种开发和诊断工具。
第八步才是优化器更新参数。
最简单的梯度下降规则为:
其中 η 是学习率。更复杂的优化器还可能结合动量、历史平方梯度、权重衰减或其他状态处理当前梯度。
因此,反向传播和优化器的分工是:
- 反向传播计算梯度;
- 优化器读取并变换梯度;
- 优化器决定更新幅度并修改参数。
反向传播不会自动修改参数,也不会决定学习率、动量、梯度裁剪或权重衰减。
即使梯度方向正确,一次实际更新也不保证损失下降。负梯度给出的是无穷小邻域中的局部下降方向:
如果学习率太大,实际更新可能超出局部线性近似成立的区域,使损失震荡甚至上升。梯度提供局部方向和变化率,具体步长仍由优化器负责。
反向传播也不负责判断训练目标是否正确。它只会忠实地计算给定损失函数的导数。如果损失函数没有表达真正目标,那么梯度仍可能在数学上完全正确,但模型会优化错误或不完整的目标。
例如,训练损失持续下降而验证指标没有改善,可能来自:
- 损失只是实际评测目标的代理;
- 模型发生过拟合;
- 训练数据与实际使用数据分布不同;
- 标签或样本权重存在问题;
- 评测指标包含损失没有体现的约束;
- 模型利用了数据中的错误捷径。
这些问题需要由数据、损失设计和独立评测解决,不能仅靠反向传播发现。
反向传播同样不保证找到全局最优解。它计算当前点的局部导数,而长期训练结果还取决于模型结构、初始化、数据顺序、优化器、学习率和随机性。
一次完整训练迭代可以整理为:
- 按训练策略清零或保留旧梯度;
- 使用当前参数执行前向传播;
- 计算标量损失并保存必要中间状态;
- 从 ∂ L/∂ L=1 启动反向;
- 逆序执行各节点的局部向量—雅可比积;
- 在共享变量处累加所有路径贡献;
- 得到各可训练参数的梯度;
- 必要时执行反缩放、非有限值检查或梯度裁剪;
- 优化器读取梯度并更新参数;
- 使用损失和独立评测判断训练是否真正改善目标。
三方职责可以概括为:
评测则负责判断模型在真正关心的标准上是否变好。
因此,反向传播的准确职责是:给定当前前向计算图、必要中间状态和最终标量损失,高效计算该损失对所有需要求导变量的梯度。它连接“损失已经算出”和“优化器可以行动”,但不替代目标设计、参数更新策略或效果评估。
13概念依赖与延伸学习路线
这一节给出反向传播的概念依赖与后续学习路线。理解反向传播之后,还需要继续回答五类问题:从什么目标开始求导、梯度如何转化为更新、深层路径为什么会削弱梯度、网络结构如何改善传播,以及大规模训练怎样权衡计算、通信与内存。
第一条依赖是损失函数。反向传播必须从一个明确的输出开始,训练中通常是标量损失:
反向传播计算的是:
因此,在讨论梯度之前必须先回答:究竟在对什么标量求导?
损失函数规定模型优化的目标。不同损失函数会产生不同的梯度,即使模型、参数和数据完全相同,更换损失后,参数收到的训练信号也可能发生变化。
损失函数与反向传播之间的关系是:
如果损失没有准确表达真实任务目标,反向传播仍然可以得到数学上完全正确的梯度,但模型会高效地优化错误或不完整的目标。因此,梯度正确不等于目标正确。
第二条依赖是梯度下降。反向传播得到:
之后,参数还没有发生变化。优化器需要把梯度转换成具体更新。最简单的梯度下降规则是:
其中 η 是学习率。
梯度提供当前参数点附近的方向和变化率,学习率决定实际步长。如果步长过小,训练可能进展缓慢;如果步长过大,参数可能越过局部下降区域,使损失震荡或发散。
在 mini-batch 训练中,当前梯度通常只是总体目标梯度的随机估计:
不同批次会产生不同梯度,因此更新中带有采样噪声。后续学习梯度下降时,需要继续理解:
- 学习率怎样控制稳定性和收敛速度;
- 批量大小怎样影响梯度噪声;
- 动量怎样整合多个步骤的梯度方向;
- 梯度裁剪怎样限制异常更新;
- 权重衰减怎样改变实际优化目标或更新规则。
反向传播负责把方向算出来,优化器负责决定如何使用这个方向。
第三条延伸是梯度消失与梯度爆炸。深层网络中的梯度由许多局部雅可比连续相乘。
设网络各层为:
则损失对早期状态 h0 的梯度包含:
如果每个局部导数在相关方向上都将梯度缩小,例如每层大约乘以 0.5,经过 n 层后,梯度尺度大约变成:
这个值会随深度指数下降,导致早期层几乎收不到有效训练信号。这就是梯度消失的一种直观来源。
反过来,如果连续雅可比在相关方向上的缩放因子长期大于 1,例如每层约乘以 1.5,梯度尺度可能增长为:
从而产生梯度爆炸。
在张量情形下,不能只看雅可比中的单个元素,还要关注雅可比对不同向量方向的缩放作用。某些方向可能被压缩,另一些方向可能被放大。
梯度消失并不表示反向传播漏掉了路径。它可能正是链式法则正确计算出的结果,只是连续局部变换使信号变得极小。理解这一问题需要继续学习激活函数、初始化、归一化、雅可比谱和深层网络的数值稳定性。
第四条延伸是残差连接。普通深层路径要求梯度连续经过所有非线性变换,而残差块增加了一条恒等捷径:
它对输入的导数为:
所以反向梯度为:
这里的 I 来自恒等路径。梯度不仅可以经过变换分支 F,还可以沿捷径直接传回输入。
从计算图角度看,残差连接正是分叉路径求和:
即使变换分支中的局部雅可比明显缩小了梯度,恒等分支仍能提供一条更短、更直接的传播路径。
残差连接没有绕过链式法则,也不能保证梯度绝不消失。它通过改变计算图结构,增加了一条导数为恒等映射的路径,从而改善深层网络中的信息与梯度传播。
第五条延伸是分布式训练中的内存权衡。训练内存通常包括:
- 模型参数;
- 前向激活;
- 参数梯度;
- 优化器状态;
- 临时计算空间。
当这些内容无法放入单个设备时,可以对不同对象采用分片、重算或跨设备调度。
参数可以在每台设备上完整复制,也可以分片保存。完整复制使每个设备都能直接访问全部参数,但每台设备都承担完整模型内存;参数分片能降低单设备占用,却需要在计算时进行额外通信。
梯度也可以在设备之间同步、求和、平均或分片。数据并行中,不同设备使用不同数据计算局部梯度,然后聚合为与整体批次一致的梯度。聚合规则必须与损失的求和或平均方式保持一致。
优化器状态通常与参数规模成比例,而且一个参数可能对应多份历史统计量。将优化器状态分片到多个设备,可以避免每个设备都保存完整副本。
激活则主要随批量大小、序列长度、网络深度和特征宽度增长。减少激活内存的方法包括:
- 使用梯度检查点,在反向时重算部分前向;
- 调整微批次大小,降低单次峰值内存;
- 进行流水线划分,让不同设备负责不同网络层;
- 沿张量或序列维度切分计算和激活;
- 在计算、通信和存储之间安排更复杂的调度。
这些方法共同面对一个基本权衡:
分布式训练没有改变反向传播的数学规则。每个节点仍执行局部向量—雅可比积,共享变量仍累加所有路径贡献。改变的是参数、激活、梯度和优化器状态存放在哪里,以及它们何时计算和如何跨设备汇总。
本节的过关标准可以用一个小型计算图验证。设:
前向传播依次计算 a、b、s、e 和 L,并保存反向需要的中间值。
反向从:
开始。经过半平方节点:
经过减法节点:
加法节点把上游梯度分别送给 a 和 b:
平方路径传回 w 的贡献为:
直接路径传回 w 的贡献为:
两条路径在共享参数 w 处相加:
这个例子同时检查三项核心能力:
- 能把复合函数拆成局部计算图;
- 能从损失种子 1 开始逐节点反传;
- 能在共享变量处正确累加分叉路径的梯度。
最后,还需要清楚说明训练中的三方分工:
损失函数回答“什么结果更好”,反向传播回答“当前参数怎样影响这个目标”,优化器回答“根据这些梯度实际迈多大一步”。
真正理解反向传播,不只是能够写出链式法则,还要知道它从什么目标出发、如何处理深层路径和分叉结构、结果怎样交给优化器,以及大规模训练系统如何为同一套数学规则安排计算、通信和内存。
| 方向 | 接下来读 | 关键问题 |
|---|---|---|
| 反向的起点 | 损失函数 | 我们究竟在对什么标量求导? |
| 梯度怎样被使用 | 梯度下降 | 方向算对后,步长和噪声怎样影响更新? |
| 深层路径为何变弱 | 梯度消失 | 连续雅可比相乘为何导致指数缩放? |
| 结构怎样帮助反传 | 残差连接 | 恒等捷径怎样提供更短梯度路径? |
| 内存如何权衡 | 分布式训练 | 激活、梯度和优化器状态怎样分片或重算? |
- Automatic Differentiation in Machine Learning: a Survey:前向/反向模式自动微分与复杂度。
- Deep Learning — Deep Feedforward Networks:计算图与反向传播。
- PyTorch Autograd Mechanics:实际自动微分图、保存张量与梯度语义。
- Training Deep Nets with Sublinear Memory Cost:激活重计算的内存—算力权衡。
计算图、手算、表格和调试流程均为本项目原创组织。