反向传播计算顺序详解:从计算图到PyTorch自动微分
2026/9/24 18:53:39 网站建设 项目流程

反向传播难的不是链式法则,难的是搞明白每一步该算什么、先算什么、后算什么。很多初学者对着公式能看懂,一让自己从头推到尾就卡壳,或者写自定义算子时梯度死活不对。这篇就用最朴素的方式把“反向传播的计算顺序”拆开讲清楚,适合正在啃反向传播的学生、准备面试的候选人,以及写 PyTorch 自定义扩展时对梯度流向拿不准的工程师。

我先把结论放在前面:反向传播的计算顺序,本质上就是计算图的“反向拓扑序”。也就是说,从输出节点开始,沿着计算图逆着正向传播的方向,一步一步往输入端走。每一步计算梯度时,它依赖的上游梯度必须已经算好。这个顺序是唯一的,不能乱。理解了这一点,后面所有细节都会变得很顺。

1. 反向传播到底在算什么

1.1 先用一个最简单的计算图理解顺序

假设有一个复合函数 y = f(g(h(x)))。正向传播时,你会先算 a = h(x),再算 b = g(a),最后算 y = f(b)。反向传播时,目标是求出 dy/dx,还有中间那些参数可能对应的梯度。

反推的顺序其实很固定:

  • 第一步,从输出开始,令 dy/dy = 1,这是反向传播的起点。
  • 第二步,算 dy/db = f'(b)。这里只用到局部导数 f'(b) 和刚才的 dy/dy。
  • 第三步,算 dy/da = (dy/db) * g'(a)。
  • 第四步,算 dy/dx = (dy/da) * h'(x)。

注意这个顺序,你不可能先算 dy/da,因为 dy/da 要用到 dy/db,而 dy/db 还没算出来。每一步都依赖上一步的结果,这就是“反向拓扑序”的含义。

打个比方,正向传播像从起点把一串链子往前甩出去,反向传播则像从链子的末端开始,一节一节把力传回来。如果你从中间开始往回拉,中间那段可能会松脱,根本传不到起点。

1.2 链式法则与梯度累积:顺序问题的起点

当计算图里出现分支时,顺序会更严格,因为还要考虑梯度累加。

举个例子,某个中间节点 a 同时被 b 和 c 两个下游节点使用,而 b、c 最终都影响 y。那根据全导数公式:

dy/da = (dy/db) * (db/da) + (dy/dc) * (dc/da)

也就是说,梯度要从两条路径分别传回来,然后在节点 a 上相加。顺序上,你必须等两条路径的上游梯度都算完,再求和。这也是 PyTorch 里多次调用 backward 时梯度会累加的原因之一,更是网络里参数共享、残差连接、多分支结构梯度计算的底层逻辑。

很多人第一次看自动微分框架的实现,会被“梯度累积”这个概念绕晕。其实它就是链式法则里“求和”二字的工程化表达。只要计算图有分支,反向梯度就一定会在某个节点汇聚,汇聚的方式就是相加。

2. 为什么计算顺序不能乱:反向拓扑序

2.1 依赖关系决定了一切

计算图本身是一个有向无环图,DAG。正向传播时,数据从输入节点流向输出节点,每个节点的输出被后续节点消费。反向传播时,梯度从输出节点流回输入节点,每个节点要计算它对上游节点的梯度,必须是它的下游梯度已经就绪。

如果把计算顺序打乱,比如某个节点的下游梯度还没算,就直接去算上上游梯度,那结果一定是错的,甚至根本没法算。现实中你不会在手工推导时犯这种错,但框架的自动微分引擎必须保证这点。PyTorch 里,每次 forward 都会记录一个 grad_fn 链,backward 时就是按照这个依赖关系反向调用的。

这里有个容易误解的点:反向计算顺序是“从后往前”,但并不是简单地把链式法则倒过来念一遍。因为计算图可能很复杂,有分支、有合并、有共享节点,你得严格按依赖关系走。这个依赖关系,就是反向拓扑排序。

反向拓扑排序的规则很朴素:每次选择一个入度为 0 的节点(在反向图中,入度表示“上游梯度还没来”的依赖数量),计算它的梯度,然后把它从图里移除,更新相邻节点的入度。重复这个过程,直到所有节点都处理完。

