升级PyTorch1.9后运行fashion-compatibility代码出现numpy采样报错如何修复
PyTorch 1.9版本升级后旧代码报错快速修复方案
问题场景
CentOS 7 GPU环境下将3年前的旧版PyTorch升级到稳定版1.9后,未修改原论文代码直接运行出现两类异常:
- 运行启动阶段首先触发
transforms.Scale弃用警告 - DataLoader加载数据阶段抛出
ValueError,错误栈指向polyvore_outfits.py的sample_negative方法中choice = np.random.choice(candidate_sets)代码行,底层报错信息为执行numpy.random.mtrand.RandomState.choice时,传入的dict_keys对象无法被识别为整数,提示输入a必须是一维数组或整数
快速修复方案
1. 修复transforms.Scale弃用警告
- 触发原因:PyTorch 0.11及以上版本将
transforms.Scale接口重命名为transforms.Resize,旧接口被标记为弃用 - 操作方法:全局搜索代码中所有
transforms.Scale调用,直接替换为transforms.Resize即可,两个接口参数规则完全兼容,无需额外调整参数
2. 修复np.random.choice参数类型报错
- 触发原因:Python 3.6及以上版本中
dict.keys()返回的是dict_keys视图对象而非列表,较高版本的numpy不支持直接将dict_keys作为参数传入np.random.choice,升级PyTorch时连带升级的numpy版本触发了该兼容性问题 - 操作方法:找到报错行
choice = np.random.choice(candidate_sets),修改为choice = np.random.choice(list(candidate_sets)),手动将dict_keys转为列表即可解决,无需修改其他业务逻辑
内容的提问来源于stack exchange,提问作者Mona Jalal
相关产品推荐
相关产品推荐

