运行PyTorch代码出现NumPy模块重复加载UserWarning的原因是什么
警告触发核心原因
该警告的本质是NumPy模块在当前运行环境中被先后导入了两次,和你编写的业务代码没有直接关联,属于运行环境和依赖版本的已知问题,具体触发逻辑如下:
- 你使用的PyTorch 1.8.1+cu111属于较旧的版本,该版本的CUDA自动混合精度(AMP)模块存在已知缺陷,模块初始化时会触发重复导入NumPy的逻辑
- 代码中导入
torchvision.transforms以及调用ToTensor()方法时,内部会先调用一次NumPy完成PIL图像到张量的转换逻辑,后续PyTorch底层初始化CUDA相关组件时,又触发了第二次NumPy导入,进而触发警告 - 部分Anaconda环境中预装的NumPy版本和PyTorch编译时绑定的NumPy版本存在轻微不匹配,也会触发重复加载检测逻辑
对运行的影响
绝大多数场景下该警告不会影响代码正常执行,你当前的代码可正常运行就是典型情况。警告中提到的潜在问题仅在极端复杂的多进程、多模块动态导入场景下才有可能出现,表现为数值计算一致性异常、微小的内存泄漏等,常规的模型训练、数据处理场景基本不会碰到。
可选解决方案
如果不想看到该警告,可任选以下方式处理:
- 直接屏蔽特定警告,在代码开头添加如下代码即可:
import warnings warnings.filterwarnings("ignore", category=UserWarning, message="The NumPy module was reloaded") - 升级PyTorch版本到1.9.0及以上,该重复导入的缺陷已经在后续版本中被官方修复
- 卸载当前环境的NumPy后重新安装匹配PyTorch 1.8.1的NumPy版本,消除版本差异导致的重复加载问题
内容的提问来源于stack exchange,提问作者liwenqiang
相关产品推荐
相关产品推荐

