ARTICLE DETAIL

资讯详情

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

深度学习训练中loss spike的成因排查与解决策略

深度学习训练中loss spike的成因排查与解决策略 做深度学习训练的工程师大概率都见过下面这种画面loss曲线前期一路下行训练已经跑了一两万步眼看就要收敛突然在某个step冒出一个尖峰loss从0.08直接跳到3.7。你还没来得及截图下一个step又跌回0.09。如果只是这样倒也罢了最怕的是尖峰之后曲线再也回不来甚至直接变成NaN。这个现象业内通常叫loss spike也就是常说的训练快收敛时loss暴增。今天这篇就把这个现象一次性聊透它为什么会发生、和优化器有什么关系、数据侧和数值侧各自扮演什么角色以及真碰到了该按什么顺序排查和处理。不管你是刚入门深度学习、正在跑开源模型还是已经在做大模型训练这篇文章的思路都能直接用上。1. 先还原现场loss尖峰到底长什么样为什么总在“快成功”时出现1.1 三种结局自动恢复、漂移、NaN训练曲线上的尖峰看起来都是向上突一下但结局天差地别。我见过的最多的情况是单step尖峰loss在一步之内从0.08跳到3.7下一步又回到0.07画在图上就像心电图上一个孤立的毛刺。第二种是尖峰之后进入一段高loss震荡区可能持续几百步才慢慢回到原水平这种相对麻烦因为这段时间内模型参数已经被推向一个不太好的区域。最坏的是第三种loss一路冲到几十甚至变成NaN之后再怎么训练都回不来只能从checkpoint恢复。先说结论方便你后面带着问题看单step尖峰多数可以忽略它通常是某个离群样本或梯度噪声造成的持续几百步的震荡需要检查优化器状态和学习率变NaN的情况几乎可以锁定是数值溢出或权重崩溃。后面几章会分别拆开讲。1.2 为什么收敛期更容易察觉尖峰为什么这类尖峰总在“快收敛”时出现一个重要原因是相对尺度。训练前中期loss普遍在2到5的范围出现0.5的波动你不会太在意到了收敛期loss已经压到0.1左右同样0.5的波动在图上就是一根刺眼的长针。所以有一部分尖峰其实是尺度效应不是模型真的崩了。但收敛期还有一个更本质的问题此时学习率通常已经被退火得非常小按理说模型应该更稳为什么还会出现大尖峰这就牵出优化器内部机制了。很多人碰到spike第一反应是“是不是学习率太大”先减lr结果减完尖峰依然在只是幅度变小了。这说明真凶不止学习率一个。我去年训练一个6层的Transformer做文本分类跑到第120k步loss已经稳定在0.02附近训练准确率到了98%以上。突然在121k步loss变成0.89吓得我赶紧看TensorBoard。这步的lr只有初始lr的0.05倍按常理根本不该发生这么大的更新。后来查出来是某个batch里混进了一条label被错标成24的样本——类别数只有12模型对该样本的预测概率几乎为0log loss直接爆了。这个案例恰好说明收敛期的spike优化器机制和数据质量各占一半责任。2. 从Adam的记账本说起优化器状态才是尖峰的真正导火索2.1 Adam更新公式里的“隐藏杠杆”要理解loss尖峰必须回到优化器的更新公式。以最常用的Adam为例m_t beta1 * m_{t-1} (1 - beta1) * g_t v_t beta2 * v_{t-1} (1 - beta2) * g_t^2 theta_{t1} theta_t - lr * m_t / (sqrt(v_t) eps)这里m是一阶矩估计v是二阶矩估计。Adam为什么在很多任务上比SGD稳因为它用v做了自适应缩放梯度大的方向步长变小梯度小的方向步长变大。但问题恰恰出在这个自适应上。v对梯度平方的移动平均beta2默认是0.999意味着它需要用大约1000步才能“消化”一个梯度平方的突变m对梯度的移动平均beta1是0.9大约10步就能响应。当训练已经收敛平时梯度都贴近0v很小。假设某一步突然来一个较大的梯度gm会迅速被推向g的方向而v因为惯性大还停留在小数值上两者相除得到的更新量会被放大到远超|lr*g|的水平。这时候Adam的一步实际效果可能顶得上普通SGD的几十步。这里给一个估算例子。正常训练中梯度尺度在0.1量级v大约在0.01量级sqrt(v)约等于0.1模型参数已经收敛所以梯度方向随机m/sqrt(v)的典型值大约在0.1左右。如果某一步突然遇到一个g10的异常梯度由于beta20.999v几乎还停留在0.01附近sqrt(v)约等于0.14而m已经跳到1的量级m/sqrt(v)约等于7比正常水平大了近百倍。注意这里还没有算lrlr只是统一缩放并不会改变“这一步相对于其他步的倍数”。2.2 为什么调低学习率不能根治既然异常更新来自m/sqrt(v)的比例失衡那么调低lr只是整体把所有步都缩小异常步虽然也会变小但正常步同样变小。如果原本正常步长已经很小——收敛期本来就是如此——再调低lr训练几乎就推不动了。也就是说你只是把这个尖峰按小了一点并没有消除它而模型在尖峰时受到的实际扰动依然是正常步长的好几倍。所以正确思路不是“爆了就减lr”而是弄清楚比例失衡的根源。这也是为什么很多大模型训练框架在后期不是单纯靠低lr来稳定loss还会配合梯度裁剪、调节beta2、增大eps等手段。它们的目标都是限制m/sqrt(v)的上限。如果你只想改一个超参数我建议优先动eps而不是lr原因后面第6章会详细讲。2.3 残留风险AdamW与weight decay现在大家普遍用AdamW比原始Adam多了解耦的weight decay。权重衰减项会在每步直接把权重向0拉一点。在正常阶段这没啥问题因为lr小一旦发生spike这步的权重更新很大weight decay会产生一种“这次更新把权重推向一个偏离点然后下一次又往0拉”的震荡感。实际表现就是尖峰之后模型要花几百步把权重“拉回来”但下游评估指标可能已经明显掉了。我自己的经验是遇到尖峰时优先检查的并不是lr而是exp_avgm和exp_avg_sqv这两个optimizer state的norm。把尖峰前后的optimizer state dump出来对比往往会看到exp_avg_sq没有跟上而exp_avg已经异常大。看到这个就可以确定是自适应比例失衡往Adam的eps、beta2方向修才对症。3. 数据批次的隐藏雷区loss暴增不一定是模型的问题3.1 收敛期的“陌生样本”效应优化器机制解释了尖峰如何被放大但没有解释那个异常梯度从哪来。大多数时候它来自数据批次。模型在收敛后对训练分布内样本的loss已经压得很低梯度也很小。这批样本可以看成“熟面孔”。此时任何一条“陌生样本”——一条标注错误的、特征异常的、或者训练分布里本来就稀有的数据——都会产生一个比正常样本大得多的梯度。在loss图上它表现为一个孤立的尖峰在优化器内部它触发了第2章说的比例失衡。所以数据和优化器其实是上下游关系数据提供“火药”Adam负责“点火”。3.2 数据管线的几类典型脏输入长时间训练的脏数据来源很杂我自己归过类最常见的是下面几类文本类任务label越界、token id被错误替换成异常值、序列截断后变成全padding、mask位置错误。这类问题在NLP预训练里最隐蔽因为很多错误不报异常只是默默把loss算大。图像类任务解码失败的图直接进模型、归一化后出现inf、某些数据增强操作产生除0。比如一张损坏的jpeg解码出来是一块纯白噪点模型对它几乎随机预测softmax输出接近均匀分布loss自然比其他正常样本高出一个量级。多模态和语音任务采样点数不齐、静音片段被当成有效样本、音频和文本的对应关系错位。采样器和分布式加载多进程shuffle种子设置不一致导致同一份数据被重复采样多个epoch里某些样本一直没被看到最后集中出现在某个batch里模型措手不及。这里还要补充一种很容易被忽略的情况周期性尖峰。如果你发现尖峰不是随机出现的而是每个epoch固定出现在同一个位置或者每隔固定步数出现一次先别查模型直接把采样器、数据加载器相关的代码翻出来。按长度排序的分桶策略非常容易造成这种现象一个桶里全是长文本下一个桶里全是短文本模型在跨桶的那一步梯度分布会发生明显变化反映在loss上就是小尖峰。另外如果某个困难样本没有做去重它会在每个epoch被反复抽到。模型第一次遇到它时loss高、产生尖峰第二次、第三次依然高只不过由于模型在上一轮尖峰后已经改变它的loss可能没那么极端。于是你看到的就是每隔固定步数出现一次小尖峰。检查方式是把loss曲线的横坐标对epoch取模看尖峰是否对齐到同一个位置。3.3 用固定种子复现法快速定位脏样本数据问题的排查其实有一套很成熟的流程核心就四个字固定种子。第一步在训练脚本里固定所有随机种子包括Python、numpy、CUDA、dataloader worker以及任何影响数据顺序的地方。第二步从spike前的checkpoint继续训练只跑一次看spike是否在同一step复现。如果能复现直接定位到该step的batch id把数据单独取出来跑一次前向和loss逐个样本算loss异常样本基本当场现形。如果不能复现说明是分布式并行或GPU非确定性造成得往all-reduce和cudnn benchmark方向查。我遇到过的一个很典型的情况是某个图像分类项目loss每到特定步数就涨一次复现后把那个batch单独拎出来发现是一条损坏的jpeg解码后变成了全零图。模型对全零图几乎等于在猜输出均匀分布loss自然高。把这个样本过滤掉以后整个训练阶段再没出现过周期性尖峰。所以数据侧问题一定要优先排除因为它排查成本最低命中率又最高。4. 数值稳定性暗坑混合精度、梯度裁剪与inf/nan的三方角力4.1 inf/nan的传播链数据问题通常只会造成大loss不会直接让训练永久性崩溃真正把尖峰变成不可恢复NaN的是数值稳定性问题。传播链一般长这样某个step出现异常大的梯度更新之后某些层的输入值尤其是attention里的score、softmax之前的值变成很大的正数或负数fp16下超出65504直接变infinf经过exp变成NaNloss变成NaN反向传播出的梯度也全是NaN下一步权重直接全NaN模型原地报废。这里最关键的是NaN一旦出现不干预就永远恢复不了。因为NaN参与任何运算都会继续传播checkpoint如果没有保存好就得回退到几天前那损失就不是一两个step的问题了。顺便提一句如果你用的混合精度是BF16因为指数范围和FP32一致几乎不会出现inf所以BF16训练中loss变成NaN的概率要低很多。但BF16尾数只有7位loss在收敛期会表现为小幅抖动看起来很像轻微的不稳定这是精度问题不是尖峰。两种问题别混为一谈。4.2 AMP动态loss scale的隐患AMP里loss会被乘上一个scaler再在fp16下做反向传播。scaler会自动变大变小以适配梯度范围。初始scaler通常很小比如128随着训练稳定会一路涨到2^15甚至2^16。训练越接近收敛loss和梯度越小scaler就越倾向于涨到高位以保留梯度精度。可问题来了——scaler越大某个异常梯度被“撑爆”成inf的概率也越大。很多同学看到训练后期loss突然变成NaN第一反应是调lr或者回滚checkpoint其实正确的第一步是去看scaler的日志。如果你在训练脚本里记录过loss scale会看到NaN出现的那一步loss scale正处于峰值附近。处理方式不是把scaler关掉而是给训练脚本加一个保护当scaler连续多次detect overflow时触发告警并保存现场。也可以把scaler的最大上限设低一点比如max_scale2**14牺牲一点小梯度的精度换取更低的溢出风险。在代码里记录scaler的当前值很简单scaler torch.cuda.amp.GradScaler(max_scale2**14) current_scale scaler.get_scale()把它写进每N步的日志里你就能在事后复盘时快速判断spike到NaN是不是scaler溢出这条链路。4.3 梯度裁剪阈值设多少是个学问梯度裁剪是防尖峰扩散的标配但很多人设阈值完全是拍脑袋。设成1.0训练前期梯度过大被剪得只剩形状设成10.0后期小梯度时代它压根不触发等于白设。更合理的做法是在训练开始后的前几百步记录grad norm的分布然后取P99作为裁剪阈值。这样既不会频繁限制正常更新又能拦住真正的离群梯度。另外要注意裁剪后再做梯度noise或weight decay的顺序不能乱PyTorch里一般是optimizer.step()前调用clip_grad_norm_。如果你配了gradient accumulation需要先等梯度累加完再裁剪千万别每个micro-step都裁剪否则梯度量级会被低估等累加完已经超过阈值了。我个人的习惯是同时监控grad norm和update norm。很多人都见过grad norm尖峰然后自己恢复但很少人记录update norm也就是参数实际被更新的幅度。一旦发现update norm也出现尖峰说明不是clip没拦住就是optimizer状态失衡如果update norm正常那loss尖峰大概率只是某个batch的logits波动对训练影响很小。5. 实战排查链路从checkpoint回溯到逐层梯度观测5.1 五步排查法排查的目的不是猜原因而是用最短时间把根因锁定在四个域之一数据、优化器、数值、分布式。我自己的排查顺序基本固定成五步。第一步保存现场。出现spike时不打断训练但要立刻dump当前step的model权重、optimizer state、lr、grad norm、batch数据的hash以及scaler状态。很多框架支持信号回调比如注册SIGUSR1 handler在任意异常时刻保存快照。没有这些后面分析就是无米之炊。第二步固定种子复现。从spike前的checkpoint继续训练但把所有随机种子固定住。如果尖峰能在同一step复现那基本可以排除GPU非确定性和分布式顺序的干扰问题就在数据或模型本身如果换一台机器后复现不出来那多半是cudnn benchmark或原子操作这类非确定性在捣乱。第三步查数据管线。这部分其实是五步里性价比最高的因为有大量spike最后都落在数据上。从复现出来的batch id出发把那个batch单独拎出来跑一遍逐个样本看loss异常样本基本当场现形。可以把所有样本的loss降序排列看前几名是不是标注错误、解码损坏、增强异常。第四步梯度归因。如果数据查不出问题给模型挂backward hook打印每层参数的grad norm。通常最先出问题的层很有规律transformer里是最后一层head和embedding层因为它们的梯度对logits和token embedding的敏感度最高。也可以用torch.autograd.detect_anomaly()跑一遍定位到产生NaN的算子。第五步看数值状态。汇总scaler日志、grad norm分布、是否有inf/NaN。这一步可以回答尖峰是不是被数值放大成了永久崩溃。5.2 分布式训练下的额外嫌疑如果你用的是数据并行或模型并行还有一个独立变量多个rank之间的交互。数据并行里最常见的是all-reduce污染某个rank上的一条脏数据产生超大梯度所有rank的梯度一汇总大家都会受影响。而由于是跨卡通信你本机上的fix可能根本解决不了别的卡的数据问题。检查方法很简单——分别打印每个rank的grad norm看哪个rank在spike step异常高。另一个容易踩的坑是不同进程的dataloader shuffle seed设置不一致导致同一个batch在多个rank间重复或错位这也会在收敛期制造尖峰。模型并行和流水线并行里则要关注batch间的层间通信、norm更新使用的全局unscale、clip的位置以及loss在micro-batch上的统计方式。偶尔loss spike只是某个micro-batch的计算顺序变了梯度累积后表现并不相同。这时候把micro-batch数量调成1试一下能很快确认是不是这个问题。5.3 一个完整的现场复盘案例拿我之前训练的一个开源中文BERT变体来说。训练到第220k步loss已经到1.62第221k步突然变成3.74之后3步回到1.6表面看问题不大。但第230k步又跳到8.2然后直接NaN。我们当时的操作链是这样的先从第229k步的checkpoint恢复固定种子跑了一次在第230k步复现了NaN。挂上detect_anomaly报错直接指向某个FFN层Function AddmmBackward0 returned nan values in its 0th output。当时第一反应是PyTorch版本问题但检查grad norm日志后发现了关键线索在NaN出现前200步grad norm从0.3慢慢爬升到2.1然后某一瞬变成inf。再往上游查发现是自定义的KL散度损失函数里没有处理target为全0的边界某个batch恰好连续采到了16条同一类别的样本模型预测和target几乎一致KL散度计算的log(0)产生-inf反向传播直接变成NaN。修复方式是给损失函数加一个数值稳定分支当target全0或预测概率为0时直接返回0同时调整采样器避免单个batch类别过于集中。修复后同样的配置跑了300k步一次尖峰都没再出现。这个案例想说明的是spike到NaN之间往往隔着一个看似不起眼的数值边界。数据侧负责“异常输入”损失函数负责“数值爆炸”优化器和scaler负责“把爆炸传播下去”。排查时缺了任何一环都会觉得是另一个玄学问题。6. 组合拳治本warmup、梯度裁剪、动态scaler与超参数修正6.1 性价比最高的四件套以下是我个人在多个训练任务里验证过、性价比最高的四个措施按优先级排。第一数据最小过滤。过滤掉任何能判定的损坏样本。更主动一点的做法是统计每个batch内部样本的loss分布若某个样本的loss稳定超过同batch中位数的5倍以上在每个epoch结束后做一次困难样本筛查。但注意不要过度过滤否则会让模型见过的分布变窄泛化反而下降。第二合理的梯度裁剪。设置基于grad norm历史分布的P99阈值而不是拍脑袋值。对需要高稳定性的训练可以再叠加一个clip_grad_value_比如单参数最大梯度设为5防止单个参数一步更新过大。第三增大Adam的eps。从默认1e-8调大到1e-4甚至1e-3可以给自适应步长加一个“下限地板”。这个技巧在NLP预训练里非常常用能显著减少收敛期尖峰。代价是训练前期有效lr会被稍微压一点可以配合延长warmup来弥补。第四自动回滚机制。检测到尖峰超过一定阈值且持续N步不恢复时自动加载spike前的checkpoint并把lr乘子设为0.3左右继续训练。实现很简单但前提是checkpoint要保存得足够频繁并同时保存optimizer state和scaler state。6.2 尖峰后的响应策略降lr、局部warmup还是直接回滚很多人问spike出现以后我是该立刻停训调参还是让它跑一步看看我的判断标准很朴素看尖峰后N步的loss能不能回到尖峰前低点的1.5倍以内。能就让它跑不能就回滚。细分下来策略是这样的单step尖峰且下一step就恢复忽略记录日志即可。尖峰后持续震荡先降lr到0.3倍或0.5倍观察500步如果恢复慢则回滚到尖峰前checkpoint再从那里开始用更低的lr训练。变成NaN直接回滚到NaN前最近的checkpoint同时检查scaler日志和grad norm确认根因后再继续。还有一种手段叫局部warmup。回滚或降低lr之后可以不要把lr维持在一个低常数而是在低lr的基础上线性回升到目标值模拟训练开始时的warmup。比如从0.1倍lr起步500步之内回到正常值。这给优化器一个重新适应的时间可以防止回滚后立刻再次spike。这里有一个训练逻辑的简单示意if spike_detected and steps_since_spike 10: load_checkpoint(best_checkpoint) new_lr base_lr * 0.3 rollback_step 0 if rollback_step 500: lr new_lr (base_lr - new_lr) * (rollback_step / 500)核心思想是把回滚之后的训练看成一个“小规模重启”而不是直接回到原来的正常状态。这个细节很多工程师会忽略结果回滚后第二个尖峰马上又来了其实是optimizer把之前的异常状态也一起加载回来了。6.3 监控与自动干预的最小实现日志是排查一切的基础。至少应该记录step、loss、grad_norm、update_norm、lr、loss_scale、每秒处理样本数。少于这些出了问题只能靠猜。我在长时间训练时都会写一个几十行的loss guard逻辑很简单维护一个“历史最低loss”变量如果当前loss超过历史最低的3倍且连续10步没有回落到1.5倍以内就触发告警并自动做两件事——保存当前现场、从最佳checkpoint恢复并降低lr。这个机制让我至少避开了两次数十卡时级别的损失。伪代码大致长这样best_loss float(inf) spike_threshold 3.0 recover_ratio 1.5 streak_limit 10 if loss best_loss: best_loss loss streak 0 elif loss best_loss * spike_threshold: streak 1 else: streak 0 if streak streak_limit: save_current_state(...) load_checkpoint(last_good_checkpoint) optimizer.param_groups[0][lr] * 0.3 streak 0这只是一个最小实现你也可以把条件改成“连续N步loss都比历史最低高2倍”或者“grad norm超过P99”等等按你的任务特点来。但有一条经验是通用的loss回到低位不代表模型没受伤。如果你在下游任务上有eval指标尖峰后一定要观察eval curve。有时候loss曲线已经恢复了但F1掉了2个点这时候就不是“跑跑就好”能解释的了强烈建议回滚。最后说个我的个人习惯训练脚本里永远会保留最近三个checkpoint每个都同时保存optimizer state和scaler state。spike这种事有点像开车遇到路面坑大部分时候颠一下就过去了但偶尔会爆胎。多留几个checkpoint多配一段监控日志成本就是几十GB磁盘收益是你在几天的训练结束后不会因为一次NaN全部重来。我现在的原则是宁可多花十分钟排查绝不多烧一天的GPU。希望这篇能帮你在下一次看到loss尖峰时不慌知道自己该先看哪里、再修哪里。
返回列表