如何检查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
相关产品推荐
相关产品推荐