2.2 从输出到输入逐层推进:一个通用推导套路

我平时手推梯度,基本按下面四步走,这套路可以用在任何网络上,包括 Transformer:

  1. 画出计算图,标清每个节点是什么运算,记录正向传播时每个中间节点的值。
  2. 从最终的损失节点开始,令损失对自身的梯度为 1。
  3. 沿着计算图逆序处理每个节点,用链式法则计算它对其直接上游节点的梯度。
  4. 如果某个节点有多条下游路径,把它接受到的所有上游梯度求和,得到总梯度,再继续向上传。

第 3 步里那句“用链式法则”,拆开就是:当前节点的上游梯度 × 当前运算对输入的局部梯度。注意这里“上游梯度”特指损失对当前节点输出的梯度,而不是对输入。很多公式写成 dy/dx = dy/du * du/dx,其中 dy/du 就是上游梯度,du/dx 是局部梯度。

这套路看起来简单,但真正动手时容易在形状和转置上卡住。我的建议是:每算一步,先确认要算的梯度是什么形状。比如输入是 (N, D) 的矩阵,参数是 (H, D),那梯度一定也是 (H, D)。形状能帮你快速判断是不是转置放错了。

2.3 中间变量的保存与释放:显存和顺序的纠缠

反向传播要用到正向传播时的中间变量,这些变量也叫激活值。比如 ReLU 需要记住哪些位置是正数,softmax 需要记住归一化后的概率,矩阵乘法需要记住输入矩阵。框架会在 forward 时把必要的东西保存下来,等 backward 时读取。

这些中间变量占显存,而且占得非常多。于是就有各种节省显存的技术,核心思路就是“重新计算”。PyTorch 的梯度检查点(checkpoint)会把一段计算图的中间激活丢掉,反向传播时再重新算一次前向来恢复这些值。恢复之后再按原来的反向顺序继续推梯度。

代价是前向计算会多跑一遍,时间变长,但显存降下来了。我训练大模型时会用 checkpoint_wrapper 包住 transformer 的每一层,配合混合精度,能把 7B 模型塞进单卡训练。理解了计算顺序,你才会明白它为什么是“用时间换显存”:反传本身需要的中间值没有消失,只是被延后到反传时重新生成。

3. 手推一个两层全连接网络

3.1 网络结构与符号约定

空谈顺序没有感觉,直接手推一个实际例子。我们构造一个非常小的网络:

  • 输入 x 是二维向量,假设一个样本 x = [1.0, 0.5]。
  • 隐藏层:z1 = W1 x + b1,激活函数用 ReLU,得到 h1。
  • 输出层:z2 = W2 h1 + b2,接 Softmax 得到预测概率 p,用交叉熵损失。

参数尺寸如下:

  • W1 是 2×2,b1 是二维。
  • W2 是 2×2,b2 是二维。

我随便给一组具体的值,方便你跟着手算验证:

  • W1 = [[0.2, -0.1], [0.3, 0.4]]
  • b1 = [0.1, -0.2]
  • W2 = [[0.5, -0.2], [0.1, 0.6]]
  • b2 = [0.0, 0.1]
  • 标签 y 的 one-hot 是 [1, 0],也就是正确类别是第 0 类。

这个例子很小,你可以拿笔算一遍,也可以直接用 NumPy 验算。后面我列出的每一步都是按照反向传播应有的计算顺序排的。

3.2 正向传播各节点的值

先算正向,顺便把中间值都记录下来,因为反传都要用。

  • z1 = W1 x + b1

    z1[0] = 0.2×1.0 + (-0.1)×0.5 + 0.1 = 0.25 z1[1] = 0.3×1.0 + 0.4×0.5 + (-0.2) = 0.30

  • h1 = ReLU(z1) = [0.25, 0.30](这里恰好全是正数,mask 是 [1, 1])

  • z2 = W2 h1 + b2

    z2[0] = 0.5×0.25 + (-0.2)×0.30 + 0.0 = 0.065 z2[1] = 0.1×0.25 + 0.6×0.30 + 0.1 = 0.305

  • softmax 概率:exp(0.065)≈1.0672,exp(0.305)≈1.3566,总和≈2.4238

    p[0] ≈ 0.4403,p[1] ≈ 0.5597

  • 交叉熵损失:L = -log(p[0]) ≈ 0.8206

