ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

反向传播与梯度下降:大模型训练底层原理与实战解析

反向传播与梯度下降:大模型训练底层原理与实战解析 反向传播Backpropagation和梯度下降Gradient Descent这两个词大概是接触大模型的人最早听到、也最容易被糊弄过去的两个概念。网上说法很多有人说反向传播就是链式法则有人说梯度下降就是沿着斜坡往下走。话都没错但如果你真的去微调一个 Llama、Qwen 这类大模型你会发现 loss 不降、loss 变 NaN、显存溢出、训着训着权重爆炸……背后全是这两个基础原理在起作用。这篇文章我打算用一套手算例子加一份 numpy 实现把反向传播和梯度下降的每个环节拆开讲透再结合大模型预训练和微调的真实场景说清楚为什么现在的大模型训练离不开这两个齿轮以及实战中你会踩到哪些和梯度相关的坑。适合两类人一类是正在学大模型原理、想看穿训练底层逻辑的初学者另一类是已经在用 PyTorch 训模型、但遇到问题只能靠调参玄学解决、想知道背后原因的朋友。1. 大模型训练的底层发动机先理解这两件事在干嘛1.1 训练的本质是找到一组好参数不管模型多大预训练还是微调训练的本质都是一件事找到一组参数让模型在训练数据上的表现尽可能好。什么叫表现好需要一个量化指标这就是损失函数Loss Function。对大模型来说最常见的损失函数是交叉熵让模型对下一个 token 预测的概率分布尽可能接近真实分布。一个大模型动辄几十亿、上百亿参数这组参数对应的损失值构成一个超高维的地形图。我们的目标是在这个地形图上找到最低点。但问题是这个地形图我们根本看不见——你不可能遍历所有参数组合去算 loss。所以只能靠梯度下降这种迭代方法一步一步往低处走。在深度学习中这两件事是一套组合拳梯度下降决定往哪个方向走、走多远反向传播负责把这个方向梯度高效地算出来。1.2 梯度下降盲人下山的数学版本想象你被蒙上眼睛丢在一座山里任务是走到山谷最低点。你看不见全局只能通过脚下的坡度判断方向哪个方向下坡最陡就往哪个方向迈一步。迈完一步重新感受坡度再迈下一步。这就是梯度下降的直觉版本。数学上梯度是一个向量指向损失函数上升最快的方向。要想让 loss 变小就往梯度的反方向更新参数。更新公式就一行θ θ - η * ∇L(θ)。这里的θ是全部参数∇L(θ)是损失对每个参数的偏导数组成的梯度向量η是学习率也就是步长。有一个很容易忽略但极其重要的点梯度是局部信息。每一步你只知道脚下那一小块地是陡是缓并不知道前方是不是悬崖。所以学习率设大了一步迈出去可能直接跨过山谷跳到对面山坡甚至飞出地图loss 变成 NaN设小了走得慢训练一天 loss 还纹丝不动。1.3 反向传播在大模型里到底算什么角色如果说梯度下降是下山策略那反向传播就是快速测量坡度的方法。一个 70B 参数的模型如果每次更新参数都要用数值微分去近似梯度每一个参数都要额外算一次前向那训练成本将是一百多亿倍的开销根本不可能实现。反向传播用了一次前向加一次反向就能在跟一次前向差不多量级的计算成本内拿到所有参数的梯度。这是深度学习和以前很多传统机器学习方法拉开差距的关键原因。而且反向传播严格依赖链式法则网络越深链子越长。一个 50 层以上的 Transformer信号要从最后一层一路传回第一层中间每经过一层就要做一次矩阵乘法误差会在这里被放大或者被磨平。这就是后面要讲的梯度爆炸和梯度消失问题的根源。所以现代大模型里几乎所有结构设计——残差连接、LayerNorm、激活函数选择——都在围绕让梯度传得更稳做文章。2. 反向传播拆解链式法则怎么一层层传回去2.1 前向传播数据在网络里的流动要理解反向传播先得把前向传播搞清楚。以 Transformer 的一层为例输入 token 序列先经过 Attention 混一下信息再过 MLP 做非线性变换每块后面跟一个残差连接和 LayerNorm。前向传播就是输入数据按这个流水线一路算下去最后输出每个 token 的概率分布再和真实标签算交叉熵损失。前向传播的时候每一层都会产生一些中间结果比如矩阵乘法之后的输出、激活函数之后的值。这些中间结果看起来是用完就扔的其实必须全部存下来。为什么因为反向传播算梯度时要靠它们。比如算∂L/∂W链式法则里必然会出现上一层输入x或者激活输出a。这就是为什么训练大模型比推理吃显存得多——推理不需要存中间结果训练必须存batch size 一大显存直接爆掉。2.2 核心数学链式法则其实就一句话链式法则的直觉很简单如果L受z影响z受w影响那么w一变L的变化量就是两段变化率的乘积。写成公式∂L/∂w (∂L/∂z) * (∂z/∂w)。反向传播做的就是把这个法则从输出层开始一层一层往前套。因为网络是分层的输出端的梯度∂L/∂ŷ是能直接算的然后根据最后一层参数和输出的关系算出∂L/∂W_last再继续往回传给前一层。这个过程就像把一条项链从最后一颗珠子开始往回捋每一颗珠子接住上一颗传来的梯度再分摊到自己身上。注意一个关键点不用把整条链从头算到尾因为每层只关心自己直接相关的这部分。这就是反向传播高效的原因——每个中间变量只被算一次用记忆化的方式逐层传播而不是对每个参数单独跑一遍前向。很多初学者学到这里会犯一个思维误区以为反向传播是在代码里求导其实它只是一个有顺序的链式法则计算过程。2.3 手算一个两层小网络把梯度算明白纸上谈兵没意思。我手算一个最简单的两层网络感受一下梯度是怎么从 loss 一路传回参数的。网络结构输入x 1隐藏层权重w1 0.5偏置b1 0激活函数用 ReLU输出层权重w2 0.8偏置b2 0。目标输出y 1损失函数用均方误差L (ŷ - y)²。前向传播四步h w1 * x b1 0.5a ReLU(h) 0.5ŷ w2 * a b2 0.8 * 0.5 0.4L (0.4 - 1)² 0.36反向传播从 loss 开始往回算∂L/∂ŷ 2 * (ŷ - y) 2 * (0.4 - 1) -1.2∂L/∂w2 ∂L/∂ŷ * ∂ŷ/∂w2 -1.2 * a -1.2 * 0.5 -0.6∂L/∂b2 ∂L/∂ŷ -1.2∂L/∂a ∂L/∂ŷ * w2 -1.2 * 0.8 -0.96经过 ReLU 时因为h 0.5 0导数等于 1所以∂L/∂h -0.96∂L/∂w1 ∂L/∂h * x -0.96用学习率η 0.1更新参数w1 0.5 - 0.1 * (-0.96) 0.596b1 0 - 0.1 * (-0.96) 0.096w2 0.8 - 0.1 * (-0.6) 0.86b2 0 - 0.1 * (-1.2) 0.12再跑一次前向h 0.692a 0.692ŷ 0.86 * 0.692 0.12 0.715L (0.715 - 1)² 0.081。一轮迭代loss 从 0.36 降到了 0.081。虽然还是一个玩具网络但完整演示了一次前向算 loss反向算梯度梯度下降更新参数的闭环。多迭代几轮loss 会越来越接近 0。这个例子还揭示了一个重要细节每一层的梯度大小取决于上游梯度 × 局部导数的乘积。如果网络很深这个乘积不断累乘下去要么指数级变大要么指数级变小。这就是后面所有梯度稳定性问题的源头。3. 梯度下降的核心参数与算法演进3.1 学习率步子迈多大直接决定成败学习率是梯度下降里最敏感、也是最需要经验的一个超参数。你可以把它理解成下山的步长。步长太短走得慢可能还没到山脚训练就结束了步长太长一脚踩空直接跨过山谷跑到对面更陡的地方甚至从山道上滚下去——对应到实际训练就是 loss 抖动、发散、最后变成 NaN。在大模型训练里学习率的设置还有一个反直觉的规律模型越大学习率通常越小。以预训练为例小模型常用1e-3甚至3e-3的学习率而几十 B 的模型往往要用到1e-4甚至更低。原因有两层一是大模型的 loss 面对参数变化更敏感参数稍微大动干戈行为就可能剧变二是大模型训练成本太高一次发散可能浪费几天时间宁可保守一点。微调场景则更夸张比如对 LLM 做全参微调主流学习率是1e-5到3e-5这个量级。预训练好的模型已经在 loss 地形的一个很深的盆地附近了微调只需要在附近小幅挪动步子一大就跳出盆地把学好的能力全忘了。这就是为什么大模型微调领域普遍流传学习率千万别超过 1e-4的说法。3.2 从 SGD 到 AdamW优化器是怎么演进的学习率设多少跟用哪个优化器强相关。最早的基础优化器叫 SGD随机梯度下降更新规则就是θ - η * g。SGD 简单靠谱但有个缺点对所有参数一视同仁碰到 loss 地形很崎岖的地方容易震荡收敛慢。后来出现了动量Momentum机制相当于给下山过程加了一个惯性如果最近几个梯度方向一致就加速冲过去如果方向反复横跳惯性会抵消一部分震荡。再后来 RMSProp 又加了自适应能力让每个参数根据自己的梯度大小单独调步长。把这两者合在一起就是 AdamAdaptive Moment Estimation。Adam 有两个核心机制一阶动量累积过去梯度的指数加权平均相当于 momentum二阶动量累积梯度平方的指数加权平均相当于给每个参数单独做归一化。直接把学习率从数值上解耦了——即使某一维梯度非常大除以二阶动量之后实际步长也不会失控。这让 Adam 成为大模型的默认选择。到了大模型时代工程师们发现 Adam 和权重衰减Weight Decay一起用会有问题L2 正则化在 Adam 的归一化机制下效果会打折扣。于是 AdamW 出现了把权重衰减从梯度里拆出来直接在参数更新时做。现在不管是 GPT 系列还是 Llama 系列预训练基本上都是 AdamW配合固定的 beta 参数和权重衰减系数比如weight_decay 0.1。3.3 大模型微调里优化器与学习率的实战配置如果你准备微调一个大模型我给一个可以直接抄作业的配置参考这是社区和开源项目里最常见的默认设置之一优化器AdamW学习率2e-5或3e-5全参微调取低一些LoRA 微调可以稍微调高一点学习率调度cosine衰减配合 3% 左右的 warmup 步数权重衰减0.01到0.1之间批次大小取决于显存但一般用梯度累积模拟到等效32到128的 batch size在这些配置下loss 曲线通常会呈现快速下降然后缓慢震荡收敛的形态。如果你用的是 cosine 调度训练后半段学习率会逐步趋近于零对 loss 做精细打磨。这里有一个很多人不知道的小细节训练结束时是不是要用最后一步的学习率还是衰减后的最小学习率对大模型效果影响不小。主流做法是用模型在验证集上表现最好的那一步做 checkpoint而不是直接取最后一步。4. 从零实现反向传播与梯度下降numpy 实操4.1 搭建最小可训练网络讲再多理论不如自己动手跑一遍。我用 numpy 实现一个两层网络手动写前向、反向和参数更新不借助任何深度学习框架。这个过程能让你真正看见梯度是怎么流动的。import numpy as np # 数据单个样本输入 x1目标 y1 x np.array([1.0]) y_true np.array([1.0]) # 初始化参数 w1, b1 0.5, 0.0 w2, b2 0.8, 0.0 lr 0.1 # ReLU 及其导数 def relu(z): return np.maximum(0, z) def relu_derivative(z): return (z 0).astype(float) for step in range(20): # 前向传播 h w1 * x b1 # 隐藏层线性输出 a relu(h) # 激活 y_pred w2 * a b2 # 输出层 loss (y_pred - y_true) ** 2 # 反向传播手动算梯度 dL_dy_pred 2 * (y_pred - y_true) dL_dw2 dL_dy_pred * a dL_db2 dL_dy_pred dL_da dL_dy_pred * w2 dL_dh dL_da * relu_derivative(h) dL_dw1 dL_dh * x dL_db1 dL_dh # 梯度下降更新 w1 - lr * dL_dw1 b1 - lr * dL_db1 w2 - lr * dL_dw2 b2 - lr * dL_db2 if step % 4 0: print(fstep {step}: loss {loss.item():.4f}, w1 {w1:.4f}, w2 {w2:.4f})运行结果大致是step 0: loss 0.3600, w1 0.5960, w2 0.8600 step 4: loss 0.0104, w1 0.4918, w2 1.2291 step 8: loss 0.0012, w1 0.4980, w2 1.3562 step 12: loss 0.0002, w1 0.4998, w2 1.4015 step 16: loss 0.0000, w1 0.5000, w2 1.4166loss 从 0.36 降到了几乎为 0。注意一个有意思的现象w1收敛到 0.5 附近就不再动了而w2一直涨到了 1.4 左右。这是完全合理的——对于这个单样本任务正确的参数不是唯一的只要w1 * w2接近 1 就能把 loss 压到 0。梯度下降只会把你带到它找到的第一条山谷不保证是唯一解这就是大模型训练中相同配置也可能得到不同模型的原因之一。4.2 代码里的关键细节这段代码最值得体会的是反向传播部分的书写顺序。你看求dL_dw1必须先算出dL_dh而dL_dh又依赖dL_dadL_da又依赖dL_dy_pred。这个顺序就是从输出到输入的传播顺序。实际框架里PyTorch 的backward()函数会自动帮你做这件事但它的底层逻辑跟我手写的一模一样在前向过程中构建计算图反向时沿着计算图的边逐层回传梯度。所以理解这段手工代码你就能理解为什么 PyTorch 有时候会报 gradient computation requires grad 或者 loss.backward() 之前要 zero_grad() 这类错误——因为计算图是无状态累积的不清理梯度就会把多轮梯度叠加在一起。还有一个细节ReLU 在负数区域的导数恒为 0。如果某个神经元的输入一直是负数它的梯度就是 0权重永远得不到更新。这就是著名的神经元死亡问题。解决办法包括换用 Leaky ReLU、GeLU 等激活函数。现在 Transformer 里大量使用 GeLU 而不是 ReLU除了性能上的考虑也和它没有把负数区域梯度完全截断有关系。4.3 把代码平移到 PyTorch 的思路如果你用 PyTorch 写同样的训练流程代码会短很多import torch import torch.nn as nn x torch.tensor([1.0]) y_true torch.tensor([1.0]) # 两层网络 model nn.Sequential( nn.Linear(1, 1), # 等价于 w1, b1 nn.ReLU(), nn.Linear(1, 1), # 等价于 w2, b2 ) optimizer torch.optim.AdamW(model.parameters(), lr0.1) loss_fn nn.MSELoss() for step in range(20): optimizer.zero_grad() y_pred model(x) loss loss_fn(y_pred, y_true) loss.backward() # 反向传播自动算所有梯度 optimizer.step() # 梯度下降更新参数PyTorch 省掉了所有手写求导的部分loss.backward()就是反向传播optimizer.step()就是梯度下降。但你完全没必要因此跳过手动实现的过程因为在框架里调参时你看到的每个报错、每个奇怪的 loss 曲线、每个学习率设置的建议背后都是这套手工逻辑在跑。把底层逻辑吃透你调参就不玄学了。5. 大模型训练中的梯度问题与排查实录5.1 梯度消失和梯度爆炸两个老冤家回到开篇说的那个点反向传播是链式法则的连乘连乘次数等于网络的层数。如果每层的局部梯度都小于 1连乘之后会指数级趋近于零这就是梯度消失如果每层的局部梯度都大于 1连乘之后会指数级发散出去这就是梯度爆炸。在深层 Transformer 里这两个问题都真实存在。比如早期 Transformer 和 RNN 时代梯度消失导致网络没法训练于是有了残差连接让梯度有一条高速公路可以直接从最后一层传回前面绕开中间的连乘路径。这招非常有效现在几乎每个大模型架构都有残差连接。梯度爆炸则常用梯度裁剪来解决设定一个阈值如果梯度范数超过阈值就把梯度整体缩回到阈值以内。大模型预训练基本都开梯度裁剪这也是为什么你会在训练代码里看到clip_grad_norm_(model.parameters(), max_norm1.0)这样一行。在真实训练中梯度范数是一个非常值得盯的指标。如果某一步的梯度范数突然暴涨到正常值的几十倍即使有裁剪兜底也说明训练可能在走向失控的边缘。我见过不少大模型训练失败案例日志里梯度范数早就出现了异常信号只是没人注意。5.2 学习率 warmup 与大模型稳定性大模型训练里有个非常普遍的操作前几千步用很小的学习率跑之后再慢慢升到目标学习率这叫 warmup。为什么要这么做原因在于初始化阶段模型参数完全是随机的此时算出来的梯度方差非常大如果一上来就用大学习率更新模型行为会剧烈抖动甚至直接发散。warmup 相当于是给训练过程一个热身期让参数先从随机状态过渡到一个相对稳定的区域再开始大步快跑。对大模型来说 warmup 不是可选项而是必需品。常见的做法是前 1% 到 3% 的训练步数做 warmup从零或者从学习率的十分之一线性增长到目标值之后接 cosine 衰减。这个调度策略在很多开源大模型训练代码里是标准配置。顺带聊一下 batch size 变大对梯度的意义梯度是多个样本梯度的平均batch size 越大梯度估计越接近真实梯度噪声越小训练越稳定。但显存放不下怎么办答案是梯度累积——先跑几个小批次把梯度累加在一起攒够了再统一更新一次参数。效果上等效于大 batch只是每一步实际更新频率变低了。LoRA 微调、QLoRA 微调里也经常配合梯度累积来模拟较大 batch这是显存受限时最实用的技巧之一。5.3 混合精度下的梯度与 loss scaling现在训练大模型几乎不可能不开混合精度AMP。混合精度的核心思路是前向和反向用 FP16 或 BF16 加速优化器状态用 FP32 保精度。但 FP16 有个问题——它能表示的最小正数范围有限梯度值往往很小在 FP16 下直接表示就会溢出变成 0梯度就蒸发了。为了解决这个问题损失缩放loss scaling登场在反向传播前先把 loss 放大若干倍梯度也跟着放大等梯度算完再缩小回去。PyTorch 的torch.cuda.amp.GradScaler就是干这个的。不过现在新一代 GPU 和框架更推荐 BF16因为它的表示范围和 FP32 一样大不需要 loss scaling 也不会梯度下溢。这也是为什么 Llama、Qwen 等新一代大模型训练几乎都用 BF16。我遇到过最典型的混合精度坑是FP16 下 loss 突然变成 NaN关掉 AMP 就正常。排查方法很简单看看是不是没有做 loss scaling或者梯度裁剪阈值设置得太激进了。而 BF16 遇到 NaN 的几率小很多一旦出现基本可以断定数据里有异常或者学习率失控了。5.4 常见问题速查表现象可能原因排查与处理loss 变成 NaN学习率过大、AMP 溢出、数据含 NaN降低学习率、开启或调整 loss scaling、检查数据loss 震荡剧烈学习率偏高、batch size 太小降低学习率、增大 batch 或梯度累积loss 完全不降学习率过低、梯度消失、数据没 shuffle 或预处理错误调大学习率、检查网络结构、验证输入输出梯度范数暴涨网络太深、训练不稳定开启梯度裁剪、检查是否缺少 LayerNorm 或残差loss 先降后升学习率调度不当、过拟合检查 cosine 调度、增加正则或提前早停微调后能力下降学习率太大破坏了预训练知识改用更小学习率或尝试 LoRA 减少可训练参数范围这张表是我在实际训练和微调中反复用到的最基本的排查思路。大模型训练的大部分故障追根溯源都能落到梯度去哪儿了、梯度是不是太大了这两个问题上。6. 几点实际操作中的体会最后聊点个人体会。第一理解反向传播和梯度下降最直接的价值不是让你能手写框架而是让你知道训练日志里的每一个数字是怎么来的。比如 lr、grad_norm、loss 曲线这些不是摆设是给你信号判断训练是否健康的。我每次启动训练第一件事就是设置好梯度范数的日志和告警而不是只盯着 loss。第二学习率策略值得花时间调。很多人微调大模型上来就一把2e-5打到训练结束其实结合 cosine 调度和 warmup效果会有肉眼可见的差异。我自己做 LoRA 微调时固定住其他超参数只调整学习率调度方式同样的数据下最终评估指标能差好几个点。这个东西在论文里一句话带过但实操里就是实打实的收益。第三能跑通反向传播并不等于能用它训好模型。训练大模型的时候那些最诡异的 bug——loss 周期性跳高、某个 token 的 loss 特别大、多个 GPU 之间梯度不同步——最后查下来往往不是反向传播本身错了而是数据和并行策略的问题。所以我的建议是反向传播和梯度下降的数学原理要懂但排查问题的重心永远放在数据、并行和数值稳定性上。把这篇文章里手算的例子自己推一遍再用 numpy 代码跑一遍你对大模型训练的理解会比看一百篇论文都扎实。至少下次有人问你为什么微调要用小学习率你可以告诉他因为预训练好的模型已经站在谷底附近步子大了容易重新爬坡。
返回列表