You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何检查PyTorch模型训练过程中是否存在NaN梯度

PyTorch 检查模型梯度NaN值的实现方案

在loss.backward()执行完成、优化器更新参数(opt.step())之前,遍历模型所有带梯度的参数,检查梯度张量是否包含NaN值即可实现需求,有两种常用实现方式:

方式1:手动遍历参数检查(推荐训练时常驻使用)

该方式性能开销极低,可以灵活自定义触发逻辑(打印提示、中断训练、定位异常层等),修改后的训练循环代码如下:

features = torch.rand((3,1))
for i in range(10):
    opt.zero_grad()
    out = net(features)
    loss = torch.mean(torch.square(torch.tensor(5) - torch.sum(out)))
    loss.backward()

    # 梯度NaN检查
    grad_has_nan = False
    for p in net.parameters():
        # 跳过被冻结、无梯度的参数
        if p.grad is not None and torch.isnan(p.grad).any():
            grad_has_nan = True
            # 如果需要定位异常层,可以在这里打印p对应的层名、参数形状
            break
    if grad_has_nan:
        print("NAN")
        # 需要中断训练直接放开下面的break即可
        # break

    opt.step()

方式2:PyTorch内置异常检测(适合调试阶段定位问题)

PyTorch提供了自动梯度异常检测接口,开启后只要反向传播过程中产生NaN梯度,会直接抛出异常并定位产生异常的前向计算算子位置,缺点是会明显拖慢训练速度,不建议正常训练时开启:

# 训练开始前开启异常检测
torch.autograd.set_detect_anomaly(True)

features = torch.rand((3,1))
for i in range(10):
    opt.zero_grad()
    out = net(features)
    loss = torch.mean(torch.square(torch.tensor(5) - torch.sum(out)))
    # 反向传播时如果出现NaN梯度,会直接触发RuntimeError中断程序
    loss.backward()
    opt.step()

补充说明

  • 手动遍历检查的性能开销极小,适合训练全程常驻使用;如果需要精准定位是哪一层出现梯度NaN,可以改用net.named_parameters()遍历,检测到异常时直接打印参数名即可快速定位问题
  • 内置异常检测会显著拖慢训练速度,仅建议定位NaN问题根源时临时开启
  • 除了梯度NaN,训练过程中出现的参数值NaN、激活值NaN问题,都可以通过torch.isnan(目标张量).any()的逻辑在对应计算节点检查

内容的提问来源于stack exchange,提问作者Penguin

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.31 00:42:27