记住这些数字,后面会用到。

3.3 反传每一步的计算顺序

反向传播的计算顺序,从损失开始,往前推:

第一步,算损失对 z2 的梯度。Softmax 加上交叉熵有一个非常简洁的梯度形式:

dz2 = p - y

所以 dz2 = [0.4403 - 1, 0.5597 - 0] = [-0.5597, 0.5597]。

这个“p 减 one-hot”的结果直接就是损失对 logits 的梯度,不需要先算 softmax 的雅可比矩阵再乘交叉熵梯度。这一技巧值得记住,很多框架实现里也是这么做的。

第二步,有了 dz2,先算输出层参数的梯度,并向上一层传递。

  • dW2 = outer(dz2, h1),也就是 dz2 的每个分量乘以 h1 的每个分量,得到形状 (2,2)。

    dW2 = [[-0.5597×0.25, -0.5597×0.30], [0.5597×0.25, 0.5597×0.30]] = [[-0.1399, -0.1679], [0.1399, 0.1679]]

  • db2 = dz2 = [-0.5597, 0.5597]

  • dh1 = W2^T dz2

    dh1[0] = 0.5×(-0.5597) + 0.1×0.5597 ≈ -0.2239 dh1[1] = -0.2×(-0.5597) + 0.6×0.5597 ≈ 0.4478

注意这里计算顺序很重要:dh1 是为继续往更上游传播而算的中间梯度,必须在一层内先算好,因为你下一步要用它来计算隐藏层的梯度。

第三步,计算损失对 z1 的梯度。因为 h1 = ReLU(z1),所以 dz1 = dh1 * mask,其中 mask 是 ReLU 前向时记录的指示向量(正数的位置为 1,非正数为 0)。

这个例子里 mask 是 [1, 1],所以 dz1 = [-0.2239, 0.4478]。

如果某个神经元输出为负数,反传梯度会直接变成 0,这就是 ReLU 死亡问题的本质,也解释了为什么反向传播顺序里需要保存 ReLU 的掩码。

第四步,继续算隐藏层参数的梯度,以及损失对输入 x 的梯度。

  • dW1 = outer(dz1, x)

    dW1 = [[-0.2239×1.0, -0.2239×0.5], [0.4478×1.0, 0.4478×0.5]] = [[-0.2239, -0.1120], [0.4478, 0.2239]]

  • db1 = dz1 = [-0.2239, 0.4478]

  • dx = W1^T dz1

    dx[0] = 0.2×(-0.2239) + 0.3×0.4478 ≈ 0.0896 dx[1] = -0.1×(-0.2239) + 0.4×0.4478 ≈ 0.2015

到这里,所有参数的梯度就算完了。你可以看到,整个流程是“输出层 logits 梯度 → 输出层参数梯度 → 隐藏层输出梯度 → 隐藏层预激活梯度 → 隐藏层参数梯度 → 输入梯度”,这个顺序像剥洋葱一样,一层一层往回走。

3.4 权重梯度的累积与代码对应

上面是单个样本的梯度。实际训练时,一个 batch 里有多个样本,梯度会在 batch 维度上求和或求平均。PyTorch 的 loss.backward() 默认把梯度累积到 .grad 里,所以每步优化前要调用 optimizer.zero_grad() 清零,否则上一轮的梯度会累加进来。

此外,如果同一个权重被多个分支使用,梯度还要跨路径累加。比如一个共享的 W 同时作用在两个输入上,那么 dW 应该等于两条路径各自贡献的梯度之和。计算顺序上,框架会先把所有路径的上游梯度求出来,再统一累加到这个参数上。这也解释了为什么手动实现多分支网络时,不能简单地在第一个分支算完就更新参数,要等所有分支的梯度都回来再做优化。

4. 自动微分框架如何决定计算顺序

4.1 动态图与静态图的差异

