PyTorch训练报错RuntimeError:Input与weight张量类型不匹配如何解决?
PyTorch 设备不匹配报错排查与解决
RuntimeError: Input type (torch.cuda.FloatTensor) and weight type (torch.FloatTensor) should be the same
该报错的核心原因是前向传播计算时,输入张量已经被迁移到GPU,但模型的权重参数仍存储在CPU上,二者设备不匹配无法完成张量运算。即使你认为已经完成了模型和输入的迁移,也可以按照以下步骤逐一排查:
排查步骤与解决方案
- 排查点1:模型迁移操作是否生效
nn.Module的.to()操作虽然是原位操作,但如果迁移后你又为模型新增了子模块/参数,新增的内容默认会落在CPU上,需要单独迁移。推荐写法:
device = torch.device("cuda" if torch.cuda.is_available() else "cpu") # 模型定义完成后统一迁移 model = YourCustomModel() model = model.to(device) # 迁移后新增子模块必须单独执行迁移 model.extra_head = nn.Linear(256, 10) model.extra_head = model.extra_head.to(device)
- 排查点2:多卡训练场景下的DataParallel处理是否正确
如果使用torch.nn.DataParallel做多卡训练,包裹后的模型参数存储在model.module属性下,包裹操作要放在模型迁移到GPU之后,避免权重设备不匹配:
# 正确多卡处理顺序 model = model.to(device) model = nn.DataParallel(model)
加载预训练权重时也要注意适配module前缀,避免权重加载到错误的对象上。
- 排查点3:预训练权重加载时是否指定了设备映射
如果加载的预训练权重是在CPU环境下保存的,直接加载不会自动迁移到GPU,需要通过map_location参数指定加载到目标设备:
# 加载时直接将权重映射到当前使用的设备 ckpt = torch.load("pretrained_weight.pth", map_location=device) model.load_state_dict(ckpt)
- 排查点4:全链路输入数据是否都完成迁移
除了输入特征,标签张量也需要同步迁移到GPU,部分损失函数的计算涉及权重与标签的运算,标签落在CPU也可能触发同类报错:
for batch_x, batch_y in train_loader: # 特征和标签同步迁移 batch_x = batch_x.to(device) batch_y = batch_y.to(device) pred = model(batch_x)
快速验证方法
在执行前向传播前添加两行打印代码,即可快速定位哪部分设备不匹配:
# 打印模型权重所在设备 print("Model weight device:", next(model.parameters()).device) # 打印输入张量所在设备 print("Input tensor device:", batch_x.device)
如果二者输出不一致,针对对应部分做迁移即可解决问题。
内容的提问来源于stack exchange,提问作者PlanetRT
相关产品推荐
相关产品推荐

