导读:本期聚焦于小伙伴创作的《训练时出现NaN怎么办?如何排查溢出与除零错误》,敬请观看详情。反向传播中途loss突然变成NaN,模型权重瞬间失效,是深度学习训练里最令人头疼的故障之一。多数情况并非随机bug,而是数值计算越界所致。当激活值或梯度超过浮点表示上限,便会溢出为inf,再经运算退化为NaN;而分母为零的归一化操作会直接产出未定义结果。本文从数值稳定性原理切入,说明如何用断言与钩子函数定位首个异常张量,并对比梯度裁剪、混合精度与epsilon平滑三类修复方案的实际表现,帮助你在训练崩塌前及时拦截风险。

在深度学习模型训练过程中,数值不稳定是导致训练失败的主要原因之一。当我们计算损失函数、梯度或者归一化统计量时,如果某些中间结果超出了浮点数可表示的范围,或者参与了不合法的运算,就会得到NaN(Not a Number)。NaN具有传播性,一旦某个张量中出现NaN,后续所有依赖该张量的计算都会变成NaN,最终表现为loss变为NaN、准确率停滞或模型参数全部失效。理解NaN产生的底层机制,是快速定位并解决训练问题的第一步。

NaN产生的根本原因:溢出与除零

从计算机浮点数标准(IEEE 754)来看,单精度(float32)和半精度(float16)的数值范围都是有限的。以float32为例,其最大正规数约为3.4e38,超过该值就会发生上溢(overflow),结果变为inf。当inf参与减法(如inf减inf)或除以inf时,就会产生NaN。在神经网络中,过大的学习率、缺乏归一化的网络层、或者错误的损失函数实现,都可能让激活值或梯度在几步之内膨胀到溢出。

除零则是另一类直接引发NaN的操作。在批量归一化(BatchNorm)或层归一化(LayerNorm)中,我们计算方差后开根号,再用特征值减去均值再除以标准差。如果方差为零(例如某层输入全部相同),且代码中没有添加微小的平滑项,那么除以零就会得到NaN。类似的还有注意力机制中的softmax分母、对比学习中的相似度归一化等场景。

下面这段伪代码展示了一个典型的除零风险点:当输入张量所有元素相等时,std为零,直接除法导致NaN。

import torch

def unsafe_norm(x):
    mean = x.mean(dim=-1, keepdim=True)
    std = x.std(dim=-1, keepdim=True)
    # 如果std为0,下面这行会产生NaN
    return (x - mean) / std

x = torch.ones(4, 8)  # 所有元素都是1,std=0
print(unsafe_norm(x))  # 包含NaN

使用钩子与断言定位首个异常张量

当训练在几百步之后才出现NaN,盲目调参效率极低。PyTorch提供的钩子(hook)机制允许我们在前向或反向传播时检查每一个模块的输出。通过注册前向钩子,我们可以打印或记录每一层输出的统计量(如最大值、是否含有NaN),从而精准找到第一个产生NaN的层。

除了钩子,在训练循环中加入断言(assert)也是一种轻量做法。例如在损失计算后立刻检查torch.isfinite(loss),若不为真则抛出错误并保存当前batch数据。配合梯度裁剪,我们可以在反向传播前拦截梯度中的inf。以下代码演示如何用钩子捕获NaN:

import torch
import torch.nn as nn

def nan_hook(module, inp, output):
    if isinstance(output, torch.Tensor):
        if not torch.isfinite(output).all():
            print(f'NaN detected in {module.__class__.__name__}')
            raise RuntimeError('NaN in output')

model = nn.Sequential(nn.Linear(10, 20), nn.ReLU(), nn.Linear(20, 2))
for name, mod in model.named_modules():
    mod.register_forward_hook(nan_hook)

x = torch.randn(3, 10)
try:
    model(x)
except RuntimeError as e:
    print(e)

在实践中,建议将钩子仅用于调试阶段,因为频繁的张量检查会带来约5%到15%的性能开销。定位到问题层后,应针对性修改该层实现或输入处理逻辑,而非全局开启检查。