PyTorch 的动态图机制,每次 forward 都会动态构建一张计算图。Tensor 上的 grad_fn 记录了它是通过什么运算产生的,这些 grad_fn 之间通过 next_functions 连成一张反向图。调用 backward() 时,PyTorch 会从输出节点的 grad_fn 出发,做一次反向拓扑遍历,依次调用每个 grad_fn 的 backward 方法,把梯度传给输入。

TensorFlow 的静态图机制则是先建立完整的计算图,再通过自动微分生成对应的反向子图。由于图是固定的,TensorFlow 可以对反向子图做更多编译优化,比如算子融合、内存复用等。但用户侧看到的效果是一致的:反向计算顺序永远是从 loss 往输入方向走。

这里有一个我经常跟人强调的点:动态图虽然灵活,但每次迭代的图都可能不一样,所以 backprop 的具体顺序依赖本次 forward 路径。PyTorch 的 Debug 工具 traceback 经常看到 grad_fn 链,就是为了让你理解当前这一次反传的顺序。

4.2 以 PyTorch 为例:一次 forward/backward 中发生了什么

下面这段代码很典型:

import torch x = torch.tensor([1.0, 0.5], requires_grad=True) w1 = torch.tensor([[0.2, -0.1], [0.3, 0.4]], requires_grad=True) b1 = torch.tensor([0.1, -0.2], requires_grad=True) w2 = torch.tensor([[0.5, -0.2], [0.1, 0.6]], requires_grad=True) b2 = torch.tensor([0.0, 0.1], requires_grad=True) z1 = x @ w1.T + b1 h1 = torch.relu(z1) z2 = h1 @ w2.T + b2 loss = torch.nn.functional.cross_entropy(z2.unsqueeze(0), torch.tensor([0])) loss.backward()

调用 loss.backward() 时,PyTorch 内部从 loss 的 grad_fn(可能是 NllLossBackward)出发,沿着 next_functions 一直往前遍历。每经过一个节点,就会调用该节点对应的反向函数。比如经过 z2 的 grad_fn(AddBackward),它会计算对 z2 输入的梯度;经过 h1 的 grad_fn(ReluBackward),它会用到前向保存的 mask,计算 dz1。最终,所有 requires_grad=True 的叶子节点(如 w1、b1、w2、b2、x)都会收到梯度。

用代码检查梯度也很简单:

print(x.grad) print(w1.grad) print(w2.grad)

如果你手推的结果和框架给出的梯度对不上,那一定是某一步转置或顺序错了。我经常这样自检。

4.3 中间值保存、释放与 checkpoint 的影响

自动微分框架里,默认前向会保存反向需要的中间张量。比如上面的 h1,它是计算 w2 梯度时需要的,所以 ReLU 会把输出 h1 保存下来。什么时候释放?在对应反向节点计算完毕后,如果这个张量不再被任何后续反向节点使用,框架才可以释放。

这个释放顺序直接决定了显存峰值。你用 torch.cuda.max_memory_allocated() 观察一次大模型训练,会发现显存占用在前向结束时达到峰值,因为前向把所有中间值都保存了。反向开始后,随着各层梯度算完,中间值逐渐释放,显存才慢慢降下来。

gradient checkpointing 的思路是:故意丢弃一部分中间值,反传到该段时再按顺序重算一次前向,得到中间值。这样前向需要保存的中间值数量大大减少,代价是计算时间变长。我在训练长序列模型时经常用 checkpoint,因为它能把显存需求砍掉一半甚至更多。

4.4 自定义 autograd.Function 的梯度顺序

如果你需要自己实现一个算子,比如自定义一个带参数的模块,就必须非常小心 backward 的返回顺序。PyTorch 规定,backward 返回的梯度元组,必须与 forward 的输入参数顺序一一对应。

class MyLinear(torch.autograd.Function): @staticmethod def forward(ctx, x, w, b): ctx.save_for_backward(x, w) return x @ w.T + b @staticmethod def backward(ctx, grad_output): x, w = ctx.saved_tensors grad_x = grad_output @ w # 注意形状匹配 grad_w = grad_output.T @ x grad_b = grad_output.sum(0) return grad_x, grad_w, grad_b

