
训练一个稍微大一点的模型时我最怕遇到两类报错一类是 CUDA out of memory还有一类是没有报错但 loss 卡在某个数值上纹丝不动。表面上看这两类问题风马牛不相及但排查到最后十有八九都会撞到同一个环节——反向传播。这个环节是深度学习的动力核心也是大模型训练里最容易被低估、最容易被背锅的部分。反向传播负责把损失函数的信息逐层传回网络计算每个参数该往哪个方向修正梯度下降则拿着这份修正指南去更新参数。二者听起来像是教科书里的两个独立章节但在大模型场景下它们被牢牢绑在一起显存不够是反向传播的激活值在作祟loss 下降不稳定是梯度噪声在起伏微调效果不理想往往是优化器选错了。这篇文章我想把大模型反向传播与梯度下降这套机制讲透从链式法则的直觉一直讲到 ZeRO、混合精度、LoRA 这些实际训练里的工程手段最后再分享几个我真实踩过的坑。适合想系统理解大模型训练原理的工程师也适合准备入局大模型微调的初学者。1. 从链式法则说起反向传播为什么是反向的1.1 计算图与局部梯度让误差坐一趟反向电梯理解反向传播先要理解一个本质问题一个参数影响最终 loss 的路径常常非常遥远。假设你有一个 100 层的 Transformer输入侧的某一层参数稍微动一点这个变化要穿过后面 99 层才能影响最后的损失函数。如果每次都从第一层开始去计算这个影响计算量会随着层数爆炸式增长。反向传播的思路是把误差看作一辆在楼层之间运行的电梯。前向传播时每一层都拿到了自己的输入并产生输出同时把这些局部关系记录在计算图中。反向传播时误差信号从最顶层的 loss 出发逐层向下传递每一层只需要计算我自己的输出对 loss 的影响再乘以上一层传下来的梯度。这个过程用的数学依据就是链式法则。我习惯用一个生活类比把网络理解成一条流水线前向传播是产品从原料端流到成品端反向传播则是质检不合格时逐道工序倒查是哪一步出了问题。每道工序只需要关注自己加工了什么东西以及上下道工序送来的中间产物根本不需要知道整条流水线长什么样。这就是局部梯度加链式法则的意义。大模型之所以能训起来本质上是这套逐层倒查机制在起作用而不是靠什么黑魔法。1.2 一个三层网络的梯度计算示例亲手算一次为了不被大模型的复杂度吓到我建议从最简单的手算开始。假设一个两层参数的网络输入 x第一层参数 W1第二层参数 W2中间过一个激活函数 σ损失函数 L。前向传播就是x - z1 W1·x - a1 σ(z1) - z2 W2·a1 - L (z2 - y)²反向传播时我们先求 dL/dz2 2(z2 - y)。这是误差信号的起点。然后沿原路倒着走dL/dW2 dL/dz2 · a1dL/dz1 dL/dz2 · W2 · σ(z1)dL/dW1 dL/dz1 · x注意计算 dL/dW1 的时候dL/dz1 是从上面传下来的W2 和 σ(z1) 都是前向传播时已经记录下来的值。所以反向传播不需要重新执行一遍全网络只需要按图索骥地做乘法和加法。这个例子最直观的价值在于你看出来反向传播的每一层梯度都只依赖当前层的局部值和上一层传下来的误差信号。这也是为什么 PyTorch 这类框架能自动求导——它只是把链式法则变成了可复用的图遍历算法。很多人一上来就啃大模型源码被几百个模块吓住其实所有复杂模型的梯度计算拆到极致都是这样一组组局部乘法。1.3 前向传播与反向传播的不对称性大模型为什么格外痛到这里你可能已经发现了计算的对称美。但工程上有个不对称的地方前向传播只需要保留最终输出而反向传播需要拿到每一层的中间激活值。换句话说反向传播为了按图索骥必须把整个计算图的中间状态都存下来。这条性质在几十层的卷积网络里也许还能忍但在动辄几十上百 GB 参数的大模型里直接决定了显存规划。早期训练大模型时最常见的 OOM不是参数本身太大而是反向传播要求的激活值缓存太大。这个问题我后面会专门展开现在只请你记住一个结论反向传播的工程代价不是来自它的数学而是来自它的记忆。理解了这一点你再看 gradient checkpointing、激活值重计算、ZeRO 这些技术就明白它们到底在省什么了。2. 梯度下降的家族谱大模型为什么普遍选择了 AdamW2.1 损失地形的比喻在高维山谷里找最低点梯度下降解决的问题很直观已知每个参数对 loss 的梯度怎么走才能更快更稳地到达最低点如果是一个碗状的二维函数用梯度走直线就行。但大模型的损失函数是一个几万亿维的曲面别说看不见连山谷和鞍点都可能是常态。我在实践中看到很多人把梯度下降理解成沿着梯度反方向走一步这个描述不算错但过于简化。在大模型里损失曲面极其崎岖有的方向梯度很大但方向变化剧烈有的方向梯度很小但一路平坦。如果所有方向都用一个固定步长去走就会遇到振荡或者停滞。这就需要梯度下降的变体来协调步长和方向。这里还有一个大模型特有的现象由于参数空间维度实在太高几乎不存在一个全局最优解让你精确找到。训练的目标更多是落入一个足够好的局部极小或平坦区域这个区域泛化能力往往更好。优化器的选择直接决定你在高维地形里能不能顺利走到这样的区域。2.2 从 SGD 到 Momentum再到 Adam每一步都在解决什么问题最朴素的 SGD 就是当前梯度反方向走一步。它在大模型上的问题非常明显如果某个维度长期给出很小的但方向稳定的梯度SGD 仍然只能一步步慢慢爬反过来如果某个维度梯度一会儿正一会儿负SGD 就会在原地抖动。Momentum 的思路是给历史梯度累积一个速度相当于给小球一个惯性让它加速通过小梯度区域又能在抖动方向上自行刹车。Adam 在 Momentum 的基础上增加了每个参数的梯度平方指数平均。它的直觉是如果一个维度过去一直梯度很大说明这个方向很不稳定就放小步长如果一直很小且稳定就放大步长。这样每一个参数都拥有属于自己的自适应学习率。这个特性对大模型这种参数规模极大的场景天然友好因为不可能用几千个调参实验去为每个参数维度特意设计步长。有意思的是Adam 在训练后期可能出现一个 bug 式的副作用由于二阶矩的累积导致有效步长变得非常小从而很难收敛到最优点附近的尖锐极小。这也就是为什么后来很多大模型训练干脆采用带权重衰减的 AdamW把参数自身的衰减与梯度更新解耦整体训练稳定性要好得多。你去看主流大模型框架的默认配置几乎清一色是 AdamW cosine 学习率调度这个组合不是拍脑袋选的背后是上述一系列失败模式逼出来的。2.3 权重衰减与 AdamW一个小改动为什么能稳定整个训练普通 Adam 里也可以加 L2 正则但它的实现方式是把权重衰减项混进梯度里而 Adam 会把这一项也做二阶矩归一化结果反而破坏了正则的原本语义。AdamW 把权重衰减从梯度里拿出去直接在参数更新时对参数本身做一次小的收缩让正则化与自适应的梯度更新互相不干扰。我在实际微调大模型时不加权重衰减往往会在训练后期看到 loss 在小范围内波动、验证集涨不上去加上一个常见的 0.01 或 0.05 权重衰减后波动幅度会明显变小。这不是玄学而是 AdamW 对高维参数空间中参数漂移的有效抑制。尤其在做指令微调时很多参数本身已经包含预训练得到的知识不希望它们被微调过程带得太远权重衰减就相当于给每步更新加了一个锚。2.4 学习率调度warmup、cosine decay 与损失函数的温度单独说梯度下降一定绕不开学习率。大模型训练的第一课往往是 warmup。刚开始几步Adam 的一阶矩和二阶矩都是瞬时的如果直接用大学习率梯度的样本噪声会被放大早期 loss 很容易飞掉。所以需要先用很小的学习率让优化器积累一些统计量再把学习率提升到目标值。后续的 cosine decay 则是一套更有节奏感的降火方式训练前期保持相对较高的学习率让模型快速探索后期逐渐降低学习率让参数在小范围内精细收敛。我在跑生成模型微调时曾经对比过恒定学习率和 cosine decay同样总步数下后者在生成质量上的提升通常更平滑。大模型里的学习率不要把它当成一个孤立的超参数它和 batch size、梯度累积步数、warmup 步数是联动的牵一发而动全身。3. 大模型训练时反向传播的真实开销显存、计算与分布式同步3.1 激活值显存反向传播的账本才是 OOM 的真凶回到开头提到的 OOM。很多人以为显存消耗等于模型参数量乘以参数精度事实上在训练模式下激活值常常会占掉 50% 以上的显存。以常见的中等规模开源模型为例纯推理时用 FP16参数量对应的显存也就十几 GB但启用训练反向传播后每一层的注意力分数、中间投影结果、层归一化统计值都要被缓存当 batch size 稍大激活值显存会轻松超过参数显存。这也是为什么工程界会反复强调 gradient checkpointing重计算。思路很简单前向传播时不把每个激活值都缓存只挑一些关键节点存下来反向传播用到中间结果时重新算一遍。相当于用额外的计算换显存实测可以把激活值显存降为原来的三分之一甚至更少。刚开始用 checkpointing 时训练时间会明显变长大约增加 20%~30%但总比 OOM 停摆要划算。在训练脚本里你只需要在模型的前向传播函数里调用torch.utils.checkpoint.checkpoint包住若干层其余交给框架处理。这类优化手段的价值恰恰体现为反向传播的账本被大幅压缩。我之前有一版代码只是把 embedding 层之后的所有 Transformer 层都用 checkpoint 包住同一个 7B 模型的训练显存就从直接爆掉降到了可用范围代价是每个 step 多算了一组前向。3.2 混合精度与 loss scaling反向传播在数值悬崖边缘跳舞大模型训练几乎不会用 FP32 跑满因为在现代 GPU 上 FP16/BF16 拥有更高的吞吐量。但半精度给反向传播带来一个麻烦梯度太小时FP16 会直接变成 0导致参数收不到更新梯度太大时FP16 又容易溢出变成 inf。解决方案是 loss scaling。前向传播计算出的 loss 先乘一个比较大的缩放因子再做反向传播。这样梯度也跟着被放大掉到可表示范围之外的概率就降低了。之后在更新参数前把梯度再除以同一个缩放因子还原真实大小。这套机制在 PyTorch 的 GradScaler 里被封装得很好但第一次用的人往往忽略了一个细节loss scaling 只保护了反向传播得到的梯度并不影响前向计算的精度。我见过有同学自己写混合精度训练循环时每一步都手动判断梯度是不是 NaN看起来严谨实际上非常容易漏掉 BF16 的细节。BF16 虽然和 FP16 有同样的显存优势但它的指数范围比 FP16 大得多正常训练里极少出现梯度下溢这也是为什么很多大模型预训练团队更偏爱 BF16。如果你刚接触大模型训练直接上 BF16 往往能省掉不少和 loss scaling 有关的麻烦。3.3 梯度累积与梯度同步反向传播在并行训练中的角色大模型单卡放不下就得用数据并行或模型并行。以数据并行为例每张卡各自跑一份模型和自己的 mini-batch前向传播完成后各自做反向传播得到各自的梯度。最后需要把这些梯度做一次 AllReduce 同步取平均后再交给优化器更新。梯度累积则是在同步之前把连续几个 micro-batch 的梯度累加再统一更新参数。它模拟了更大的 batch size却不需要一次性把所有数据塞进显存。这里有个关键点梯度累积后必须把梯度除以累积步数否则等效学习率会偏大极容易导致训练不稳定。我见过不止一个朋友因为忘了这个除法loss 一路下滑但模型质量反而变差——本质上是用了过大的等效 batch size优化轨迹变得很激进。在分布式训练里反向传播还有一个隐藏的通信开销每个梯度张量计算完就要立刻参与 AllReduce这会导致大量细碎的小通信拖慢整体吞吐。工程上常用的做法是梯度分桶把若干层的小梯度攒成一个桶再做同步能显著减少通信次数。这些都说明在真正的工业级训练里反向传播不只是一个算法概念还是一个需要精心编排的分布式流程。4. 微调场景下的降本增效LoRA 如何改变反向传播的规则4.1 全量微调中优化器状态的隐形成本很多人知道微调大模型很贵但没细算过钱花在哪里。全量微调不仅要对所有参数求梯度还要为每个参数维护优化器状态。以 AdamW 为例每个参数要保存一份一阶动量、一份二阶动量再加上参数本身和梯度接近 4 倍于模型参数大小的显存要被占用。一个 7B 模型用 FP16 全量微调单是优化器状态就是 28GB 以上很多消费级显卡根本扛不住。这也是为什么真正做微调的人很少直接对所有参数跑完整 AdamW而是优先考虑参数高效微调方法。在理解 LoRA 之前先要明白这个隐形成本的构成。梯度下降本身需要保存历史统计量而大模型的参数规模让这些统计量的显存开销变得完全不可忽略。4.2 LoRA 的核心思想冻结主干只更新低秩矩阵LoRA 的基本假设是大模型微调时的参数更新其实分布在一个很低的秩子空间里。于是它冻结原模型全部参数只在某些模块旁路注入一个低秩矩阵 A×B。反向传播时梯度只需要流经这条低秩旁路原始主干自身的参数不需要更新也就省掉了为它们维护优化器状态的开销。我在一块 24GB 显存的卡上做过对比实验用 FP16 全量微调 7B 模型刚起训练就 OOM换成 rank8 的 LoRA 后显存峰值明显下降一个 epoch 的时间也缩短了不少。虽然最终效果不等于全量微调但在多数指令微调场景下差距可以做到非常小。而且 LoRA 的权重可以在训练结束后合并回主干推理时完全不需要额外的旁路开销这是它在生产环境里特别受欢迎的原因之一。LoRA 还有一个额外好处你可以随时把训练好的低秩矩阵保存为一个小文件发给别人时不需要推送整个模型。这一点在协作场景里特别实用——相当于只传输变化的增量而不是完整的模型快照。4.3 一个更省的视角LoRA 与 gradient checkpointing 的组合LoRA 训练时同样会遇到反向传播保留激活值的问题。许多人以为冻结了主干就能随便放大 batch size其实主干虽不更新但它的中间激活值对旁路的梯度计算依然是必需的。这时候最好的做法通常是把冻结主干和 gradient checkpointing 叠加使用让主干层被部分重计算反而能获得接近全层训练可比的吞吐量。这个组合让我真正理解了为什么 PEFT 工具库要把 LoRA 适配器和 checkpointing 固化为默认推荐。它们不是两个独立技巧而是共同作用在反向传播这条流水线上一个减少参数更新规模一个减少中间缓存规模。如果你卡在本地部署大模型微调时的显存瓶颈最优先尝试的往往不是换更大的卡而是把这两个开关同时打开。5. 踩坑复盘最常见的 loss 不下降、梯度爆炸与反向传播 bug5.1 排查 loss 变成 NaN 的完整链路早年间我训练一个中等规模的 Transformer 时遇到了一次非常诡异的 NaN前几十步完全正常到第 60 步时 loss 突然变成 NaN而且每次都在同一个位置附近爆。当时的日志里没有任何异常分布式梯度同步也没有报错。顺着反向传播的思路往下查我先看 loss 是否在某一刻变成负数或无穷大把 loss 值打点后发现它在第 58 步开始出现 1e20第 60 步直接 NaN。这说明是前向传播溢出了。但为什么同样的数据在前几个 epoch 没事后来定位到是数据 loading 顺序恰好让某些 batch 里出现极端异常值而前向传播在手写实现时没有做 clip到反向传播时梯度以指数方式放大最终击穿 FP16 的可表示范围。这件事给了我两个教训第一反向传播的问题往往会先表现为前向计算或者数据问题第二用 FP16 训练时给 loss 加一个适当的梯度裁剪或数值监控比事后追查 NaN 要高效得多。现在我在训练脚本里一定会记录每步的 loss、grad norm以及输入数据的 min/max 统计量。出现异常时先看这些指标定位到是数据侧还是模型侧再决定动刀的位置。5.2 grad norm 监控与梯度裁剪的实际操作大模型训练通常把 grad norm 作为最重要的健康指标之一。计算非常简单把所有参数的梯度拼起来求 L2 范数。正常的训练曲线里grad norm 应该在某个范围内缓慢波动如果某一刻它突然比历史均值大超过 10 倍就说明出现了梯度尖峰。梯度裁剪gradient clipping是应对尖峰的标准动作设置一个阈值比如 1.0如果 grad norm 超过阈值就把所有梯度按比例缩放使其范数刚好等于阈值。注意裁剪的是整体范数不是逐个参数裁剪否则会破坏不同层之间的梯度比例。我通常在 debug 阶段会把阈值放得比较小0.5~1.0在稳定期再慢慢放大。还有一点容易被忽略梯度裁剪应该在梯度同步之后、优化器更新之前执行。如果先裁剪再同步每个 rank 的裁剪阈值不同最后同步出来的梯度范数可能仍然很大。这个顺序问题在分布式训练里特别值得留意我踩过一次之后现在都会在训练循环里明确标注三段顺序反向传播 - 梯度同步 - 裁剪 - 更新。5.3 用数值梯度验证反向传播的正确性有时候写了一个自定义层或者改了损失函数loss 却纹丝不动不能立刻怀疑是学习率问题很可能是反向传播算错了。最可靠的验证办法就是 numerical gradient挑选几个参数维度把它们分别扰动一个很小的 epsilon然后用 (L(xeps) - L(x-eps)) / (2*eps) 与反向传播给出的梯度对比。我在实现一个自定义 attention 偏置层时就用这个方法抓到了一个维度错位的 bug。当时数值梯度和解析梯度相差了约 10%对照后发现是对某个维度用了错误的展开逻辑。数值梯度虽然慢但在模型较小、参数较少时完全够用是每一个想自己写训练代码的人必备的调试工具。做一个简单的验算流程固定一份随机输入跑一次前向和反向拿到所有参数的解析梯度。只挑一个参数把它的值加一个微小扰动重新前向得到 loss1。再把它减一个微小扰动重新前向得到 loss2。把 (loss1 - loss2) / (2*eps) 和解析梯度对比误差应该在 1e-4 量级以内。这套流程验证的是反向传播计算链是否正确而不是模型是否收敛。它解决的是梯度下降拿着错误的方向图的问题。方向图错了后面无论怎么调学习率都没用。5.4 学习率与超参数的一点个人经验最后说点有实证的部分。对于大模型微调我常用3e-5 起步、1e-5 左右微调作为基准但这只是起点。更重要的原则是当 loss 不降时先看梯度是否正常再看优化器状态是否被正确重置最后才去动学习率。很多人一上来就把学习率调大两个数量级结果 loss 直接飞到 NaN这完全是浪费算力。我自己在复现一个开源对话模型的微调时曾经因为忘记加载优化器状态导致前几百步都在用随机初始化的 Adam 动量重新累积训练曲线莫名其妙地下滑又回升。这种问题只有在把反向传播和梯度下降看成一条完整链路时才会注意到——反向传播负责给出正确的梯度而优化器状态负责让梯度下降的第一步走得稳。如果你也正准备从头复现一个模型训练不妨先拿 10M 参数的玩具模型跑一遍手推反向传播每个张量的形状。这个习惯在后来我处理大模型训练日志时帮了很大忙——所有复杂的显存优化、梯度同步、混合精度问题本质上都回到同一件事反向传播有没有拿到它该拿的梯度梯度下降有没有用对这份梯度。就这么简单。