首页 / 资讯中心 / 文章详情

反向传播与梯度下降实战:从原理到调参避坑指南

反向传播与梯度下降实战:从原理到调参避坑指南 ★ FEATURED ARTICLE
1. 从一次训练翻车说起为什么反向传播和梯度下降值得死磕刚接触大模型那会儿我干过一件现在想起来都脸红的事。当时用一个小型Transformer做文本分类训练了整整两天loss曲线跟心电图似的上下乱跳准确率死活卡在随机猜测的水平。我以为是数据有问题换了三份数据集又怀疑模型结构写错了对着论文逐行核对了两遍。最后一位做优化的朋友看了一眼我的学习率——0.1配的还是SGD。他沉默了三秒说了一句让我记到现在的话“你这不叫训练叫在参数空间里蹦迪。”这件事让我彻底明白一个道理大模型的参数量再大、架构再花哨真正让它从一堆随机数变成能说人话的智能体的就是反向传播和梯度下降这对组合拳。反向传播负责算清楚“每个参数该往哪个方向调、调多少”梯度下降负责“实际去调”。一个管诊断一个管治疗缺了谁模型都学不会东西。这篇文章我打算把这两个东西从头到尾拆一遍。不是教科书式的公式罗列而是从一个实际调过模型、踩过坑的人的角度讲清楚它们到底在干什么、为什么这么干、实操中哪些参数会要命、遇到问题怎么排查。不管你是刚入门想搞明白大模型基础理论的新手还是已经能跑通微调但说不清底层逻辑的开发者应该都能从里面找到对自己有用的东西。涉及到的链式法则、学习率设置、梯度累积这些关键词我都会结合具体场景展开尽量让每个概念都能落到“你下一步该怎么做”上面。2. 反向传播到底在传什么把链式法则讲成人话2.1 一个生活化类比工厂质检的追责链条先别急着看公式。想象一个工厂流水线原料经过五道工序变成成品每道工序都有几个可调的旋钮。成品不合格你要找出是哪道工序、哪个旋钮的问题。但麻烦在于你只能看到最终成品的合格率中间每道工序的具体情况你没法直接测量。反向传播干的就是这件事从最终的错误出发沿着流水线倒着走算出每个旋钮对最终错误的“责任大小”。如果某个旋钮稍微转一点成品合格率就大幅提升那这个旋钮的“责任”就大应该多调如果转了跟没转一样那就不用怎么动它。这个“责任大小”在数学上就是偏导数。而“沿着流水线倒着走、逐层计算责任”的方法就是链式法则。链式法则说白了就是如果A影响BB影响C那A对C的影响等于A对B的影响乘以B对C的影响。一层一层乘下去就能把最终误差的“责任”分配到每一个参数上。2.2 计算图反向传播的施工图纸实际实现中框架不会真的去“倒着走流水线”而是先把整个前向计算过程画成一张计算图。每个节点是一个运算加法、乘法、激活函数等每条边是数据流动的方向。前向传播时数据从输入流到输出同时每个节点把自己的输入输出记下来反向传播时误差从输出端往回传每个节点根据自己记下的信息计算局部梯度再往前一个节点传。这里有个关键点很多人会忽略计算图是动态构建的。PyTorch之所以灵活就是因为每次前向传播都会重新建一张图。这意味着你可以在模型里写if-else、循环、递归只要前向能跑通反向就能自动算。但代价是每次迭代都有建图的开销这也是为什么静态图框架比如早期的TensorFlow在推理部署时更有优势——图建一次可以反复用。2.3 梯度消失与梯度爆炸反向传播的两个经典死法链式法则有个天然缺陷它是连乘。如果每一层的局部梯度都小于1乘个几十层之后梯度就趋近于0了这就是梯度消失反过来如果都大于1乘几十层就爆炸了这就是梯度爆炸。我在实际项目里遇到过这两种情况。梯度消失的典型表现是靠近输入的层参数几乎不更新loss降到一个平台就下不去了。梯度爆炸则是loss突然变成NaN或者参数值瞬间变得巨大。对于深层网络这两个问题几乎是必然要面对的。常见的应对手段有这么几个。残差连接让梯度可以走捷径不用每层都连乘归一化层BatchNorm、LayerNorm把每层的输出分布拉回标准范围间接控制梯度大小梯度裁剪在梯度超过阈值时按比例缩小防止爆炸。这些手段在大模型里几乎是标配Transformer里每个子层都有残差连接和LayerNorm不是没有道理的。注意梯度裁剪的阈值不是越大越好。设得太大会失去保护作用设得太小会让正常的大梯度也被砍掉导致训练变慢。一般从1.0开始试根据loss曲线调整。3. 梯度下降的家族谱系从SGD到AdamW3.1 最朴素的SGD为什么它慢但依然有人用随机梯度下降SGD的逻辑极其简单每次拿一个batch的数据算梯度然后参数沿着梯度的反方向走一步。步长就是学习率。它的优点是内存占用小、每次更新计算快而且有理论上的收敛保证。但缺点也很明显因为每次只看一个batch梯度方向噪声很大loss曲线会剧烈震荡遇到峡谷型的地形一个方向陡一个方向平它会来回横跳收敛很慢。那为什么现在还有人在用SGD因为在某些任务上SGD的噪声反而有助于跳出局部最优最终泛化性能可能比自适应方法更好。另外在分布式训练里SGD的通信量最小扩展性最好。所以不是SGD不行是要看场景。3.2 动量法给梯度下降装上惯性SGD的震荡问题一个直观的解法是加动量。想象一个球从山坡滚下来如果每次都只看当前坡度决定方向遇到小坑就卡住了但如果球有惯性它会带着之前的速度冲过小坑。动量法的做法是维护一个速度变量v每次更新时让v等于“之前的速度乘以一个衰减系数加上当前的梯度”。这样梯度方向一致的维度会加速方向反复变化的维度会相互抵消。实际效果就是收敛更快、震荡更小。动量系数一般设0.9这个值在大多数任务上都比较稳。3.3 自适应学习率Adam和它的变体们SGD和动量法都有一个共同问题所有参数共用一个学习率。但实际中有些参数需要大步伐有些需要小步伐。Adam的思路是给每个参数单独维护一个学习率根据它历史梯度的大小自动调整。具体来说Adam同时维护梯度的一阶矩均值和二阶矩方差的指数移动平均然后用二阶矩的平方根来归一化学习率。梯度大的参数学习率自动变小梯度小的参数学习率自动变大。再加上偏差修正训练初期也不会因为矩估计不准而跑偏。Adam的默认参数是学习率0.001、beta10.9、beta20.999、epsilon1e-8。这套参数在大多数任务上都能work所以成了大模型训练的默认选择。但Adam也有问题二阶矩的估计在训练后期可能不准导致学习率震荡。AdamW把权重衰减从梯度更新里拆出来单独做解决了L2正则和Adam自适应学习率相互干扰的问题现在是大模型微调的首选优化器。优化器核心思想适用场景典型学习率SGD沿梯度反方向走分布式训练、追求泛化0.01-0.1SGDMomentum加惯性冲过小坑计算机视觉0.01-0.1Adam每参数自适应学习率大多数任务默认0.001AdamWAdam解耦权重衰减大模型微调1e-5到1e-43.4 学习率调度什么时候该踩刹车学习率设成固定值就像开车一直踩同样力度的油门。上坡时不够力下坡时又冲太快。学习率调度就是根据训练进度动态调整学习率。最常见的策略是预热余弦退火。预热是在训练最开始用很小的学习率逐步升到设定值避免一开始梯度噪声太大把参数带偏。余弦退火是训练后期让学习率按余弦曲线慢慢降到接近0让模型在最优解附近精细搜索。大模型训练几乎都用这套组合预热步数一般是总步数的1%到5%具体看batch size和任务难度。另一个实用技巧是梯度累积。显存不够放不下大batch时可以分多次前向反向把梯度累加起来再更新一次。效果等价于大batch但显存占用小。这里有个坑梯度累积时学习率要不要跟着放大我的经验是如果只是为了让显存跑得动学习率保持不变如果是为了模拟更大batch的训练动态可以适当放大但不要线性放大一般按sqrt缩放比较稳。4. 实操从零手写一个反向传播和梯度下降4.1 用NumPy实现一个两层网络光看公式容易飘我习惯用NumPy手写一遍把每个矩阵的维度、每个梯度的来源都搞清楚。下面是一个两层全连接网络做回归的完整实现核心就是前向传播、计算损失、反向传播、更新参数四步。import numpy as np # 数据100个样本每个10维特征 np.random.seed(42) X np.random.randn(100, 10) y np.random.randn(100, 1) # 网络结构10 - 32 - 1 W1 np.random.randn(10, 32) * 0.01 b1 np.zeros((1, 32)) W2 np.random.randn(32, 1) * 0.01 b2 np.zeros((1, 1)) lr 0.01 for epoch in range(1000): # 前向传播 z1 X W1 b1 a1 np.maximum(0, z1) # ReLU z2 a1 W2 b2 loss np.mean((z2 - y) ** 2) # 反向传播 dz2 2 * (z2 - y) / X.shape[0] dW2 a1.T dz2 db2 np.sum(dz2, axis0, keepdimsTrue) da1 dz2 W2.T dz1 da1 * (z1 0) # ReLU导数 dW1 X.T dz1 db1 np.sum(dz1, axis0, keepdimsTrue) # 梯度下降更新 W1 - lr * dW1 b1 - lr * db1 W2 - lr * dW2 b2 - lr * db2 if epoch % 200 0: print(fEpoch {epoch}, Loss: {loss:.4f})这段代码里dz2是损失对第二层输出的梯度dW2是损失对第二层权重的梯度da1是损失对第一层激活输出的梯度dz1是损失对第一层线性输出的梯度。每一步都是链式法则的具体展开。跑一遍你会发现loss从初始的1.0左右降到0.01以下说明反向传播和梯度下降确实在正常工作。4.2 关键维度检查反向传播最容易出错的地方手写反向传播时维度对不上是最高频的错误。矩阵乘法里(m,n) (n,p) (m,p)反向传播时每个梯度的维度必须和对应参数的维度一致。比如dW2的维度必须是(32,1)和W2一样db2的维度必须是(1,1)和b2一样。我自己的检查习惯是写完一行梯度计算立刻在纸上标出每个矩阵的维度确认乘法合法、结果维度正确。这个习惯帮我省了无数调试时间。另外数值梯度检验也是验证反向传播正确性的好方法对某个参数加一个很小的扰动看loss变化多少和解析梯度对比。如果相对误差在1e-6以下说明反向传播实现正确。4.3 学习率调参实战从loss曲线读出问题学习率是梯度下降里最需要调的参数。我总结了一个从loss曲线判断学习率是否合适的经验法则loss震荡剧烈、不下降学习率太大参数在最优解附近跳来跳去。试着除以10。loss下降极慢、曲线几乎平学习率太小每次更新走得太近。试着乘以3到10。loss先降后升学习率在后期太大把已经找到的好参数又踢飞了。加学习率衰减。loss变成NaN梯度爆炸了。先检查学习率是不是太大再加梯度裁剪。实际调参时我一般先用一个较大的学习率跑几百步看loss有没有下降趋势如果有再逐步减小找到下降最快又不震荡的那个值。这个过程叫学习率范围测试比盲猜高效得多。5. 大模型时代的反向传播显存、精度与分布式5.1 激活值重计算用时间换显存大模型反向传播最大的瓶颈不是计算量而是显存。前向传播时每一层的激活值都要存下来反向传播时才能算梯度。一个几十层的Transformer激活值占的显存可能比模型参数本身还大。激活值重计算也叫梯度检查点的思路是前向传播时只存少数几个关键层的激活值反向传播需要用到中间激活值时从最近的检查点重新前向算一遍。这样显存占用大幅降低代价是计算量增加约30%。在大模型训练里这个交换几乎总是值得的因为显存不够根本跑不起来多花点时间至少能跑。5.2 混合精度训练FP16和BF16怎么选混合精度训练是另一个省显存、加速计算的手段。核心思想是前向和反向用16位浮点数算参数更新用32位浮点数做。这样显存占用减半计算速度提升同时保持参数更新的精度。FP16和BF16的区别在于动态范围。FP16的指数位少能表示的数值范围窄容易溢出或下溢所以需要损失缩放——把loss放大一个系数反向传播后再缩回来。BF16的指数位和FP32一样多动态范围大基本不会溢出但尾数位少精度略低。现在大模型训练更倾向用BF16省去了损失缩放的麻烦稳定性更好。5.3 分布式训练中的梯度同步当模型大到一张卡放不下时就要用数据并行或模型并行。数据并行里每张卡有一份完整的模型副本各自算自己那部分数据的梯度然后AllReduce把所有卡的梯度求平均再各自更新参数。这里有个细节梯度同步的通信量等于模型参数量。一个70亿参数的模型每次迭代要同步70亿个浮点数通信开销巨大。所以实际中会用梯度累积减少同步频率或者用通信压缩比如量化梯度降低通信量。这些优化手段的底层逻辑还是梯度下降——只不过把“算梯度”和“用梯度”拆到了不同的设备上。提示分布式训练时如果各卡的loss差异很大先检查数据划分是否均匀再检查梯度同步是否正常。有时候是某张卡的梯度没参与AllReduce导致参数更新不一致。6. 常见问题排查速查表现象可能原因排查方向解决方法loss变NaN梯度爆炸检查梯度范数降低学习率、加梯度裁剪loss不下降学习率太小或太大打印梯度范数和参数更新量做学习率范围测试靠近输入的层不更新梯度消失检查各层梯度范数加残差连接、用LayerNorm训练loss降但验证loss升过拟合对比训练和验证曲线加正则、早停、增数据显存溢出激活值太大看显存占用分布激活值重计算、混合精度多卡训练速度不升反降通信瓶颈看GPU利用率和通信时间梯度累积、通信压缩微调后模型输出重复学习率太大检查微调前后参数变化降低学习率到1e-5级别这张表里的每一条我都在实际项目里遇到过。最想强调的是梯度范数监控——很多人只看loss不看梯度。其实梯度范数是更早的预警信号loss还没出问题的时候梯度范数可能已经在异常波动了。养成打印梯度范数的习惯能帮你提前发现很多问题。7. 几个让我少走弯路的实操心得第一不要迷信默认学习率。Adam的0.001在小型任务上通常没问题但大模型微调时往往需要更小的学习率1e-5到5e-5是常见范围。我一般会先用1e-5跑几百步看loss下降情况再决定要不要放大。第二梯度累积的步数不是越多越好。累积步数太多相当于batch size太大可能会降低模型泛化能力。而且累积期间参数不更新训练效率会下降。一般累积4到8步就够了除非显存实在紧张。第三反向传播的正确性值得花时间验证。尤其是自己实现新层或者自定义损失函数时数值梯度检验是必须做的。我见过太多人因为反向传播写错训练了几小时才发现模型根本没在学。第四学习率预热在大模型训练里几乎是必须的。模型参数随机初始化时梯度方向噪声很大直接上大学习率容易把参数带偏。预热几百步再升到目标学习率训练稳定性会好很多。第五监控梯度范数比监控loss更有前瞻性。loss是结果梯度是原因。梯度范数突然变大往往预示着loss即将爆炸梯度范数趋近于0说明模型已经学不动了。把梯度范数加到训练日志里长期看能省很多排查时间。这些经验没有什么高深的理论都是实际跑模型时一次次翻车换来的。反向传播和梯度下降的数学原理几十年前就定型了但怎么用好它们怎么根据具体任务调参、排错、优化这些才是真正区分“能跑通”和“跑得好”的地方。
阅读完成 · 觉得有帮助?
咨询建站