这里 forward 的输入顺序是 x, w, b,所以 backward 必须返回三个梯度,顺序也是 grad_x, grad_w, grad_b。如果你把顺序写反了,PyTorch 不会报错,但梯度会错误地传给别的参数,训练结果会莫名其妙地发散。这种 bug 特别隐蔽,所以我的习惯是先写一个单步测试,用 torch.autograd.gradcheck 验证自定义算子的梯度是否正确。

5. 常见问题与排查技巧实录

5.1 梯度为 NaN、为 0、不更新的典型原因

很多人问,为什么我的 loss 一开始就是 NaN?第一个要查的就是反向传播这条链上的数值稳定性。Softmax 里如果直接算 log(softmax(z)),z 很大时 exp(z) 会上溢,log 里会变成 inf,梯度自然就是 NaN。正确做法是用 log_softmax 再套 NLLLoss,或者在交叉熵计算里做好数值修正。这是计算图表达方式带来的顺序问题:先做一次稳定的前向,反向梯度顺序才能正确且数值可靠。

梯度为 0 的常见原因有三个:ReLU 输出负区间的神经元永远不激活;初始化导致所有激活值过大,softmax 饱和,梯度极接近 0;或者某个中间节点被 detach 了,梯度被硬生生切断。

梯度不更新,先确认三件事:requires_grad 是否打开,loss 是否含有梯度路径,以及 optimizer.zero_grad() 是否调用了。如果忘了 zero_grad,梯度会跨 batch 累积,表现就是 loss 经常突然跳一下。这和反向传播本身的“梯度累积顺序”是直接相关的:你连续 backward 了两次,第一次的梯度还没清零,第二次又加上去,optimizer 拿到的是一个混合了多个 batch 的梯度。

5.2 用 gradcheck 验证梯度计算顺序是否正确

手推梯度容易翻车,尤其是形状复杂的算子。PyTorch 提供一个数值梯度检查工具:

from torch.autograd import gradcheck input = torch.randn(1, 2, dtype=torch.float64, requires_grad=True) w = torch.randn(2, 2, dtype=torch.float64, requires_grad=True) b = torch.randn(2, dtype=torch.float64, requires_grad=True) test = gradcheck(MyLinear.apply, (input, w, b), eps=1e-6, atol=1e-4) print(test)

gradcheck 会用中心差分算数值梯度,再和你实现的 backward 输出的梯度做对比。如果顺序错了,数值梯度能正确反映“输入 e_i 对 loss 的影响”,而你的解析梯度会张冠李戴,两者对不上,gradcheck 就会返回 False。这个工具对我排查自定义算子帮助很大。

注意 gradcheck 对输入类型有要求,通常要用 float64,因为 float32 下数值差分误差会比较大,可能误报。

5.3 共享参数与多分支的梯度累积顺序

权重共享在实际模型里很常见,比如 Siamese 网络,或者某些结构里同一个矩阵在不同时间步被反复使用。这样的参数一次 forward 会经历多次正向路径,反传时它的梯度是各路径梯度之和。PyTorch 会自动累加,但复杂度在于,如果路径的深度不同,梯度到达参数的顺序不同,累加结果可能受到浮点误差影响,但不会影响最终值。真正危险的是在有分支的情况下提前修改参数,这会破坏后续路径的反向计算顺序,导致梯度算错。

所以正确做法永远是:先完成一轮完整的 forward,再完整 backward,不要在一个层刚算完梯度就立刻更新参数。等所有梯度都累积到 .grad 里,调用 optimizer.step(),最后 zero_grad。这个顺序是深度学习训练流程的黄金法则。

5.4 混合精度下的顺序与数值稳定性

混合精度训练时,前向用 FP16,梯度通常用 FP32 累积。如果框架没有做好梯度累加的精度保持,比如在 FP16 里连续加很多个小梯度,可能因为数值范围不够出现精度损失甚至下溢。这也是为什么 AMP(Automatic Mixed Precision)在反向传播的梯度累加处特别小心:梯度暂时存在 FP32 buffer,只在需要时才转成 FP16。

