ARTICLE DETAIL

资讯详情

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

反向传播计算顺序详解:手推三层网络梯度传播

反向传播计算顺序详解:手推三层网络梯度传播 1. 顺着算还是反着算反向传播到底在算什么东西很多人学反向传播第一反应是去看那堆梯度公式或者直接打开PyTorch调一个loss.backward()就算完事。但真到了自己手写网络、自己实现自定义算子或者要排查梯度异常的时候最先卡住你的往往不是公式本身而是“这一层该先算哪一步、后算哪一步”。反向传播的计算顺序听起来像个不起眼的问题实际上决定了你的梯度是正确还是错乱。先说结论反向传播不是一个笼统的“从后往前算”就能概括的过程它是一套严格按依赖关系展开的链式法则求导流程。具体到某一层你要先算这一层的输入残差也就是传给前一层的梯度再算这一层的参数梯度而整个网络的反向过程则是从损失函数出发依据正向传播时的计算路径一步步把残差从输出层传回输入层。为了把顺序讲透我在这篇文章里统一用一个三层全连接网络做例子手推一遍完整过程。这个网络结构是输入层2个神经元隐藏层3个神经元激活函数用ReLU输出层1个神经元激活函数用Sigmoid损失函数均方误差MSE为什么选这个例子因为它足够简单每一层的计算都能手算验证但又包含了“线性变换激活函数损失函数”这些最常见的组件完全够用来说明计算顺序的核心逻辑。1.1 一个三层小网络的定义先明确符号。输入向量是(x[x_1, x_2]^T)隐藏层的权重矩阵是(W^{(1)})形状是3×2偏置是(b^{(1)})形状是3×1。隐藏层的线性输出是(z^{(1)}W^{(1)}xb^{(1)})经过ReLU激活后得到(a^{(1)}\max(0, z^{(1)}))。输出层的权重是(W^{(2)})形状是1×3偏置是(b^{(2)})形状是1×1。输出层的线性输出是(z^{(2)}W^{(2)}a^{(1)}b^{(2)})经过Sigmoid后得到(\hat{y}\sigma(z^{(2)}))。损失函数是(L\frac{1}{2}(y-\hat{y})^2)。这里用1/2是为了求导后系数整洁不影响顺序。一次完整的训练迭代包含两个阶段正向传播按“输入→隐藏层→输出层→损失”的顺序计算所有中间值反向传播则按完全相反的依赖顺序从损失出发“先输出层→再隐藏层→最后输入层”逐层计算梯度。1.2 反向传播本质是链式法则按层展开为什么说“按层展开”是关键因为链式法则本身只描述了“梯度沿着计算路径相乘”但它没有明确告诉你先乘哪个、后乘哪个。而实际工程中你不可能把所有路径上的梯度一次性全部算出来——那样内存会爆炸计算量也会重复。你只能一层一层地传。具体来说反向传播维护的核心变量是“残差”也叫误差项记为(\delta^{(l)})表示损失对第(l)层线性输出(z^{(l)})的偏导数。这个残差的计算是逐层递推的你只有先算出(\delta^{(l1)})才能算出(\delta^{(l)})。而每层的参数梯度对(W)和(b)的偏导又依赖于该层的输入激活值和残差。这就是整个反向传播计算顺序的骨架先算后面层的残差再算前面层的残差参数梯度与当前层残差同时算但必须先有当前层残差才能算参数梯度。1.3 中间变量与残差的命名习惯为了后续推导不混乱我按工程惯例定一套命名。每个层有两个关键变量一个是“线性输出”也叫logits一个是“激活输出”。隐藏层分别叫(z^{(1)})和(a^{(1)})输出层叫(z^{(2)})和(\hat{y})。反向传播时我们关心的是(\delta^{(2)}\frac{\partial L}{\partial z^{(2)}})和(\delta^{(1)}\frac{\partial L}{\partial z^{(1)}})。这套命名和你将来读源码时遇到的cache、grad、delta等变量是一一对应的。在深度学习框架里正向传播过程中保存下来的中间值就是为了反向传播时按依赖顺序取用。2. 计算图是理解顺序的第一性原理如果你只记住一个工具来理解反向传播顺序那就是计算图。计算图本质上是把一次前向计算过程画成一个有向无环图每个节点是一个操作或者一个变量每条边是数据流向。反向传播的顺序完全由这个图的拓扑序决定——从输出节点开始沿反向边逐节点计算梯度。有了计算图你就不需要背公式了只要看数据流向就知道下一步该算什么。这是我看过的最可靠的理解方式尤其是处理复杂网络结构时。2.1 把矩阵运算拆成一条条图上的边还是用上面的网络举例。一次完整的前向传播对应计算图上的一条条边x → 线性层1 → z(1) → ReLU → a(1) → 线性层2 → z(2) → Sigmoid → ŷ → MSE → L把矩阵运算拆成这种链式结构后你会清晰看到(z^{(2)})依赖于(a^{(1)})所以反向传播时只有先算出(\frac{\partial L}{\partial z^{(2)}})和(\frac{\partial L}{\partial a^{(1)}})才能继续往前传。每一步的输入都是“上一步已经算好的梯度”。这种拆解还有一个好处——你把“对某层的梯度”拆成了“对线性输入的残差”和“对激活输入的梯度”两级。计算顺序上网络每一层内部也是先算对(z)的残差再通过它算对(a)的梯度前一层的残差再依赖这个梯度。这正是“逐层递推”四个字的具体含义。2.2 手动展开五条边的反向顺序我用上面的例子把反向传播需要计算的所有偏导列出来按依赖顺序排序损失对输出(\hat{y})的偏导(\frac{\partial L}{\partial \hat{y}} \hat{y} - y)Sigmoid层的反向得到(\delta^{(2)}\frac{\partial L}{\partial z^{(2)}}\frac{\partial L}{\partial \hat{y}}\cdot \sigma(z^{(2)}))线性层2的反向先得到(\frac{\partial L}{\partial a^{(1)}}\delta^{(2)}\cdot W^{(2)})再得到参数梯度(\frac{\partial L}{\partial W^{(2)}}\delta^{(2)}\cdot (a^{(1)})^T)(\frac{\partial L}{\partial b^{(2)}}\delta^{(2)})ReLU的反向得到(\delta^{(1)}\frac{\partial L}{\partial z^{(1)}}\frac{\partial L}{\partial a^{(1)}}\odot \mathbb{1}[z^{(1)}0])线性层1的反向得到(\frac{\partial L}{\partial x}\delta^{(1)}\cdot W^{(1)})参数梯度(\frac{\partial L}{\partial W^{(1)}}\delta^{(1)}\cdot x^T)(\frac{\partial L}{\partial b^{(1)}}\delta^{(1)})看到没有这五步的顺序是严格线性的第3步依赖第2步第4步依赖第3步第5步依赖第4步。你不可能先算第5步因为(\delta^{(1)})还没影。这就是“反向传播计算顺序”的核心含义每一步的输入是上一步的输出环环相扣顺序错了就全错了。2.3 为什么顺序错了结果就错有读者可能会想反正都是链式法则我把所有偏导先求出来再统一相乘不行吗数学上当然可以但工程上不行。原因有三点第一中间变量没算出来后面的链式乘法根本无从开始。比如你没有先算(\frac{\partial L}{\partial a^{(1)}})就不知道怎么把这个梯度继续往前传给ReLU层。第二链式法则里某些梯度是矩阵乘法顺序反了维度直接对不上。比如(\frac{\partial L}{\partial a^{(1)}} \delta^{(2)} W^{(2)})这个式子如果调换乘法顺序变成(W^{(2)}\delta^{(2)})维度可能还是能乘但数值就完全错位了。第三内存效率问题。实际训练时你不会把整条计算路径上的所有梯度一次性算好存着而是用“边算边丢”的方式只保留当前层反向需要的残差缓存。如果顺序乱了缓存就接不上数值自然就飘了。3. 正向顺序和反向顺序像穿衣服和脱衣服上一节从计算图的角度理清了全局顺序这一节我把手推过程完整走一遍。我的体会是正向传播像穿衣服先穿内层再穿外层反向传播像脱衣服先脱外层再脱内层。这个类比虽然简单但能准确反映依赖关系——你脱外套时不需要先解开内衣扣子但你脱内衣时必须先把外套脱掉。3.1 逐层手推一整套反向步骤为方便手算我取一组具体数值。输入(x[1.0, 2.0]^T)真实标签(y0.5)。初始化权重(W^{(1)}\begin{bmatrix}0.1 -0.2 \ 0.3 0.4 \ -0.5 0.6\end{bmatrix})(b^{(1)}[0.05, -0.05, 0.1]^T)(W^{(2)}[0.2, -0.3, 0.5])(b^{(2)}[0.02])第一步正向传播隐藏层线性输出(z^{(1)}_10.1\times1(-0.2)\times20.05-0.25)(z^{(1)}_20.3\times10.4\times2(-0.05)1.05)(z^{(1)}_3-0.5\times10.6\times20.10.8)ReLU激活后(a^{(1)}_10)因为-0.25被置0(a^{(1)}_21.05)(a^{(1)}_30.8)输出层线性输出(z^{(2)}0.2\times0(-0.3)\times1.050.5\times0.80.020.105)Sigmoid激活(\hat{y}\sigma(0.105)\approx0.5262)损失(L\frac{1}{2}(0.5-0.5262)^2\approx0.000343)第二步反向传播从输出层开始先算损失对(\hat{y})的梯度(\frac{\partial L}{\partial \hat{y}}\hat{y}-y0.0262)Sigmoid的导数(\sigma(z)\sigma(z)(1-\sigma(z)))。所以(\delta^{(2)}0.0262\times0.5262\times(1-0.5262)\approx0.006528)第三步反向传播到线性层2即输出层的线性变换先算损失对隐藏层激活值的梯度(\frac{\partial L}{\partial a^{(1)}}\delta^{(2)}\cdot W^{(2)}0.006528\times[0.2,-0.3,0.5][0.001306,-0.001958,0.003264])注意顺序是(\delta^{(2)})在前、(W^{(2)})在后这样乘出来的结果形状才是1×3正好对应隐藏层的3个神经元。再算线性层2的参数梯度(\frac{\partial L}{\partial W^{(2)}}\delta^{(2)}\cdot (a^{(1)})^T0.006528\times[0,1.05,0.8]^T)([0, 0.006855, 0.005222])严格说对每个权重分量单独求(\frac{\partial L}{\partial b^{(2)}}\delta^{(2)}0.006528)第四步反向传播到ReLU层ReLU的导数是分段常数(z0)时为1(z\leq0)时为0。所以(\delta^{(1)}_10.001306\times\mathbb{1}[-0.250]0)(\delta^{(1)}_2-0.001958\times1-0.001958)(\delta^{(1)}_30.003264\times10.003264)第五步反向传播到线性层1权重梯度(\frac{\partial L}{\partial W^{(1)}}\delta^{(1)}\cdot x^T\begin{bmatrix}0 \ -0.001958 \ 0.003264\end{bmatrix}\times[1,2]\begin{bmatrix}0 0 \ -0.001958 -0.003916 \ 0.003264 0.006528\end{bmatrix})偏置梯度(\frac{\partial L}{\partial b^{(1)}}[0, -0.001958, 0.003264]^T)到这里反向传播全部梯度都算完了。整个过程严格的顺序是从输出端(\hat{y})开始依次是Sigmoid残差→输出层参数梯度→激活值梯度→ReLU残差→隐藏层参数梯度。每一步都建立在前一步的基础上没有任何跳步的可能。3.2 维度检查顺序推导的安保网上面的手推过程如果你跟着算一遍会发现一个特别有用的副产品——每一步都可以通过维度来检验是否正确。这是我在实际写代码时最依赖的校验手段。以(\frac{\partial L}{\partial a^{(1)}}\delta^{(2)}\cdot W^{(2)})为例(\delta^{(2)})形状是1×1(W^{(2)})形状是1×3乘出来是1×3正好对应隐藏层3个神经元的梯度。如果你不小心把顺序写反了变成(W^{(2)}\cdot\delta^{(2)})虽然1×3乘1×1维度上也能乘但语义上就变成了“把输出层权重按输入维度缩放”结果完全不对。更明显的例子是(\frac{\partial L}{\partial W^{(1)}}\delta^{(1)}\cdot x^T)。(\delta^{(1)})是3×1(x^T)是1×2乘出来是3×2正好和(W^{(1)})的形状一致。如果你写成(x^T\cdot\delta^{(1)})那是1×2乘3×1维度都不匹配直接报错。所以维度检查不仅是验证手段更是在复杂网络里推导正确计算顺序的导航仪。3.3 没有残差缓存反向顺序寸步难行手推完一遍你会发现反向传播需要的输入只有两类一是后一层的残差二是当前层正向传播时的输入值比如(z^{(1)})和(a^{(1)})。这就是为什么所有深度学习框架在正向传播时都会“顺手”保存一批中间变量PyTorch的ctx.save_for_backward、TensorFlow的tape记录本质干的是同一件事把反向传播必需的残差缓存按依赖顺序存好留给反向阶段取用。具体到我们的例子反向时需要的缓存有(z^{(1)})ReLU反向计算(\delta^{(1)})要用判断哪些神经元处于激活状态(a^{(1)})输出层权重梯度计算要用(x)隐藏层权重梯度计算要用(z^{(2)})Sigmoid反向要用这个列表非常重要。如果你自己写一个自定义层却没有在正向阶段保存这些中间值反向阶段就会因为缺少输入而卡死。很多初学者在这一步翻车——写自定义层时只实现了正向计算忘了save_for_backward结果一跑反向传播就报错。4. 三种最容易在顺序上翻车的反向传播场景基础网络推完我把范围扩大到实际项目中更容易出问题的三种情况。这三种情况各自对应一个“热搜词”也是我平时被问得最多的问题。它们共同的坑都在计算顺序上但翻车方式完全不同。4.1 softmax加交叉熵的反传顺序先看softmax和交叉熵组合。很多教程会把这两层的反向传播分别展开但实际上在工程实现中它们几乎总是合并成一步计算原因就是计算顺序和数值稳定性。单独的softmax反向需要先算出softmax输出的雅可比矩阵再与上游梯度相乘。单独交叉熵对softmax输入求导也有一串公式。但如果按“先softmax反向再交叉熵反向”的顺序分两步做数值误差会非常大而且计算量成倍增加。实际实现中正确顺序是先得到softmax输出概率(p)再计算交叉熵损失然后直接合并求导得到(\delta\frac{\partial L}{\partial z}p-y)其中(y)是one-hot标签。这个合并后的公式极其简洁顺序上跳过了“中间层的雅可比矩阵”这一步。如果你坚持分步算就要先构造softmax的雅可比矩阵再和上游梯度做矩阵乘法最后还要处理交叉熵的系数。不仅效率低而且数值容易飘。这在工程上是个经典的“顺序优化”案例反向传播时能合并的相邻层尽量合并既省计算又稳数值。我举个具体例子。假设softmax输出的概率分布是(p[0.7,0.2,0.1])真实标签是第0类(y[1,0,0])上游梯度损失对softmax输出的梯度是(\frac{\partial L}{\partial p})。如果分两步走你需要计算3×3的雅可比矩阵[ J_{ij}p_i(\delta_{ij}-p_j) ]然后乘以(\frac{\partial L}{\partial p})。但如果直接合并最终残差就是(p-y[-0.3,0.2,0.1])。这两者数学上等价但合并之后的计算量少了几个数量级而且永远不会出现“0×∞”这类数值问题。4.2 BPTT随时间展开的反向顺序循环神经网络的反向传播叫BPTTBackpropagation Through Time它的计算顺序比普通前馈网络更容易搞混因为多了一个时间维度。BPTT的核心思路是把RNN按时间步展开成一个深层的“伪前馈网络”然后在这个展开图上做反向传播。展开后你会发现它和普通网络最大的差别是同一套权重矩阵(W)在多个时间步上被复用了因此梯度会在时间方向上累积来自未来时间步的梯度必须按时间倒序逐层流回。这里最容易犯的顺序错误是先算当前时刻的参数梯度再算沿时间反向的梯度。正确的顺序应该是先算当前时刻的残差(\delta_t)再沿时间步反向递推出(\delta_{t-1},\delta_{t-2},\dots)最后把所有时间步上对同一参数(W)的梯度相加更新。我在初期写BPTT时吃过亏为了“省事”我先在时间步(t)对应的展开层上算了参数梯度然后才递推前一时刻的残差。结果发现当前时刻的参数梯度中含有前一时刻残差项因为(W)既连接了(h_{t-1}\to h_t)也直接连接了输入到输出残差没算出来梯度就不完整导致数值错误。后来我强迫自己严格按“先残差递推后参数累计”的顺序写才稳定下来。4.3 BatchNorm层的前向统计量与反传顺序BatchNorm层的反向传播是最容易让人怀疑“我算的是不是错了”的场景因为它的梯度不仅依赖于当前样本的输入还依赖于整个batch的统计量均值和方差。而统计量是正向传播时先算好的因此反向传播的顺序必须保证先利用当前batch的均值/方差缓存算出对每个样本输入(\hat{x}_i)的残差再通过残差累积出对均值/方差的偏导最后用它们计算对(x_i)的最终残差。这里顺序如果乱了最常见的症状是训练时梯度爆炸或梯度消失特别是在batch size比较小的时候。我自己排查过两次这种情况最后发现都是因为反向实现时先算了“对(x_i)的最终残差”但中间用到了尚未累积完成的“对均值和方差的偏导”导致数值错误。正确顺序我展开一下从上游拿到(\frac{\partial L}{\partial \hat{x}_i})先算(\frac{\partial L}{\partial \sigma_B^2}\sum_i\frac{\partial L}{\partial \hat{x}_i}\cdot(x_i-\mu_B)\cdot\frac{-1}{2}(\sigma_B^2\epsilon)^{-3/2})再算(\frac{\partial L}{\partial \mu_B}\sum_i\frac{\partial L}{\partial \hat{x}_i}\cdot\frac{-1}{\sqrt{\sigma_B^2\epsilon}})最后算(\frac{\partial L}{\partial x_i}\frac{1}{\sqrt{\sigma_B^2\epsilon}}\cdot(\frac{\partial L}{\partial \hat{x}_i}-\frac{1}{m}\frac{\partial L}{\partial \mu_B}-\frac{x_i-\mu_B}{m}\cdot\frac{\partial L}{\partial \sigma_B^2}))前两步必须先算因为第四步的最终残差同时依赖这两项。如果你先算第四步前面两项还没累积梯度自然就错了。5. 实训跑通一个三层网络并检查每一步中间值理论讲得再多不如直接跑一段可验证的代码。我在这里用Python手写一个极简的三层网络反向传播不加任何自动微分库让你能直观看到每一步的中间值到底长什么样以及顺序错位会带来什么后果。5.1 核心代码结构与残差缓存import numpy as np # 正向传播 x np.array([1.0, 2.0]) y np.array([0.5]) W1 np.array([[0.1, -0.2], [0.3, 0.4], [-0.5, 0.6]]) b1 np.array([0.05, -0.05, 0.1]) W2 np.array([[0.2, -0.3, 0.5]]) b2 np.array([0.02]) # 正向 z1 W1 x b1 a1 np.maximum(0, z1) z2 W2 a1 b2 y_hat 1 / (1 np.exp(-z2)) loss 0.5 * (y - y_hat) ** 2 print(正向传播中间值) print(fz1 {z1}) print(fa1 {a1}) print(fz2 {z2}) print(fy_hat {y_hat}) print(floss {loss})这段代码的输出结果和我们在第3节手推的一模一样。你看到z1、a1、z2、y_hat这些中间值就是后续反向传播的“残差缓存”。如果你要写一个自定义层这些值就必须在正向阶段被保存下来。5.2 反向传播的分步实现与打印接下来是反向传播的完整实现我特意把每一步都拆开打印# 反向传播第一步损失对 y_hat 的梯度 dL_dy_hat y_hat - y # 反向传播第二步Sigmoid 残差对 z2 的梯度 sig_deriv y_hat * (1 - y_hat) delta2 dL_dy_hat * sig_deriv # 反向传播第三步线性层2的梯度 dL_da1 delta2 * W2 # 形状 (1,3) dL_dW2 delta2.T a1.reshape(1, -1) # 形状 (1,3) dL_db2 delta2 # 反向传播第四步ReLU 反向 delta1 dL_da1 * (z1 0) # 形状 (1,3)注意是逐元素乘法 # 反向传播第五步线性层1的梯度 dL_dW1 delta1.reshape(3, 1) x.reshape(1, 2) # 形状 (3,2) dL_db1 delta1 print(反向传播中间值) print(fdL_dy_hat {dL_dy_hat}) print(fdelta2 {delta2}) print(fdL_da1 {dL_da1}) print(fdL_dW2 {dL_dW2}) print(fdL_db2 {dL_db2}) print(fdelta1 {delta1}) print(fdL_dW1 {dL_dW1}) print(fdL_db1 {dL_db1})你运行这段代码时会看到delta1的输出有个很有趣的现象第一个隐藏神经元的残差是0因为它在正向传播时z1是负数ReLU直接把它压死了。这正是反向传播计算顺序中比较关键的一环——ReLU的反向必须知道正向时是激活还是抑制状态而这些信息就只能从正向缓存的z1里拿。5.3 顺序对但数值不稳该检查什么如果代码跑出来梯度数值异常比如NaN、Inf、或者某个梯度特别大而你的计算顺序确认无误那大概率是以下三种情况第一数值溢出。常见于sigmoid或softmax的中间值。sigmoid函数的输入如果绝对值很大np.exp(-z)会直接变成0或Inf。稳妥做法是在实现时加上数值稳定的分支比如当z 0时用1/(1exp(-z))当z 0时用exp(z)/(1exp(z))。第二学习率过大。在梯度正确的前提下如果参数更新步长太大训练曲线会像过山车一样震荡甚至发散。我的经验是先用1e-3这样的小学习率跑通一轮确认梯度数值合理后再慢慢调大。第三缓存没有正确保存。这是最隐蔽的坑。你自定义的层如果在正向传播时忘了保存某个中间值反向时直接报错或计算出错。检查方法很简单把正向传播的中间值打印出来和反向用到的缓存对一遍确保每个量都在。6. 多年踩坑后我养成的“反传顺序直觉”文章写到这儿基本原理和实例都讲完了。最后分享几条我在实际项目中踩过坑之后总结的经验这些不是课本上写的但比很多理论都更实用。6.1 反向传播时哪一层的参数先更新先说参数更新的顺序。你可能觉得反传算完所有梯度后参数更新顺序反正都一样因为每一层各更新各的。但实际工程中尤其是做梯度裁剪、分层学习率、或者某些优化器实现时参数更新的顺序会影响结果因为更新一层后后续层的梯度计算一般不会受影响如果所有梯度都已经算好的话但如果使用动量类优化器且部分层有共享参数先更新谁就变得很关键。我的经验是严格按反向传播的顺序更新参数也就是输出层先更新、隐藏层后更新。这和梯度计算顺序一致是最符合直觉的做法也不容易引入隐藏的bug。6.2 梯度累积与多任务损失的顺序如果你在做多任务学习或者用梯度累积模拟大batch size顺序问题就更加微妙。常规做法是先对每个任务的损失分别调用反向传播然后把梯度累加到一起最后统一更新参数。关键在于累加的顺序。从数学上看梯度累加满足交换律先加谁后加谁都一样。但从数值精度上看从小梯度往大梯度上累加比从大梯度往小梯度上累加更稳定。所以我的习惯是先计算损失较小的任务梯度再计算损失较大的任务梯度最后统一累加。这样做能减少浮点舍入误差的累积。6.3 最后分享一个调试口诀多年下来我给团队培训时经常说一句话“反向传播的顺序就是正向传播的倒序每前进一步只依赖已经算好的上一步。”如果你在排查梯度问题时没有头绪就按这个口诀把计算图重新画一遍标出每一步的输入和输出然后逐层检查中间值是否符合预期。我靠这个口诀排查过不下十次诡异的梯度问题几乎每次都有效。再补充一个小技巧如果你不想每次手画计算图可以直接打印正向传播时缓存的所有中间值和反向传播时算出的所有梯度放在一张表里逐行对比。哪一行对不上问题就出在哪一行。正向缓存: z1, a1, z2, y_hat 反向梯度: dL_dy_hat, delta2, dL_da1, delta1, dL_dW2, dL_dW1只要这张表里的每一行都符合你手推的数值那说明计算顺序完全正确。如果某一行出现NaN、Inf、或者量级明显不对就从那一行对应的反向步骤开始回溯基本能在十分钟内定位问题。
返回列表