
反向传播和梯度下降这两个词几乎每个接触大模型的人都在各种文章里见过但真正能把它们串起来讲清楚数据到底怎么在神经网络里流动、参数到底怎么被更新的人并不多。我见过太多人背下了链式法则求导和沿梯度反方向更新这两句话可一旦被问到为什么学习率设大了会震荡、设小了又不动、梯度累积到底在解决什么问题、为什么大模型训练里几乎没人用纯批量梯度下降就答不上来了。这篇内容就是冲着这些具体问题来的我会从计算图的视角把反向传播拆开再把梯度下降的几种变体和实际训练中的调参经验讲透适合已经了解神经网络基本结构、想真正搞懂训练底层逻辑的读者也适合正在做大模型微调、被学习率和显存问题折磨的从业者。1. 从一次前向计算说起数据是怎么流过网络的要讲反向传播必须先讲清楚前向传播在算什么。很多人一上来就啃链式法则的公式结果脑子里没有数据流的画面公式就成了纯粹的符号游戏。我习惯先建立一个具体的、小到可以手算的例子把前向过程走一遍后面反向的时候才有对照。1.1 一个三层网络的完整前向过程假设我们有一个极简的网络输入层两个节点中间一个隐藏层两个节点输出层一个节点。输入是 x10.5、x20.8权重和偏置我随便给一组初始值激活函数用 Sigmoid。前向传播做的事情本质就是加权求和再过激活函数一层一层往前推。隐藏层第一个节点的计算是这样的z1 w11x1 w21x2 b1假设 w110.1、w210.2、b10.05那么 z1 0.10.5 0.20.8 0.05 0.05 0.16 0.05 0.26。过 Sigmoid 得到 a1 1/(1e^(-0.26)) ≈ 0.5646。第二个节点同理算出 a2。然后输出层再拿 a1、a2 做一次加权求和加激活得到最终的预测值。这个过程看起来平平无奇但关键在于前向传播的每一步计算都在构建一张计算图。z1 依赖 x1、x2、w11、w21、b1a1 依赖 z1输出依赖 a1、a2。这张图是有向无环的每个节点记录了它由谁算出来。反向传播之所以能高效工作就是因为这张图被缓存下来了求导的时候不需要重新推导整个表达式只需要沿着图的边往回走。提示理解计算图是理解反向传播的前提。如果你把网络看成一堆孤立的公式反向传播就是天书把它看成一张有依赖关系的图反向传播就是在这张图上做一次有序的回溯。1.2 为什么必须缓存中间结果这里有个容易被忽略的工程细节前向传播时缓存下来的中间值比如 z1、a1在反向传播时会被反复用到。以 Sigmoid 为例它的导数可以写成 a*(1-a)也就是说只要前向时把激活后的输出 a 存下来反向时就不用重新算一遍 Sigmoid。这就是为什么深度学习框架在训练时显存占用远大于推理——训练要保存大量中间激活值供反向使用。我实测过一个对比同一个模型推理时显存占用可能只有 2GB但开启训练、保存全部激活值后能飙到 8GB 以上。这也是梯度检查点Gradient Checkpointing技术存在的理由——用计算换显存只存部分中间值其余的在反向时重算。这个取舍在大模型微调里非常常见后面讲梯度累积时还会提到它。1.3 损失函数把预测变成一个有方向的标量前向传播最后输出的是预测值但预测值本身没法直接告诉我们错了多少。损失函数的作用就是把这个预测和真实标签之间的差距压缩成一个标量。回归任务常用均方误差 MSE分类任务常用交叉熵。为什么强调标量因为反向传播的起点必须是一个标量。如果损失是个向量我们就没法定义梯度的方向。交叉熵损失在分类任务里之所以流行除了它和最大似然估计的天然联系还有一个实用原因它配合 Softmax 使用时导数形式极其简洁就是预测概率减去真实标签计算量小且数值稳定。这个细节在推导反向传播时会体现出来也是为什么框架里 Softmax 和交叉熵经常被合并成一个算子实现。2. 反向传播的本质链式法则在计算图上的有序回溯现在进入正题。反向传播经常被讲得很玄其实它的数学内核只有一个链式法则。难点不在数学而在于怎么组织计算顺序让求导不重复、不遗漏、效率最高。2.1 链式法则复合函数求导的拆解逻辑链式法则说的是如果 y 是 u 的函数u 又是 x 的函数那么 y 对 x 的导数等于 y 对 u 的导数乘以 u 对 x 的导数。写成公式就是 dy/dx (dy/du) * (du/dx)。放到神经网络里损失 L 对某个深层权重的导数需要经过一长串中间变量。比如 L 依赖输出 aa 依赖 zz 依赖权重 w那么 dL/dw (dL/da) * (da/dz) * (dz/dw)。每一段都是简单的局部导数乘起来就是全局导数。这就是为什么反向传播能处理任意深度的网络——再深的网络求导也只是把一串局部导数连乘起来。但这里有个陷阱连乘会导致数值问题。如果每一段的导数都小于 1乘几十上百次之后梯度就趋近于 0这就是梯度消失如果都大于 1就会指数级放大这就是梯度爆炸。Sigmoid 的导数最大值只有 0.25深层网络里连乘几次就衰减得厉害这也是后来 ReLU 取代 Sigmoid 成为主流激活函数的核心原因之一。2.2 从后往前为什么反向比正向更高效一个自然的疑问是既然链式法则可以正着用也可以反着用为什么一定要从后往前算关键在于复用。假设损失 L 对某一层输出的梯度已经算出来了那么这一层所有参数的梯度都可以基于这个梯度直接算不需要各自从头推导。如果从前往后算每算一个参数的梯度都要重新走一遍到损失的路径计算量会爆炸。反向传播的精髓就是从损失出发逐层往回传梯度每一层的梯度只算一次然后分发给所有依赖它的参数。我用一个具体数字说明效率差异。一个 10 层的网络如果从前往后逐个参数求导每个参数都要走一遍完整路径复杂度是层数的平方级别而反向传播每个节点只访问一次复杂度是线性级别。层数越深差距越夸张。这就是为什么反向传播在 1986 年被系统化提出后直接推动了神经网络的复兴——它把训练成本从不可接受降到了可以接受。2.3 梯度是怎么一层层传回去的具体到每一层反向传播做的事情可以拆成三步。第一步接收来自上一层的梯度 dL/d_output。第二步用这个梯度乘以本层激活函数的局部导数得到 dL/d_z。第三步用 dL/d_z 分别算出对权重、偏置的梯度以及要继续往前传的 dL/d_input。以全连接层为例如果前向是 output input W b那么反向时dL/dW input^T dL/doutputdL/db dL/doutput 按批次求和dL/dinput dL/doutput W^T。这三个公式是框架里全连接层反向实现的核心几乎所有深度学习库的底层都是这么写的。注意 dL/dinput 是继续往前传的部分它让梯度能穿过这一层到达更浅的层。注意dL/db 需要对批次维度求和因为偏置在每个样本上是共享的。这个细节在手动实现反向传播时特别容易漏漏了会导致偏置梯度偏大训练不稳定。2.4 一个手算例子把梯度传一遍还是用前面那个三层网络。假设损失对输出层激活后的梯度是 0.3输出层用的是 Sigmoid激活前值 z_out 对应的激活后值是 a_out。那么 dL/dz_out 0.3 * a_out * (1 - a_out)。假设 a_out 0.6则 dL/dz_out 0.3 * 0.6 * 0.4 0.072。接着算输出层权重梯度dL/dw_out dL/dz_out * a_hidden。假设 a_hidden 0.56则 dL/dw_out 0.072 * 0.56 ≈ 0.0403。然后继续往前传dL/da_hidden dL/dz_out * w_out再乘以隐藏层 Sigmoid 的导数得到 dL/dz_hidden以此类推。手算一遍之后你会发现整个过程没有任何魔法就是老老实实地按链式法则乘下去。框架做的事情无非是把这套流程自动化、向量化、并行化。理解了这一点再看 PyTorch 里的loss.backward()你就知道那一行代码背后发生了什么。3. 梯度下降家族从批量到随机再到小批量梯度算出来了接下来就是用它更新参数。梯度下降的核心思想极其朴素沿着梯度的反方向走一小步因为梯度指向的是损失上升最快的方向。但走多大一步、用多少数据算梯度这两个问题衍生出了一整个家族的方法。3.1 批量梯度下降稳但慢批量梯度下降BGD每次更新都用全部训练数据算梯度。优点是梯度方向准确、更新稳定、损失曲线平滑。缺点是每次更新都要遍历整个数据集数据量大时慢到无法接受。我做过一个粗略估算假设训练集有 100 万条样本模型前向加反向一次处理一条样本耗时 1 毫秒那么批量梯度下降每更新一次参数就要 1000 秒接近 17 分钟。而训练一个像样的模型可能需要几万次更新这个时间成本完全不可行。所以批量梯度下降现在基本只出现在教科书里或者数据量极小的场景。3.2 随机梯度下降快但抖随机梯度下降SGD每次只用一个样本算梯度。更新频率极高一秒能更新很多次而且单样本带来的噪声有时候反而有助于跳出局部极小值。但问题是抖动严重损失曲线像心电图而且单个样本的梯度方向未必代表整体方向。这里有个反直觉的点SGD 的噪声不完全是坏事。有研究表明适度的梯度噪声能帮助模型逃离尖锐的局部极小值找到更平坦的极小值而平坦极小值通常泛化性能更好。所以后来的一些方法比如某些带噪声的优化器是主动往梯度里加噪声的而不是消除噪声。3.3 小批量梯度下降工程上的最优解小批量梯度下降Mini-batch GD是前两者的折中每次用一小批样本比如 32、64、256 条算梯度。它既利用了矩阵运算的并行能力又保持了较高的更新频率梯度方向也比单样本稳定得多。批量大小的选择是个经验活。太小比如 1、2梯度噪声大、GPU 利用率低太大比如上万单次更新慢、显存吃紧而且梯度太准反而可能降低泛化。实践中 32 到 256 是最常见的区间大模型训练因为显存限制往往用更小的批量配合梯度累积来模拟大批量。方法每次用多少数据更新频率梯度稳定性典型场景批量梯度下降全部样本极低极高小数据集、教学随机梯度下降1 条样本极高极低在线学习、流式数据小批量梯度下降32~256 条高中等绝大多数深度学习任务3.4 学习率那个最需要调的参数学习率决定了每次沿梯度方向走多远。它的影响非常直接太大参数更新步子迈太大损失会震荡甚至发散太小更新太慢训练半天不动还容易卡在不好的区域。我常用的一个判断方法是观察损失曲线。如果损失上下剧烈跳动、甚至越来越大学习率大概率偏大如果损失下降极其缓慢、几乎是一条平线学习率大概率偏小。理想情况下损失应该在前若干步快速下降然后逐渐放缓曲线整体平滑。学习率还和批量大小有耦合关系。一个经验法则是线性缩放批量扩大 k 倍学习率也大致扩大 k 倍。原因是批量越大梯度的方差越小估计越准可以放心走更大的步子。但这个规则不是绝对的大批量时往往需要配合 warmup预热来避免训练初期不稳定。4. 让梯度下降真正好用的那些改进纯 SGD 在实际训练里问题不少学习率难调、在鞍点附近停滞、不同参数方向上的尺度差异大。于是有了动量、自适应学习率、梯度累积等一系列改进。这些方法不是花架子每一个都针对一个具体的痛点。4.1 动量给梯度下降装上惯性SGD 的一个问题是在峡谷形的地形里会来回震荡前进缓慢。动量法Momentum的思路是维护一个速度变量它是历史梯度的指数加权平均更新时用这个速度而不是当前梯度。打个比方普通 SGD 像是一个没有惯性的小球每走一步都被当前坡度完全决定方向带动量的小球则像滚下山的球即使遇到小坑也能靠惯性冲过去。数学上速度 v β*v (1-β)*g参数更新用 v 而不是 g。β 通常取 0.9意味着速度大致是最近 10 步梯度的平均。动量带来的好处有两个一是抑制震荡方向上的来回摆动二是加速一致方向上的前进。在损失曲面有狭长峡谷时效果尤其明显。4.2 自适应学习率让每个参数有自己的步长不同参数的梯度尺度可能差好几个数量级。比如词嵌入层的梯度可能很小而输出层的梯度可能很大。用同一个学习率要么前者学不动要么后者震荡。自适应方法AdaGrad、RMSProp、Adam的核心思想是为每个参数维护一个独立的学习率根据它历史梯度的大小自动调整。Adam 是目前最流行的选择它同时结合了动量和自适应学习率还做了偏差修正。它的更新规则大致是用一阶矩估计梯度方向用二阶矩估计梯度尺度然后两者相除。Adam 的默认学习率 1e-3 在大多数任务上都能work这也是它受欢迎的原因——省心。但 Adam 也不是万能的。有研究发现在某些任务上 Adam 的泛化性能不如带动量的 SGD尤其是在计算机视觉领域。所以你会看到很多论文里视觉任务用 SGD MomentumNLP 和大模型任务用 Adam 或它的变体 AdamW。AdamW 的关键改动是把权重衰减从梯度更新里解耦出来这个细节对训练稳定性影响很大。4.3 梯度累积小显存跑出大批量的效果这是大模型训练里绕不开的技术。前面说过大批量训练梯度更稳但大批量吃显存。梯度累积的思路是把一个大 batch 拆成若干个小 batch分别前向反向算出梯度累加起来等攒够了一个大 batch 的量再统一更新一次参数。比如你想用 batch size 256但显存只够 32那就设累积步数为 8。每处理 32 条样本算一次梯度累加 8 次后更新一次参数。效果上等价于用 256 的批量训练但显存占用只有 32 批量的水平。这里有个容易踩的坑梯度累积时损失要除以累积步数。因为默认的损失是求和或求平均累加 8 次梯度相当于把损失放大了 8 倍不除的话等效学习率就变大了 8 倍训练会不稳定。我见过不少人配了梯度累积但忘了这一步结果训练直接发散排查半天才发现是这里的问题。提示梯度累积和梯度检查点是两个不同的技术。前者用时间换显存多次小批量累加后者用计算换显存不存中间激活反向时重算。两者可以叠加使用是大模型微调的标准配置。4.4 学习率调度训练全程动态调整固定学习率往往不是最优的。训练初期参数离最优解远可以走大步训练后期接近最优解需要小步微调。所以实践中几乎都会用学习率调度策略。最常见的几种阶梯衰减每隔若干轮降一次、余弦退火学习率按余弦曲线从大到小平滑下降、线性预热加衰减先从小学习率线性升到最大再逐渐降下来。大模型训练里预热 余弦退火几乎是标配。预热的作用是避免训练初期梯度估计不准时步子迈太大余弦退火则让训练后期平稳收敛。我个人的经验是预热步数一般设总步数的 1% 到 5%具体看批量大小和任务难度。批量越大、任务越难预热可以适当长一点。5. 训练中那些真实会遇到的坑理论讲完了但真正训练模型时问题往往出在细节上。这一节我挑几个高频问题把排查思路和解决方法讲清楚。5.1 损失不下降先别急着改模型损失不下降是最常见的问题但原因可能有很多。我的排查顺序是这样的先看数据有没有问题标签是否对齐、输入是否归一化再看学习率是否合适试着调大调小各一个数量级然后看梯度是否正常有没有出现 NaN 或全零最后才怀疑模型结构。梯度全零通常意味着激活函数饱和或者某处梯度被截断了。梯度 NaN 则往往是学习率太大或者数据里有异常值。我习惯在训练脚本里加一段梯度监控代码每隔若干步打印一次梯度的范数一旦发现异常能第一时间定位。5.2 梯度消失与爆炸的识别和处理梯度消失表现为浅层参数几乎不更新模型学不到东西梯度爆炸表现为损失突然变成 NaN 或者数值巨大。识别方法很简单打印每一层梯度的范数看是否随层数加深而指数级衰减或增长。处理梯度爆炸最直接的手段是梯度裁剪Gradient Clipping设定一个阈值梯度超过就按比例缩回去。这个在大模型训练里几乎是必备的尤其是 RNN 和 Transformer 类模型。梯度消失则要从结构上解决比如换用 ReLU 类激活函数、加残差连接、用 BatchNorm 或 LayerNorm。5.3 批量大小、学习率、显存三者的平衡这三个变量是相互制约的。批量大梯度稳但吃显存学习率大收敛快但可能不稳定显存有限就得在批量和模型规模之间取舍。我的实操策略是先确定模型规模受显存上限约束再用梯度累积把有效批量凑到目标值然后根据有效批量设定学习率最后用预热和调度保证训练稳定。这个顺序能避免反复调整。如果显存实在紧张梯度检查点是最后的救命稻草代价是训练速度下降 20% 到 30%。5.4 从零实现一遍反向传播的价值最后说个可能被低估的建议如果你真的想搞懂反向传播找个周末用 NumPy 从零实现一个两层网络的反向传播不用框架纯手写。你会被迫面对每一个矩阵的维度、每一个导数的符号、每一次求和的轴。这个过程很痛苦但走完之后你对框架里那些 API 的理解会完全不一样。我自己就是这么过来的。第一次手写的时候偏置梯度忘了对批次求和训练怎么都不收敛debug 了一整晚。但正是那次经历让我彻底记住了这个细节。框架帮你屏蔽了这些但也屏蔽了你对底层逻辑的直觉。偶尔回到裸奔状态收获往往比想象中大。梯度下降和反向传播这对组合撑起了整个深度学习的训练体系。从 1986 年的反向传播到今天的 AdamW 加梯度累积加余弦调度核心思想没变变的是工程上越来越精细的打磨。理解这些打磨背后的为什么比记住任何一个公式都重要。