PyTorch 0.4.0无溢出/下溢警告,训练出现NaN值如何解决?
嘿,我来帮你搞定这个头疼的数值问题!PyTorch 0.4.0作为比较早期的版本,默认确实不会主动触发像exp(10000)得到inf这类数值溢出的警告,这给训练中的NaN排查带来了不少麻烦。下面是几个实用的解决思路,亲测有效:
一、先把数值警告和溢出检测开起来
1. 手动添加张量检查钩子
PyTorch 0.4.0没有内置的全局数值警告开关,那我们就自己写个小工具来监控张量:
import torch def check_nan_inf(tensor, tensor_name="unknown"): has_nan = torch.isnan(tensor).any().item() has_inf = torch.isinf(tensor).any().item() if has_nan: print(f"⚠️ NaN detected in {tensor_name}!") # 这里可以加断点,方便你定位代码位置 import pdb; pdb.set_trace() if has_inf: print(f"⚠️ Inf detected in {tensor_name}!") import pdb; pdb.set_trace()
然后在关键运算后调用它,比如你的exp运算:
input_tensor = torch.Tensor([10000]) output = torch.exp(input_tensor) check_nan_inf(output, "exp_result")
这样下次再出现inf或NaN,就能立刻抓住它的踪迹。
2. 开启Python浮点警告辅助
除了PyTorch张量的检查,还可以开启Python的浮点异常警告,帮你更早发现数值异常:
import numpy as np import warnings np.seterr(all='warn') warnings.filterwarnings("default", category=RuntimeWarning)
这个能捕获一些Python层面的浮点问题,和上面的张量检查配合使用效果更好。
3. 能升级PyTorch的话尽量升
如果你的项目依赖允许,强烈建议升级到PyTorch 1.x及以上版本(比如1.5+)。后续版本在数值稳定性上做了很多优化,默认会对严重的数值问题发出警告,还能通过torch.set_warn_always(True)强制开启更多细节警告,省掉很多手动检查的麻烦。
二、一步步排查训练中的NaN来源
1. 逐层检查模型输出
在模型的forward函数里,对每一层的输出都调用check_nan_inf,比如卷积层、激活层、全连接层之后,这样就能定位到哪一层最先出现NaN/inf,缩小排查范围。
2. 检查损失和优化器配置
- 损失函数:比如交叉熵损失如果输入的logits数值过大,会导致计算时溢出变成inf,进而变成NaN;如果是自定义损失,要检查是否有除以零、开根号这类容易出问题的操作。
- 优化器:学习率过高是NaN的头号元凶!先试试把学习率降到原来的1/10甚至1/100,看看问题是否消失。另外,Adam这类优化器在初始阶段如果梯度异常大,也可能导致参数更新出问题,可以试试先用SGD热身几轮。
3. 排查输入数据
别忽略数据本身的问题!检查训练数据里有没有NaN、无穷大或者极端值(比如某个特征突然出现1e10的数值),这些脏数据会在模型运算中被放大,最终导致NaN。可以先对输入做归一化/标准化,把数值限制在合理范围内。
三、给容易溢出的函数做“安全包装”
像torch.exp()这种天生容易溢出的函数,我们可以给它加个数值截断的安全层:
def safe_exp(x): # float32的最大值约为1e38,exp(88.72)≈1e38,所以把输入截断到这个值以下 x = torch.clamp(x, max=88.72) return torch.exp(x)
这样就不会直接得到inf,后续运算也不容易变成NaN了。
内容的提问来源于stack exchange,提问作者ElleryL

