ARTICLE DETAIL

资讯详情

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

深度学习训练快收敛时loss突然暴增:诱因分析与排查指南

深度学习训练快收敛时loss突然暴增:诱因分析与排查指南 如果你跑过分类、检测或者大模型微调八成遇到过这种场景loss明明已经跌到很低训练眼看要收敛了突然一个step跳上去一个数量级然后又稀里糊涂地降回来。上礼拜我还帮同事排查了一个SegNet的训练任务第47个epoch loss都到0.15了一个跳变直接顶到1.9再往后才慢慢回落。这种“快收敛时loss突然暴增”的现象在深度学习模型训练里比随机初始化时的发散更让头疼因为整个训练流程看起来都没毛病就是某一步坏了事。这篇文章我会按实际排查思路来写先讲清楚为什么偏偏在“快收敛”阶段容易出这种幺蛾子再逐个拆解常见诱因然后给一套可以直接抄的代码级排查流程最后是预防手段和个人的一些经验。适合正在被“训练不收敛”折磨的算法工程师、在啃深度学习的同学也适合拿halcon这类深度学习工具做训练、被loss曲线折腾过的工业视觉从业者。1. 先搞清楚“快收敛时暴增”为什么值得单独讲1.1 正常的loss曲线到底长什么样很多新手觉得loss曲线应该是一条平滑的下降直线这认知从一开始就跑偏了。实际上随便拿ResNet在CIFAR-10上跑一遍记录每个step的loss你会看到一条像心电图一样的线——整体往下走但局部有很多小尖刺。这是因为我们用的是mini-batch随机梯度下降每个batch是全体数据的一个随机子集梯度本身就是带噪声的估计loss自然也跟着抖。关键是怎么区分“正常抖动”和“异常暴增”。我的经验是看两条东西一是尖峰的高度二是尖峰之后是否回落到原水平。正常抖动一般在一个很小的幅度范围内波动比如整体趋势在0.2到0.1之间下降时单个step跳到0.4、0.5不罕见但很快会跌回趋势线附近。异常暴增则是尖峰的高度明显超出正常波动的几倍甚至几十倍而回落后又可能比之前还高一点或者干脆再也不回来了。val loss的情况还不太一样。如果train loss正常下降val loss在某个点开始持续走高这是过拟合不是我们今天聊的“暴增”。真正要处理的是train loss本身在快收敛阶段忽然跳上天这意味着优化过程出了实质性问题不能简单用“正常波动”来解释。1.2 为什么偏偏是“快收敛”这个时间点最容易出事这个问题你想明白了排查方向就清晰了。核心原因有三层第一层是学习率和损失函数地形的关系。我们在用梯度下降时学习率的上限其实由损失函数在当前点的“平滑程度”决定。越接近局部最优点损失函数往往越像一个狭长的山谷不同方向上的曲率差异非常大——数学上就是Hessian矩阵的条件数变大。在这种地形里之前能稳定收敛的学习率现在可能已经超过该方向允许的上限等于每走一步都会过头形成振荡甚至发散。就好比你开着一辆车在开阔公路上可以把速度开到120进了盘山窄路还保持120那必然冲出弯道。第二层是后期梯度噪声的比例急剧变大。训练初期模型离目标远梯度方向一致性很强每个batch算出来的梯度都大致朝同一个方向相当于“信号强度大、噪声小”。到了后期模型已经接近最优解真正的梯度信号本身已经非常小但mini-batch采样引入的随机噪声并没有同比例下降。此时一个包含难样本、噪声标注或者类不均衡极端样本的batch其梯度方向可能完全盖过真实梯度方向直接把参数推出好的区域。第三层是训练过程中累计的状态本身出了问题比如优化器里的momentum项记了一笔“历史账”遇到后期的小梯度时这笔账反而变成惯性把参数往错误方向多推了一步。还有BNBatch Normalization在训练后期的running statistics估算已经比较稳定一旦某个batch的样本分布偏移前向传播的数值就会整体漂移造成loss瞬间飙高。这三层原因单独出现也好叠加出现也罢都决定了同一件事在快收敛阶段训练系统对外界干扰的容忍度是最低的平时“无所谓”的小问题这时候都会被放大成loss尖峰。1.3 train loss暴增和val loss上升根本不是一回事在一些交流群里我看到有人把“val loss上升”和“train loss暴增”混在一起讨论这是两个需要完全不同的处理思路的问题。val loss上升而train loss持续下降这是模型开始记住训练集特有模式、泛化能力下降的信号解决手段是加大数据增强、加正则化、降低模型容量、提前停止本质是“不要继续拟合下去”。train loss在快收敛时暴增则是优化数值层面的问题和泛化没有直接关系。你加大增强或者加正则化对于这个问题往往没什么效果。动手之前先在日志里分清是训练集的loss在跳还是验证集的loss在涨。我看过太多人在val loss上升时去调优化器或者在train loss暴增时去改网络结构最后折腾一礼拜发现方向完全错了。2. 按出现频率排序的6个诱因以及它们各自的“指纹”2.1 学习率在后期仍然偏大最常见也最容易被误判这是我最先排查的对象不是因为原理复杂而是因为太容易混淆。很多框架默认的StepLR scheduler是每30个epoch把学习率乘0.1如果设置不当在快要收敛的阶段学习率仍然保持着初始值那loss暴增几乎是必然的。判断方法其实很直接看loss尖峰是不是出现在同一位置附近、并且方向一致。比如在固定的第N个epochloss从0.2跳到1.8然后经过几十个step逐渐回到0.25再过一段时间又在同一个epoch附近重复跳。这是典型的“学习率过大导致周期性越过最优区域”的表现。处理手段也直白把学习率直接降到当前值的十分之一如果尖峰消失那基本坐实。注意这里要做一个“对照试验”不要同时改scheduler、改batch size、改数据增强不然你根本不知道是哪个改动救回来的。我自己的习惯是发生了loss暴增先不动代码只把初始学习率乘0.1跑20个epoch看曲线是否恢复平滑这个信息量足够帮你定位问题。2.2 脏数据在训练后期“爆雷”这是最容易忽视的训练初期模型还在学明显的模式对个别错误标注或异常样本不敏感。到了后期模型已经能强行拟合训练集里的绝大多数样本那些难样本和错误标注开始变成主导梯度方向的因素。比如一个分类任务里某张图片被标成错误的类别前期模型反正也分不对这个“错误监督信号”和模型自身的错误混在一起看不出来。后期模型已经把正确样本学透了这张脏样本的梯度就成了一个巨大的“对抗信号”直接把loss顶上去。怎么验证有个笨办法但非常有效在loss暴增的那个step附近暂时冻结模型参数重新加载这个batch的数据用同样的模型参数再算一次loss。如果这个batch每次都稳定地产生高loss那说明确实是这个batch内部的数据有问题。接下来把batch展开逐个样本过一遍找出loss最高的那几个单独看一下标注和图像内容。我遇到过几次发现是某张图上同时出现了两个类别目标标了一个漏了一个这种样本越训练越会成为“钉子户”。2.3 自定义loss函数里藏着一颗“地雷”很多网络不是简单用CrossEntropyLoss而是加了各种辅助项比如focal loss、logit adjustment loss、一致性正则、梯度惩罚。这些loss项在训练早期和中期可能都表现得很平稳问题往往在后期爆发。拿focal loss举例它的设计初衷是降低易分类样本的权重、让模型关注难样本。训练后期大部分样本已经非常容易分类focal loss对这些样本的梯度趋近于零整个loss的主要贡献集中在少数几个“极难样本”上。如果这些难样本里面混杂着噪声标注甚至outlier那loss项就会非常高而且由于focal loss里面的调制因子存在数值上限参数稍微偏移一点就会把loss放大很多。logit adjustment loss也有类似特征它把类别先验频率加进logit里在类别极度不均衡的数据集上非常有效。但训练后期尾部类别的logit可能被压缩得很极端一旦某个batch里尾部类别的样本占多数loss尖峰就轰然出现。这种问题排查时最直接的手段是把loss组成拆开分别记录每一项在暴增时刻的数值。你不需要在训练循环里加多少代码每次backward之前把各个loss项的item()值打个log就行。谁跳了问题就在谁身上。2.4 混合精度训练和BN在“最后一公里”拖后腿现在但凡显存紧张大家都会开混合精度训练。FP16的优点是快、省显存但缺点是数值范围比FP32小得多。训练初期梯度大不太容易出问题后期梯度变得很小很容易低于FP16能表示的最小正数变成0或者某些中间计算溢出为nan再传导到loss就炸了。如果你用了自动混合精度AMP可以观察loss暴增时是不是伴随“梯度为nan”的警告。如果是最简单的处置是把这个阶段的数值灵敏度降下来或者对关键张量改用FP32。还有一个容易忽略的点不同厂家的GPU对FP16的支持能力有差异有的卡甚至对BF16和FP16的转向支持不一样在小batch下更容易踩坑。BN的问题则集中在小batch size场景。训练早期BN的running statistics变动大batch之间的统计量差异被模型适应了后期模型已经依赖一个相对稳定的统计量来做推理结果某个batch里的样本恰好偏亮或偏暗、尺寸和位置分布偏移BN层输出的数值会整体移动导致loss曲线猛跳一下。这种情况在目标检测里特别常见尤其是训练batch size小于8的时候。2.5 优化器状态变成了“惯性事故”动量和二阶矩估计是Adam这类优化器收敛快的关键但也可能成为后期暴增的帮凶。动量项相当于给更新方向加了历史速度的惯性训练后期真实梯度很小动量积累的方向如果和当前梯度方向冲突就会出现“急刹车”甚至“掉头跑”的情况。尤其是训练中途学习率已经通过scheduler调小了但动量项还没来得及缩水参数会被“带”出去一段距离loss随之暴涨。另一个被低估的因素是分布式训练下的batch不一致。比如你用4张卡本来全局batch是256数据并行时单卡实际上是64。如果数据分布不均匀某张卡上某个step恰好拿到一堆困难样本同步梯度时这张卡的梯度分量异常大整体更新方向就会被它带着跑。小batch越大越容易踩这种问题所以多卡训练时的loss暴增先检查不同卡上的数据分布是否一致再考虑用梯度裁剪兜底。2.6 数据顺序、shuffle状态和多轮训练中的“偶然”还有一个不算少见的坑dataloader的shuffle设成了False或者多轮训练时数据顺序固定。模型会“背下来”数据的出现顺序学到一种和真实数据分布无关的时序规律。一旦某个位置的batch是整个数据集里最难的一块loss尖峰会规律地出现在固定的step位置上。这种问题非常隐蔽因为它只在固定的训练步数出现看起来像scheduler的问题其实只要把dataloader的shuffle打开、或者换一个seed重新排一下数据顺序尖峰就没了。我在工业界项目里还遇到过一种纯属偶然的情况某个batch里的样本恰好全是相似场景模型在这个batch上的输出特征高度相关导致BN统计量和梯度方向出现共振loss短暂暴涨。这种“偶发尖峰”没有任何规律重启训练不一定复现通常不用特别处理。3. 没有现成日志时的整套排查流程照着做就行3.1 第一步给训练脚本加一个“波形监护仪”如果你手头的训练脚本没有记录每个step的loss也没有保存梯度和参数范数那遇到问题就两眼一抹黑。我的建议是先用一个简单的logger把这几项记下来哪怕只是每分钟打印一行也比啥都没有强。下面这个代码片段可以直接插进训练循环作用是计算梯度范数它在定位后期loss暴增时非常好用import torch def compute_grad_norm(model): total_norm 0.0 for p in model.parameters(): if p.grad is not None: param_norm p.grad.data.norm(2) total_norm param_norm.item() ** 2 return total_norm ** 0.5 # 在loss.backward()之后、optimizer.step()之前调用 grad_norm compute_grad_norm(model) print(fstep {step}, loss {loss.item():.6f}, grad_norm {grad_norm:.6f})正常情况下训练后期grad_norm应该随loss一起变小。如果loss暴增的同时grad_norm也飙高说明是梯度爆炸类问题如果grad_norm没怎么变、loss却飙了那更可能是前向传播数值或者loss函数本身的问题。这个方向性判断能帮你省掉一半的排查时间。3.2 第二步用“黑匣子”方式定位肇事batch只知道“第47个epoch、第1034个step暴增”还不够你得能回放当时发生了什么。简单做法是在loss暴增的step先不更新参数保存模型当前状态然后重新拿这个batch的前向结果算一次loss看是不是稳定复现。如果复现就把这个batch单独存下来逐个样本排查如果复现不了说明问题和模型参数的即时状态有关和数据batch无关。具体代码可以这样处理# 在loss异常高于某个阈值时触发 if loss.item() spike_threshold: torch.save(model.state_dict(), spike_model.pth) # 保存当前batch的data和target方便回放 torch.save({data: data, target: target}, spike_batch.pt) # 暂停训练先把现场固定下来 debug_mode True这个“现场固定”的思路和航空事故后找黑匣子一样先保留证据再分析原因。不要急着改参数重启训练那会丢掉最关键的信息。3.3 第三步用“三个二分”缩小范围把问题缩小到数据、模型、优化器三个层面方法是用对比实验做二分。数据层面把数据增强全关掉、shuffle打开固定seed重新跑到暴增点。如果暴增消失说明问题大概率在数据增强或数据顺序仔细检查增强策略里是不是有什么操作在后期把样本破坏得太厉害。模型层面换一个更小的骨干网络、或者干脆暂时冻结大部分层只训练分类头跑到同样的位置看是否还暴增。如果不暴增重点检查模型结构和BN层的数值稳定性。优化器层面在同一个模型上分别试SGD、Adam、AdamW各自固定seed跑到暴增点谁的曲线崩了就说明优化器和学习率scheduler的配合有问题。这套流程看起来很土但实际效果比直接看曲线猜原因靠谱得多。每一步都只改变一个变量结果也容易解读。4. 不止是亡羊补牢预防loss暴增的几个“体操动作”4.1 给学习率一个“坡道”而不是急刹车很多框架里默认的StepLR是在某个epoch直接把学习率乘以0.1这种“断崖式”的衰减方式在模型快速逼近最优解时很可能产生阶段性的参数抖动。更平滑的做法是warmup加余弦退火。warmup的思路是在训练初期用很小的学习率跑一段时间等梯度方向统计得比较稳了再逐步加到目标学习率。余弦退火则是让学习率在每个周期内平滑衰减避免突然拐弯。这两个配合起来能非常有效地压低后期loss尖峰。我习惯用下面这个公式做余弦退火衰减到目标学习率的1/100左右import math def cosine_lr(step, total_steps, init_lr, warmup_steps0, min_lr1e-5): if step warmup_steps: return init_lr * (step 1) / (warmup_steps 1) progress (step - warmup_steps) / max(1, total_steps - warmup_steps) return min_lr 0.5 * (init_lr - min_lr) * (1 math.cos(math.pi * progress))单独说一句快收敛时暴增很多时候不是“学习率这个值不对”而是“学习率在该降的时候没来得及降”。如果你用了StepLR建议在训练后半段观察曲线趋势如果临近切换点loss已经有抬头迹象就把scheduler的切换位置提前。4.2 梯度裁剪的阈值不是随便填的梯度裁剪几乎是处理loss暴增的万金油但很多人用不明白。clip_grad_norm_的原理是算所有参数梯度组成的全局范数如果这个范数超过max_norm就按比例缩放所有梯度保持方向不变、大小压下来。PyTorch里一行代码搞定torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm1.0)max_norm设多少合适这取决于任务和loss量级。我的经验是分类任务常见1.0到5.0目标检测里因为loss通常更大可以放宽到10.0或更高大模型微调用1.0比较稳GAN这类对抗训练反而不能裁剪太狠否则判别器学不动。确定阈值的实操方法是先不裁剪打印每个epoch的grad_norm分布。假如前期grad_norm中位数是0.5偶尔冲到5.0那max_norm设1.0就是合适的选择。阈值设太小会导致有效梯度被频繁压缩训练变慢设太大则起不到兜底作用。4.3 EMA模型权重是最后的救命稻草EMAExponential Moving Average是指维护一份模型参数的滑动平均训练过程中每隔几步把当前模型参数的“一部分”融进这份平均参数里。你在推理和保存最终模型时用的是这份平均参数而不是训练时那个会抖动的参数。EMA之所以对loss暴增特别有效是因为它天然做了平滑。就算训练过程中某一步参数被坏梯度带飞了EMA里也只计入了一小部分不会受到致命影响。等到峰值过去EMA会逐渐“赶上”好的参数区域。实际调模型的时候我经常遇到训练loss曲线在最后阶段疯狂抖动但EMA版本的模型在测试集上效果依然很好续训时用EMA参数做初始化能避免很多莫名其妙的尖峰。代码实现非常简单torch.no_grad() def ema_update(model, ema_model, decay0.999): for ema_p, p in zip(ema_model.parameters(), model.parameters()): ema_p.data ema_p.data * decay p.data * (1 - decay)decay设多少要看模型的更新频率和数据集规模。小数据集、模型收敛快decay可以设0.99动辄训练几十万步的大模型设0.999到0.9999更合适。4.4 给loss组成加一道“保险丝”除了梯度裁剪另一个被忽视的预防手段是把loss里的敏感项做数值保护。比如自定义loss里有一项是梯度惩罚这一项在backward时会算二阶导数数值很容易在某些点上变得巨大。学会用detach()把不需要回传梯度的分支摘掉用clamp()限制loss项的数值范围用min()/max()避免极端值参与后续计算。mode_loss mode_loss.clamp(max10.0) # 限制单项loss的上限这种做法看起来很简单但如果不做一旦某个单项loss跳到1000整个训练过程都得陪葬。保险丝的另一个形态是对loss为nan或inf的情况做自动“熔断”直接跳过这个step的优化。虽然会浪费一点计算但比整轮训练报废要划算。if not torch.isfinite(loss): optimizer.zero_grad() continue5. 最后分享一点我的个人体会深度学习模型训练里“不收敛”和“收敛后拔刺”是两种完全不同的问题。前者通常意味着模型结构或者数据和标签对不上后者往往只是优化过程在最后阶段变得极其脆弱一点点扰动都能被放大。快收敛时loss暴增这件事绝大多数情况下不是你的模型架构不行而是训练流程里的某个环节在“刚开始不影响、临近收敛才露馅”——这可能是一个涨得不合理的学习率一道不该出现的脏标注或者一个数值不稳定的自定义loss项。所以我的排查顺序一直没变过先看scheduler再看数据再看loss组成最后看数值精度。如果你连日志都没有第一步永远是先加上梯度范数监控没有数据就别分析原因全靠猜就是浪费时间。还有一个实战里验证过很多次的小技巧如果loss尖峰恰恰发生在“loss已经掉到很低”之后的某个epoch先怀疑学习率没跟着降九成的case都会落在这里。如果尖峰出现后训练又自动恢复到之前的水准可以不用太紧张但你最好把保存checkpoint的频率调高一点以防下一次尖峰直接把模型推到发散区域连恢复的机会都没有。毕竟训练一次大模型那么贵不该在最收尾的时候功亏一篑。
返回列表