做深度学习训练的工程师,大概率都见过下面这种画面: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,模型对该样本的预测概率几乎为0,log 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_{t+1} = theta_t - lr * m_t / (sqrt(v_t) + eps)这里m是一阶矩估计,v是二阶矩估计。Adam为什么在很多任务上比SGD稳?因为它用v做了自适应缩放:梯度大的方向步长变小,梯度小的方向步长变大。但问题恰恰出在这个自适应上。
v对梯度平方的移动平均,beta2默认是0.999,意味着它需要用大约1000步才能“消化”一个梯度平方的突变;m对梯度的移动平均beta1是0.9,大约10步就能响应。当训练已经收敛,平时梯度都贴近0,v很小。假设某一步突然来一个较大的梯度g,m会迅速被推向g的方向,而v因为惯性大还停留在小数值上,两者相除,得到的更新量会被放大到远超|lr*g|的水平。这时候Adam的一步,实际效果可能顶得上普通SGD的几十步。
这里给一个估算例子。正常训练中,梯度尺度在0.1量级,v大约在0.01量级,sqrt(v)约等于0.1,模型参数已经收敛所以梯度方向随机,m/sqrt(v)的典型值大约在0.1左右。如果某一步突然遇到一个g=10的异常梯度,由于beta2=0.999,v几乎还停留在0.01附近,sqrt(v)约等于0.14,而m已经跳到1的量级,m/sqrt(v)约等于7,比正常水平大了近百倍。注意这里还没有算lr,lr只是统一缩放,并不会改变“这一步相对于其他步的倍数”。
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_avg(m)和exp_avg_sq(v)这两个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直接变inf;inf经过exp变成NaN,loss变成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_scale=2**14,牺牲一点小梯度的精度,换取更低的溢出风险。
在代码里记录scaler的当前值很简单:
scaler = torch.cuda.amp.GradScaler(max_scale=2**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尖峰时,不慌,知道自己该先看哪里、再修哪里。