调用Pytorch自定义summary方法提示输入与权重张量类型不匹配如何解决
报错原因
你遇到的报错是因为设备不匹配:summary函数硬编码将输入张量转换为CUDA浮点类型torch.cuda.FloatTensor,但你加载的resnext50_32x4d模型默认存放于CPU内存,前向传播时输入和模型参数的设备、类型不匹配就触发了该错误。
解决方案
你可以任选以下一种方法解决:
- 方案1:将模型迁移到CUDA设备(无需修改summary代码)
调用summary函数前,添加一行代码把模型移到CUDA上即可:resnext50_32x4d = resnext50_32x4d.cuda() - 方案2:修改summary代码,自动适配模型所在设备(更通用)
找到summary函数中硬编码dtype的行:
替换为以下代码,自动读取模型参数所在的设备,匹配对应数据类型:dtype = th.cuda.FloatTensor
修改后不管模型放在CPU还是CUDA上都可以正常运行,无需额外调整。# 自动获取模型第一个参数所在的设备 device = next(model.parameters()).device dtype = th.cuda.FloatTensor if device.type == 'cuda' else th.FloatTensor - 方案3:仅CPU运行场景(无CUDA环境)
直接把summary函数中的dtype定义改为CPU浮点类型即可:dtype = th.FloatTensor
补充说明
这段代码是较早版本的PyTorch实现,现在新版本PyTorch已经不需要使用Variable封装张量,可以直接给张量指定device参数,写法更简洁,不过现有代码不需要调整也能正常运行。
内容的提问来源于stack exchange,提问作者user11717481
相关产品推荐
相关产品推荐

