
如果你在深度学习中遇到过 Loss 突然变成 NaN 的情况应该能理解那种困惑——明明昨天训练得好好的今天改了一行代码或者换了一条数据loss 直接从 2.3 飞到 NaN整个终端刷满了红色告警。更头疼的是NaN 的结果往往是链式传播的loss 变成 NaN反向传播的梯度变成 NaN所有参数跟着变成 NaN最后保存出来的模型权重全是一堆 nan等于白训。这个问题的糟心程度我在各类项目里深有体会。无论是自己从零手写 DETR、SegFormer还是用 MMDetection、MMRotate、PaddleOCR、EasyOCR、nnU-Net 这些开源工具箱甚至是跑 LoRA 微调大模型NaN 都像一个挥之不去的幽灵。而且 自定义 Loss 函数 这个场景尤其容易中招框架自带的 CrossEntropyLoss、DiceLoss 通常做了大量数值稳定性处理你自己写的那几十行 loss 代码却没做任何保护一旦计算路径里出现 log(0)、除零、inf 相减NaN 就顺理成章地出现了。我写这篇内容就是想把自己这些年排查 NaN 的经验完整梳理一遍。核心聚焦在自定义 Loss 函数的数值稳定性同时把模型侧、数据侧、训练策略侧的常见诱因一并讲清楚。不管你是刚入坑深度学习、正在用 YOLOv8 训练自己的数据集的小白还是已经在用 Mask2Former、mmrotate 跑 DOTA 数据、甚至在做 LoRA 微调和增量训练实战的进阶玩家这篇文章都能给你一套可以直接抄作业的排查方法和防御习惯。1. NaN 不是“灵异事件”先从数值计算的角度看清它出现的原因很多人一看到 loss 变成 NaN第一反应是模型结构有问题、代码写错了、环境坏了。但我要说一个反直觉的结论NaN 很少是随机冒出来的它是一系列数值事件在浮点精度约束下累积的结果。在 IEEE 754 浮点数标准里NaN 是“Not a Number”的缩写它表示一个未定义或不可表示的数值。正常情况下只有少数几种运算会产生 NaN。1.1 产生 NaN 的几种典型运算路径我总结了一下在深度学习训练里NaN 的来源基本逃不出下面这几类0 / 0分母为零时除零操作。在 loss 函数里最常见的就是loss torch.mean(a / b)如果 b 在某个 batch 里恰好为 0整个损失就是一个除零错误。inf - inf正无穷减正无穷。比如log(0) -inf如果 sigmoid 输出饱和到 0torch.log(pred)就是 -inf如果再叠加一个 -inf减去它就等于加上了 inf于是得到 NaN。0 * inf零乘以无穷。在 mask 操作里经常发生某个位置被 mask 成了 0但它的 loss 分支仍然计算出了 inf乘上 0 变成 NaN。sqrt(负数)在自定义 loss 里如果用了平方根比如欧氏距离的变体而内部数值因为浮点误差变成了一个极小的负数sqrt 就会输出 NaN。超过浮点范围float32 的最大值约是 3.4e38float16 的最大值只有 65504。当梯度或者中间变量在指数级增长时很容易直接溢出成 inf后续运算再一步就变成 NaN。拿生活类比的话这就像你在计算器上按了0 除以 0计算器不会给你报错只会给你显示一个“数学错误”的符号。深度学习框架也一样它不会在 NaN 出现的那一瞬间停下来告诉你“这里出错了”而是默默地把这个 NaN 传递下去直到你在终端里看到 loss 变成 nan才意识到出了问题。1.2 理解 forward 和 backward 两个阶段的 NaN这里有个关键点NaN 可能出现在前向传播阶段也可能只出现在反向传播阶段。前向阶段的 NaN 最容易被观察到因为你直接把 loss 打印出来就是 nan。反向阶段的 NaN 比较隐蔽——你打印 loss 的时候还是 1.23但 backward 之后发现梯度已经是 NaNloss 下一轮迭代就直接变成 NaN 了。这种“前向正常、梯度 NaN”的情况往往是因为 loss 函数对某个中间量的导数在极端输入下不存在或者溢出。最经典的例子是手写 softmax log前向传播时你算出了正确的概率分布log 之后也没有问题但在反向传播时梯度公式里出现了除以 softmax 分母的项当某个 logit 特别大导致 softmax 输出几乎全为 0 时梯度就会除零变成 NaN。所以排查 NaN 时不能只看 loss 本身还要看梯度的状态。这一点我在第 3 节会展开讲完整的定位链路。1.3 一个重要区分一开始就是 NaN 还是训练中途才变 NaN遇到 NaN我第一个会问自己它是从第一步迭代就出现还是训练了几百步之后突然出现这两种情况的排查方向完全不同。如果从第一步就是 NaN大概率是数据输入或初始化的问题比如输入数据里本身带有 NaN、标签类别超出了 num_classes 范围、归一化时除的是零方差、或者自定义 loss 的分母在一开始的随机预测下本来就是 0。这种问题通常在运行第一个 batch 时就会爆发。如果训练中途才突然变 NaN那多半是数值累积和训练策略的问题学习率过大导致梯度爆炸、混合精度下 loss scale 处理不当、模型某些层参数在长期训练后退化到极端值比如 BN 的 running variance 变成 0、或者是数据流中偶然出现了异常样本。这种问题更隐蔽因为你可能需要复现很多次才能抓到那个触发条件。2. 自定义 Loss 函数的高危写法我反复踩过的那些坑既然题目聚焦在“定义 Loss 函数”这一节我必须把那些容易引入 NaN 的写法逐个点名。很多坑不是运气不好踩到的而是代码范式本身就有问题只是数据或者训练状态一旦波动问题就暴露了。2.1 对可能为 0 的分布取 log交叉熵实现里的经典错误手写交叉熵损失是 NaN 的头号来源。新手版错误写法是这样的# 极其危险的写法 loss -torch.mean(target * torch.log(pred))这行代码有三个问题。第一pred如果经过 softmax它的值域是 (0, 1)但可能非常接近 0。比如模型对某个类别很自信时softmax 输出可能只有 1e-8这时torch.log(pred)就是 -18 左右虽然还没到灾难程度但如果 pred 因为浮点下溢直接变成 0log(0) -inf整个 batch 的 loss 就变成 inf反向传播后参数全部 NaN。第二target如果是 one-hot 编码绝大多数位置是 00 * (-inf)这个操作在浮点运算里直接得到 NaN。第三正常实现的 cross-entropy 应该对 logits 做 log-softmax而不是对 softmax 的输出做 log。我见过不少人在实现 Focal Loss、Dice Loss 变体时重蹈这个覆辙。比如这样# Focal Loss 手写pt 下溢时 CE 部分还是会炸 ce -torch.log(pt) # pt 可能为 0 focal alpha * (1 - pt) ** gamma * ce正确做法要么直接用 PyTorch 内置的F.cross_entropy要么至少对 logits 走F.log_softmax并且用clamp保护对数输入。内置 API 之所以稳定是因为它在数学上做了等价变换把log(softmax(x))合并成了x - logsumexp(x)logsumexp 在极大值时会先减掉最大值再算指数从而避免了溢出。2.2 除以一个可能为 0 的分母Mask 和类别缺失问题在分割、检测、对比学习场景里我们经常需要做 mask 归一化。比如只对前景像素计算 loss然后除以前景像素数# 危险写法 loss (pred - target).pow(2).sum() / num_fg_pixels如果某个样本完全没有前景目标——这在训练前期非常常见尤其是你刚换了一个新数据集、标注还没完善或者 DOTA 这类旋转目标数据集里某些类别本身就稀有——那么这个num_fg_pixels恰好为 0整个 loss 就是0 / 0 NaN。我在用 mmrotate 训练 DOTA 数据集时就遇到过这个问题。旋转框的回归分支里经常用delta / (w * h)这种方式做尺度归一化如果某个 gt 框的宽或高因为标注错误变成 0分母为 0回归 loss 直接 NaN。而且这种错误往往不是全量数据的是某几张图里的几个框触发的所以特别难复现。防御型的写法是做一个安全除法denom num_fg_pixels.clamp(min1) loss (pred - target).pow(2).sum() / denom或者用torch.where(num_fg_pixels 0, loss_value, torch.zeros_like(loss_value))把无效样本的 loss 置为 0。但要注意torch.where两个分支都会计算如果loss_value在分母为 0 时计算出 NaNtorch.where不会救你——因为 NaN 已经被计算出来了只是被掩盖在未选择的分支里梯度回来后它还是会作妖。2.3 对负数开平方根或做非整数次方距离度量的隐藏雷区自定义 loss 里常写torch.sqrt(distance)或者torch.pow(distance, 0.5)用于拉近特征向量的距离或做边界惩罚。理论上 distance 应该非负但由于浮点运算的舍入误差某些操作会让一个本来应该为 0 的距离变成一个很小的负数比如 -1e-9。不要小看这个 -1e-9。sqrt(-1e-9)在 PyTorch 里不会报错但会输出 NaN然后你的 loss 就神秘地变成了 nan。TensorFlow 也一样tf.math.sqrt对负数是 NaN还好tf.sqrt内部可能处理成 NaN 而不是报错但结果同样是灾难。防御写法是先clamp_min(0)再开根号distance torch.nn.functional.relu(distance) # 先非负 loss torch.sqrt(distance)或者干脆避免开根号直接用平方距离既省了计算量又避开了 NaN 风险。很多实现里 MSE 比 RMSE 更常用除了数学性质更好数值稳定性也是重要原因。2.4 忽略 padding 位置标签里塞了无效值NLP、OCR、序列标注场景里padding 是常规操作。如果你自定义了一个序列 loss 函数但忘记传入 padding mask那么在 padding 位置上模型会收到无效的 label 和毫无意义的梯度。比如 EasyOCR 训练自己模型时标签序列是变长的padding 部分通常填 0。如果 loss 函数里没设ignore_index0模型即使预测对了 padding 位置的输出也会被当成错误去惩罚导致 loss 大得离谱。更极端的情况是label 里某个索引因为字典构造错误超出类别数模型的输出概率在那个位置趋近于 0log 之后得到 -infloss 变成 NaN。PaddleOCR 的训练里也有类似的坑文本识别分支用 CTC Loss如果字典顺序不统一标签里的字符索引和模型的 class 数量不匹配或者ignore_index设置错误CTCLoss 的 forward 阶段就可能因为 logits 存在 -inf 而输出 NaN。正确的做法是在构造 dataloader 时同步生成valid_mask在 loss 计算里用 mask 过滤同时给 CTC 或 CE 损失传ignore_index参数。2.5 混合精度下的缩放问题AMP 和 NaN 的暧昧关系很多人写了正常的 loss 代码一开自动混合精度AMP就 NaN关了就好。这不一定是你代码错了而是混合精度训练时 loss 被梯度缩放器GradScaler缩放缩放因子在溢出时会跳过更新但自定义 loss 里若有较大的中间值计算图里就可能生成 NaN。float16 的动态范围只有 float32 的约千分之一。如果你的 loss 函数里出现了大数值中间量比如计算 logits 和 label 的 inner product 后再除以一个很小的温度参数得到的中间值很容易超过 65504 的上限变成 inf。AMP 的 GradScaler 虽然会动态调节 loss 的缩放因子来避免梯度下溢但它管不了你 loss 函数内部的中间数值溢出。所以我现在的习惯是自定义 loss 函数里的中间张量尽量保持 float32只在必要的时候才让它参与 AMP 缩放。具体做法是在 loss 计算前手动把关键张量.float()化最后再把结论转回输入 dtype。另外PyTorch 的torch.autocast默认对大多数算子用 float16 计算如果你的 loss 里有不受 autocast 保护的敏感操作比如torch.log、torch.exp最好在autocast语境里显式包一层torch.cuda.amp.custom_fwd(cast_inputstorch.float32)。3. 一次完整的 NaN 定位链路从打印到钩子的排查过程当 NaN 真的发生的时候很多人第一反应是打开 loss 代码反复盯但盯代码往往看不出所以然。我建议你按下面这条链路来排查先确定故障范围再逐步缩窄到具体张量最后找到触发那一步。3.1 第一步区分问题的空间位置在开始复杂排查之前先确定 NaN 出现在哪个阶段。这一步只花两分钟但能砍掉一半的排查方向。如果打印出来的 loss 第一轮就是 NaN优先检查输入数据有没有 NaN、标签是否非法、初始化和 loss 函数本身的边界条件。如果 loss 是逐步增长然后突然跳到 NaN优先怀疑学习率过大、梯度爆炸、混合精度溢出、或数据流中的偶发异常样本。如果 loss 打印正常但模型参数很快变成 NaN需要检查梯度因为问题在 backward 阶段甚至有可能出现在优化器更新后。定位手段很简单在训练循环里插入一段检查for batch_idx, batch in enumerate(train_loader): preds model(batch) loss criterion(preds, batch[target]) if not torch.isfinite(loss): print(fLoss is not finite at batch {batch_idx}: {loss}) torch.save(batch, bad_batch.pt) # 保存触发异常的 batch break很多框架里尤其是检测、分割模型一个数据 batch 里有大量“合法但边界”的样本保存 bad_batch 可以让你反复用同一份数据去复现非常利于后续排查。3.2 第二步利用 PyTorch 的异常检测机制PyTorch 提供了一个能定位到具体张量运算的调试工具torch.autograd.set_detect_anomaly(True)。torch.autograd.set_detect_anomaly(True) # 要在 forward 之前设置开启之后一旦反向传播图里出现 NaN 或 InfPyTorch 会直接抛出异常并且告诉你出错的操作位置例如RuntimeError: Function LogBackward returned nan values in its 0th output.它还能定位到是哪个模块的哪个计算导致了 NaN。缺点是它会显著拖慢训练速度而且有时会报错在CopyBackwards之类的底层操作上需要你再往上追一层。但对那种“loss 正常、梯度 NaN”的隐蔽问题detect_anomaly是目前最有效的定位工具。不过要注意detect_anomaly本质上是检查 backward 时输入的张量里是否有 NaN/Inf它不能告诉你“NaN 究竟是在哪一行 forward 代码里产生的”。如果 forward 阶段已经算出 NaN你打印 loss 时就看到了。所以它主要帮你锁定的是 backward 链路。3.3 第三步用钩子检查每个中间激活和参数梯度当detect_anomaly报错位置不够明确时比如报在一个通用算子我会在模型的每个子模块上挂 forward/backward hook打印哪些层的输入输出包含非有限值。这个方法虽然土但最直接。具体做法def register_nan_hook(module, name): def forward_hook(module, input, output): if isinstance(output, torch.Tensor) and not torch.isfinite(output).all(): print(fNaN in forward output of {name}) elif isinstance(output, (list, tuple)): for i, o in enumerate(output): if isinstance(o, torch.Tensor) and not torch.isfinite(o).all(): print(fNaN in forward output of {name}[{i}]) def backward_hook(module, grad_input, grad_output): for i, g in enumerate(grad_output): if isinstance(g, torch.Tensor) and not torch.isfinite(g).all(): print(fNaN in backward grad_output of {name}[{i}]) module.register_forward_hook(forward_hook) module.register_full_backward_hook(backward_hook) for name, module in model.named_modules(): register_nan_hook(module, name)运行几步后终端会打出类似NaN in forward output of backbone.layer3.2.conv2的信息你就知道故障来自哪个子网络。如果 forward hook 全都没输出但 backward hook 有输出说明问题出在 loss 对某个中间量的梯度计算上这时候就要回 loss 函数里找。3.4 第四步检查梯度范数与参数状态定位到层之后还要确认是“梯度爆炸”还是“梯度本身变成了 NaN”。我通常在训练循环里加一段梯度检查total_norm 0.0 for p in model.parameters(): if p.grad is not None: param_norm p.grad.detach().data.norm(2) total_norm param_norm.item() ** 2 total_norm total_norm ** 0.5 print(fgrad norm: {total_norm}) if not torch.isfinite(total_norm): print(Gradient norm is NaN/Inf, stop training and inspect)如果total_norm在变成 NaN 之前已经冲上了 1e10 甚至 1e20那基本可以断定是梯度爆炸修复方向是降低学习率、加梯度裁剪、加权重初始化调整。如果total_norm是直接一步跳到 NaN没有经历过极端大值那更可能是在 loss 计算里某处产生了 NaN通过链式法则把 NaN 传给了所有梯度。3.5 平时做的数据探针训练前先跑 10 step其实我最推荐的“定位”是防患于未然的做法每次定义完自定义 loss 后先别急着全量训练跑一个 10-step 的 smoke test。用固定的随机种子相同的输入张量打印 10 步的 loss 和 grad norm。如果 10 步内没有 NaN再把训练跑起来。这个习惯帮我省了无数时间。因为 NaN 很多情况下是“数据触发”的——某个 batch 恰好出现了极值或非法标注才把 loss 里的隐患引爆。smoke test 虽然不能完全排除这种偶发性但至少能保证模型结构、loss 的常规路径是健康的。换数据集、改 loss、改预处理之后都值得先跑这 10 步。4. 修复的关键不在 Loss 本身学习率、数据和训练策略同样致命排查完到底是不是 loss 代码的“锅”之后你会发现真正导致 NaN 的元凶经常不在 loss 函数内部。它可能在优化器的学习率上在数据预处理里甚至在 BN 层的统计量中。这一节我逐个讲清楚并给出能直接用的修复手段。4.1 学习率过大最容易被忽视的头号元凶我见过特别多这样的情况loss 函数本身写得没问题但学习率从 1e-4 改到 3e-4 之后训练到第 2000 步 loss 突然变成 NaN。原因很简单学习率过大导致梯度更新幅度过大参数冲到了某个极端区域再往前推进一步输出就溢出了。特别是模型里有 attention 模块或 BN 层时参数对学习率的敏感度极高。解决办法有几种。先降低学习率比如从 3e-4 退到 1e-4这是最直接的。如果还想用较大的学习率那就加 warmup预热前 1000 步把学习率从 0 线性增长到目标值让模型在最早期的高波动阶段走得更稳。我在训练 YOLOv8 和 SegFormer 时warmup 几乎是标配不只是为了收敛精度更是为了避 NaN。另外如果是自定义 loss 里某些项天然数值较大比如检测回归里的坐标差动辄几百可以在 loss 函数里对这个 term 乘以一个小的权重系数比如 0.1让梯度对参数的冲击变小。很多检测模型里loss_cls和loss_bbox的权重就是这么配的不一定是为了调精度也可能是为了稳训练。4.2 优化器本身要设置好Adam 的 eps 不是摆设很多人在用 AdamW 时会把eps默认成1e-8。这个小数值在 float32 下勉强能用但如果你开了 AMPfloat16 梯度可能被缩放后下溢到接近 0分母的sqrt(v) eps里 v 是历史梯度的平方均值可能在某个参数位上小到 1e-12这时候eps1e-8是可以兜住底部的。但如果你用 SGD 或 Momentum 这类优化器本身没有 eps 兜底梯度一旦爆炸就全 NaN所以 SGD 训练时梯度裁剪几乎是必须的。我强烈建议自定义 loss 训练时至少在优化器上加一个梯度裁剪。PyTorch 里一句话torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm1.0)max_norm怎么选我通常先看 grad norm 的正常范围。正常训练时 grad norm 一般在 1 到 50 之间如果某一步突然冲到 1e8裁剪到 1.0 能阻止参数瞬间飞出。但要注意裁剪只是止损它不解决根本问题。如果频繁触发裁剪还是要回去调学习率或检查 loss 设计。4.3 数据输入里藏着的 NaN预处理的一个小疏漏有一次我在训练分割模型时从第一步 loss 就是 NaN。我盯着 loss 函数看了半天又用detect_anomaly定位最后发现是输入数据里已有 NaN。原因在数据预处理某些遥感影像的波段里存在 0 值我在做归一化时用(x - mean) / std归一化而某些像素的归一化结果是 NaN——因为某块区域全是无效值分割掩膜里对应的 label 恰好没被过滤掉。这类问题用torch.isnan(batch_data).any()检查就能发现。我给的数据检查代码是这样的assert not torch.isnan(batch[image]).any(), Image contains NaN assert not torch.isnan(batch[target]).any(), Target contains NaN assert batch[target].max() num_classes, Label out of range尤其是用 mmsegmentation 训练 Cityscapes 这类数据集时label 里有一个 255 的 ignore 索引。如果你没在CrossEntropyLoss(ignore_index255)里设置 ignore_index255 作为类别索引会查到一个不存在的类概率模型输出在该位置趋近于 0log 之后就是大数甚至 -infloss 直接飞掉。这类问题通常表现为 loss 从第一轮就很高或者前几步后炸掉。4.4 BN 的计算陷阱当方差变成 0 或负数BatchNorm 在训练模式下会统计当前 batch 的均值和方差然后对数据做标准化。这里有一个冷门的 NaN 源头如果某个 batch 里某个通道的方差恰好为 0BN 的标准化公式(x - mean) / sqrt(var eps)虽然会因为 eps 兜底而不至于除零但当 var 下溢成负数理论上 var 非负但浮点下溢后可能表现为 0加上 eps 后依然极小输出会被放大到极大继而引发后续溢出的 NaN。这种情况在 batch size 很小、或者输入特征退化时容易出现。所以如果你发现开了 BatchNorm 后训练不稳定有两条路一是把 batch size 调大一点让每个通道的统计量更稳二是把eps参数从默认的 1e-5 调到 1e-4 或更大牺牲一点点精度换取稳定性。注意PyTorch 的 BatchNorm2d 默认eps1e-5我一般会习惯性改成1e-4尤其是用自定义 loss 训练检测或分割模型时。4.5 fp16 混合精度训练里的经典对策如果你使用的是 NVIDIA 显卡的自动混合精度AMP有一个更隐蔽的坑AMP 会动态调整 loss 缩放因子当检测到梯度溢出时会把缩放因子减半并跳过这次 optimizer 更新。正常来说这能保护模型不被梯度溢出毁掉。但在自定义 loss 场景下问题在于如果你擅自修改了 loss 的数值尺度——比如把 loss 乘以 1000 想放大梯度——那么在 AMP 下grad scale 和你的放大因子会叠加作用极容易导致梯度溢出触发跳过。我在项目中见过有人为了“平衡多个 loss 项”给某个 loss 乘以 100 的权重然后训练变得极其不稳定。正确做法是保持各 loss 项在同一量级权重尽量在 0.1 到 10 之间。如果一个 loss 天然就比另一个大几个数量级先在 loss 内部做归一化而不是靠外部乘一个大权重去硬拉。此外AMP 的GradScaler状态要记得保存和恢复在断点续训时如果只恢复了模型权重而没恢复 scalerloss scale 会从默认值重新起步可能造成早期不稳定。5. 不同模型与框架下的 NaN 高频现场从 YOLO 到 Transformer 的对比热词里出现了 YOLOv8、Mask2Former、mmrotate、PaddleOCR、EasyOCR、nnU-Net、LoRA 等大量真实训练场景。我自己在这些工具箱里都踩过 NaN与其抽象地讲不如把它们放在一张表里做个对比每个场景的核心诱因和修复建议一目了然。训练场景常见 NaN 表现核心诱因首选修复手段YOLOv5/YOLOv8 自训练loss 早期正常几百步后突然 NaN学习率过高 数据里存在异常 box宽高为 0加 warmup、降低初始 lr、数据增强里过滤非法 boxSegFormer / mmsegmentation 训练 Cityscapes开启训练后第一个 epoch loss 就是 inf 或 NaNlabel 中 255 没被 ignore交叉熵查表越界CrossEntropyLoss(ignore_index255)或自定义 loss 里过滤 ignore 索引Mask2Former 训练推理正常但训练 loss 中 mask loss 突然 NaNmask 二分类分支里正负样本极度不均Focal Loss 中的log(pt)下溢使用内置sigmoid_focal_loss而非手写、对pt做 clampmmrotate 训练 DOTA某几个 step 后回归 loss NaN旋转框回归中分母出现 0宽高为 0或角度跨周期导致 loss 突变对分母clamp_min、角度回归用中心度加权或sin/cos编码PaddleOCR / EasyOCR 自训练loss 从高到爆或者直接 NaN字典索引超出类别数、padding 位置没被 mask检查字典映射、给 CTC/CE loss 传ignore_indexnnU-Net 训练标注Dice Loss 突然 NaN某一层 GT 全为 0前景区域为空Dice 分母为 0在 loss 里对前景体素数clamp_min(1)空目标样本的 Dice 置 0LoRA 微调大模型loss 在 100~1000 步后变成 NaN学习率偏大 float16 下中间激活溢出、或梯度裁剪缺失使用 bf16 替代 fp16、降低 LoRA rank/learning rateword2vec / 负采样训练early loss 就是 NaNsigmoid 输出饱和到 0log(0) -inf使用F.logsigmoid或F.binary_cross_entropy_with_logits避免手写 logDETR 训练自己数据整体 loss 训练中某 step NaN匈牙利匹配中 cost 矩阵全为 inf或 no-object 分支 CE 中的 logits 极端检查注意力 mask 和目标数量、对 no-object cost 设置有限值逐个解释几个细节。YOLO 系改动 anchor 或标签时最容易埋雷的是 box 的w, h变成 0。数据增强里比如 random perspective、rotation如果不小心把某些小目标的框翻转或裁剪到完全消失生成的 target box 宽高可能为负或 0而 YOLO 的 GIoU/CIoU loss 在 delta 计算时以w * h作分母直接引发 NaN。修复方法是在 collate 时过滤掉宽高小于 1 像素的目标。Mask2Former 和各类检测分割模型mask loss 通常用 sigmoid focal loss 或 BCE。手写 Focal Loss 时不保护pt 1 - p当预测概率p极度接近 1 时(1 - p)下溢为 0那log(0)又来了。这类问题哪怕你写了(1 - pt) ** gammagamma 次方本身没问题但后面的交叉熵项依然会炸。所以我建议直接复用mmdet或torchvision里现成稳定实现不要自己造 Focal Loss 的轮子。mmrotate 的 DOTA旋转框的表示有多种形式如果用角度回归的方式当 anchor 和 gt 的角度差跨过周期边界比如 180 度和 -180 度角度 loss 可能会产生一个异常大的梯度。而 DOTA 里大长宽比的舰船、油罐目标很多角度回归如果没做好周期处理训练后期容易突然 NaN。修复方案是改用gliding vertex或sin/cos编码或者在 loss 里对大角度差做截断。LoRA 微调大模型热词里有 LoRA 训练。LoRA 场景下 NaN 的原因经常是 fp16 的精度不足以支撑某些模型层在适配器更新时的中间值。很多人在 LoRA 微调时用 4-bit 量化中间的compute_dtype设成 float16容易遇到 loss 在几百步后变 NaN。我的经验是能上 bf16 就上 bf16bf16 虽然精度低一点但动态范围大得多几乎不会溢出。另外 LoRA 论文里有一句很重要的话LoRA 的初始化是 B0、A 随机高斯这保证初始时 delta W 为 0。如果你的代码在初始化时不注意直接把 A 和 B 都随机初始化模型一开始的输出就偏离原始权重loss 可能非常大配合大学习率就炸了。nnU-Net 的 Dice Loss医疗分割场景里经常遇到某些切片完全没有目标区域。Dice 的公式是2 * intersection / (union)分母是预测和真实区域之和当真实区域为 0 且预测也为 0 时0 / 0 NaN。正确做法是给分母加smooth比如smooth1e-5同时对“GT 全空”的样本直接把 Dice loss 记作 0不参与梯度回传。这张表的价值在于当你遇到 NaN 时先找到对应场景的第一可能原因就按表里的建议去改命中率非常高。我在多个项目里试过至少能解决 70% 的问题。6. 从训练第一天就开始做的防御性编程习惯说到最后我想分享几个我从无数次 NaN 事故里沉淀下来的习惯。这些习惯不需要花很多时间但能把 NaN 从“偶发的夜半惊魂”变成“训练流程里一个可预期的检查点”。6.1 自定义 Loss 的单元测试用玩具数据先算一遍每写一个新的 loss 函数我都会先构造一个玩具 Tenso r,手动算一遍理论上的输出再和代码输出对比。比如写一个 Dice Loss我就构造一个pred torch.tensor([0.9, 0.1])、target torch.tensor([1.0, 0.0])手算 Dice 应该是2 * 0.9 / (0.9 1.0 0.1 0.0) ≈ 0.9跑一遍看对不对。这个测试的价值不只是验证公式它还能让你在训练之前就暴露“分母会不会为零”“log 会不会取到 0”这类边界问题。更好的做法是把这些断言写成 pytest 用例每次修改 loss 后自动跑一遍。我在自己的项目里维护一个test_losses.py里面覆盖了空目标、全零预测、极大 logits、极小 logits 四个边界场景。这套测试帮我拦下了至少五次 NaN 事故。6.2 在 loss 出口加一道 finite 断言即使写了各种保护我还是建议在 loss 最终返回前加一句assert torch.isfinite(loss), fLoss is not finite! Components: cls{cls_loss}, reg{reg_loss}, aux{aux_loss}这句话会直接把三个子项的数值打出来。万一真的炸了你至少能立刻知道是分类分支炸的还是回归分支炸的还是辅助头炸的。这比面对一个光秃秃的 NaN 要省出半天排查时间。当然训练正常时这句话不消耗任何计算资源可以放心留着。6.3 复杂 Loss 函数里拆分子项再组合一个复杂的自定义 loss往往由多个子项组成。很多人在一个函数里写了 30 行最后返回total_loss loss_1 loss_2 loss_3。我强烈建议把每个子项单独算好、单独打印而不是只打印总 loss。loss_cls ... loss_box ... loss_iou ... loss_total loss_cls loss_box * 0.5 loss_iou * 0.3这样一旦训练中 total loss 变 NaN你的日志里能看到loss_cls0.23, loss_boxnan, loss_iou0.12马上把矛头指向 box 分支。很多时候 NaN 只藏在某个子项里其他子项是正常的但组合之后整个 loss 就 NaN 了。6.4 断点续训时检查三个状态如果训练已经跑了很久才出现 NaN你多半需要回滚到最近的正常 checkpoint。这里有三个状态必须一并恢复模型参数、优化器状态包括 Adam 的 exp_avg 和 exp_avg_sq、如果是 AMP 还要恢复 GradScaler 的 state。以前我偷懒只恢复模型参数结果从同一个 checkpoint 继续训练loss 很快就又 NaN原因是优化器的方差估计还停留在旧状态学习率预热也在错误的位置重启。在自定义 loss 的项目里这三件套尤其重要因为自定义 loss 对应的梯度分布往往不像内置 loss 那样平滑优化器状态丢失更容易导致训练波动。6.5 心态调整NaN 不等于模型废了最后说一点心态问题。很多朋友一看到 NaN 就想从头重新写模型其实不必。NaN 是训练过程中的一个“状态错误”不是模型架构的“死刑判决”。如果 loss 是中途变 NaN可以从最近一个正常的 checkpoint 恢复把学习率降一半加上梯度裁剪大概率能继续训练下去。如果是第一步就 NaN优先检查数据输入不要急着改模型结构。我见过最惨的情况是有人因为 NaN 连续重写了三遍 YOLO 的自定义 loss最后发现是数据集的某些标签文件里出现了空字符串解析后 box 坐标全是nan。数据侧的问题让 loss 背了很多锅。所以请记住排查时把“数据是否干净”放在模型和 loss 之前这一步永远不亏。最后再分享一个小技巧自定义 Loss 的调试我这些年养成的最有价值的小习惯其实特别简单每个自定义 loss 都保留一个“裸奔模式”——即返回一个额外的小字典包含所有中间子项的数值。这样训练日志里能直接看到各子项的走向。比如loss_components {cls: loss_cls.item(), box: loss_box.item(), iou: loss_iou.item()}在TensorBoard或终端里单独显示。这个习惯帮我在 DOTA 训练中一次就发现其实不是回归分支导致 NaN而是分类分支里某个稀有类别的样本在某个 batch 中完全没有分类 loss 的 mask 出了问题。如果没有子项拆分我可能会在回归 branch 上浪费一整天。另一个小技巧是如果在 PyTorch 里对 NaN 感到无助时试着把张量移动回 CPU 并在 NumPy 里手动算一遍同样公式的 forward。NumPy 的表现更接近数学直觉而且你可以一行行 print 中间值很多在框架黑盒里看起来“莫名其妙”的 NaN在 NumPy 里就是标准的log(0)或0/0。训练路上的 NaN拦住了很多人但一旦你掌握了这套排查链路和防御习惯它就只是一个普通的训练日志关键词而已。希望这篇文章能帮你在下一次遇到 NaN 时少走点弯路。