三类修复方案的原理与对比

针对溢出问题,梯度裁剪(gradient clipping)是最常用的手段。它通过限制梯度范数上限,防止参数更新步长过大引发激活爆炸。在PyTorch中可使用torch.nn.utils.clip_grad_norm_实现。该方法不改变前向数值分布,仅约束反向传播,因此适用性广,但无法解决前向除零问题。

混合精度训练(mixed precision)使用float16加速计算,但半精度动态范围极小,更容易溢出。此时需配合损失缩放(loss scaling):将损失乘以一个大常数,使梯度移入半精度可表示区间,反向后再缩放回来。若未正确设置缩放因子,反而会更频繁地出现NaN。相比之下,在归一化操作中添加epsilon(如1e-5)是消除除零最直接的方式,几乎所有官方实现都默认带此参数。

我们通过一个对比表格总结三种方案的特点:

方案解决的主要问题额外开销使用建议
梯度裁剪梯度溢出默认开启,范数设1到10
损失缩放半精度下溢与溢出用AMP自动管理
epsilon平滑除零所有除法分母必加

综合来看,防御NaN应当从前向数值安全(加epsilon、使用稳定激活函数如gelu替代不成熟自定义函数)和反向稳定(梯度裁剪、AMP)两端同时入手。每次修改训练脚本后,先用小数据跑五十步并开启NaN钩子验证,可大幅降低后期调试成本。

构建可复用的训练健康检查流程

为避免重复编写调试代码,可以将NaN检测封装为训练器的一部分。例如在每一个epoch开始前重置计数器,在backward之后扫描梯度,若发现非有限值则保存当前状态并退出。这种流程化设计让团队成员在接手项目时能立刻获得清晰的错误上下文,而不是面对一条毫无信息的NaN loss日志。

另一个常被忽略的点是数据预处理。输入特征中若含有缺失值(NaN)或极端离群点,会直接穿透到网络内部。因此在数据加载器中加入torch.nan_to_num或基于分位数的截断,能从源头消灭一类隐患。结合前文提到的钩子与裁剪策略,便可形成从数据、前向、反向到参数更新的完整防护链。

当模型结构复杂、包含自定义CUDA算子时,还需确认算子内部是否正确处理了边界情况。有些第三方算子未做分母保护,调用后即污染整个计算图。此时应在算子外包裹一层Python端的数值校验,或向维护者提交补丁。只有将检查意识融入工程习惯,才能彻底摆脱训练NaN的反复折磨。

NaNoverflowdivide_by_zero修改时间:2026-08-13 12:42:35

免责声明:​ 已尽一切努力确保本网站所含信息的准确性。网站内容多为原创整理与精心编撰,观点力求客观中立。本站旨在免费分享,内容仅供个人学习、研究或参考使用。若引用了第三方作品,版权归原作者所有。如内容涉及您的权益,请联系我们处理。
内容垂直聚焦
专注技术核心技术栏目,确保每篇文章深度聚焦于实用技能。从代码技巧到架构设计,为用户提供无干扰的纯技术知识沉淀,精准满足专业提升需求。
知识结构清晰
覆盖从开发到部署的全链路。AI、前端、编程、数据库、服务器、建站、系统层层递进,构建清晰学习路径,帮助用户系统化掌握开发与运维所需的核心技术。
深度技术解析
拒绝泛泛而谈,深入技术细节与实践难点。无论是数据库优化还是服务器配置,均结合真实场景与代码示例进行剖析,致力于提供可直接应用于工作的解决方案。
专业领域覆盖
精准对应开发生命周期。从前端界面到后端编程,从数据库操作到服务器运维,形成完整闭环,一站式满足全栈工程师和运维人员的技术需求。
即学即用高效
内容强调实操性,步骤清晰、代码完整。用户可根据教程直接复现和应用于自身项目,显著缩短从学习到实践的距离,快速解决开发中的具体问题。
持续更新保障
专注既定技术方向进行长期、稳定的内容输出。确保各栏目技术文章持续更新迭代,紧跟主流技术发展趋势,为用户提供经久不衰的学习价值。