另外,在不同设备或不同 reduction 顺序下,floating point 的求和顺序不同,结果也可能差 1e-7。如果你复现论文时发现 loss 曲线和官方不完全一致,先别急着怀疑自己代码,看看是否是计算顺序带来的浮点误差。只要在合理的误差范围内,就不用管。

6. BPTT:随时间反向传播的计算顺序

6.1 把时间维展开成计算图,顺序就清楚了

循环神经网络 RNN,包括 LSTM、GRU,训练时的核心算法叫随时间反向传播,简称 BPTT。它本质上就是把时间维展开成一个很深的计算图,然后在展开后的图上做标准的反向传播。所以 BPTT 的计算顺序依然是“从后往前”,只不过多了一个时间方向。

展开到第三步的简单 RNN:

  • h1 = tanh(W_h h0 + W_x x1 + b)
  • h2 = tanh(W_h h1 + W_x x2 + b)
  • h3 = tanh(W_h h2 + W_x x3 + b)

每个时刻都会输出一个预测,损失可能是各时刻损失之和。反传时,先算最后一个时刻 h3 的梯度,然后通过 W_h 回传到 h2,在 h2 处还要加上第 2 时刻的输出损失带来的梯度,再继续回传到 h1。整个过程就像一个深度为 T 的前馈网络,只是隐藏层之间共享 W_h。

6.2 一个三步展开的 RNN 手推思路

用 T=3 的例子,如果损失 L = L1 + L2 + L3,那么反向顺序可以描述为:

  1. 计算 L3 对 h3 的梯度,设为 g3。
  2. 把 g3 通过 tanh 的局部导数传到 h3 的输入,再通过 W_h 传到 h2,得到一部分 g2。
  3. 把 L2 对 h2 的梯度加到 g2 上,这就是 h2 的总梯度。
  4. 再通过同样的方式回传到 h1,同时加上 L1 对 h1 的梯度。
  5. 最后,对 W_h 的梯度等于每个时间步贡献之和,对 W_x 的梯度也类似。

可以看到,BPTT 的顺序和前馈网络的逐层回传没有本质区别,只是把“层”换成了“时间步”。在实际实现中,PyTorch 的 LSTM 会为每个时间步创建一次反向节点,所以反向时会逐时间步回溯。

6.3 长序列中的梯度爆炸与截断 BPTT

BPTT 一个臭名昭著的问题是梯度爆炸或梯度消失。原因在于,损失对较早时间步的梯度,需要在时间方向上连乘一长串转移矩阵的导数。如果矩阵的特征值大于 1,梯度指数增长,很快就变成 NaN;如果小于 1,梯度指数衰减,早期时间步几乎学不到东西。

工程上常用的缓解手段有三个:

  • 梯度裁剪:对所有参数梯度的范数设一个上限,比如 max_norm=1.0。如果梯度范数超过这个值,就整体缩放。这个操作要放在 optimizer.step() 之前,顺序不能反。
  • 截断 BPTT:不把整个长序列一次性反传到第一个时间步,而是每 k 步截断一次,只反传最近 k 步。PyTorch 的 LSTM 里通常用 detach 切断 h 的图,这相当于人为控制反传深度。我处理上千长度序列时会设计类似策略,否则显存和时间都吃不消。
  • 门控结构:LSTM 和 GRU 通过遗忘门、输入门控制梯度流,能在一定程度上缓解梯度消失。这也是为什么实际任务中很少用朴素 RNN。

从顺序角度看,梯度裁剪的位置也有讲究。应该在整个反向传播完成后,对所有参数梯度统一做范数裁剪,再更新参数。如果边反传边裁剪,等价于改变了梯度量级,可能影响收敛。

最后分享一个我自己的习惯:每次实现新模型,我都会先手动推一遍至少一层网络的正向和反向顺序,用一个小矩阵把每步梯度形状写下来,再和框架输出对照。过程很枯燥,但真的能帮你建立对“反向传播计算顺序”的直觉。哪怕后来用惯了 PyTorch 的自动微分,一旦遇到梯度异常,我还是会回到这个最朴素的手推流程,沿着计算图一步步查,通常很快就能定位问题。

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询