PyTorch如何便捷迁移含子模块与包装模块的模型至CUDA?
问题:PyTorch模型整体迁移至CUDA失败的解决方法
场景与代码
想要将PyTorch模型整体迁移至CUDA运行,代码示例如下:
import torch import torch.nn as nn class SubModel(nn.Module): def __init__(self): super(SubModel, self).__init__() self.conv1 = nn.Conv1d(in_channels=1, out_channels=1, kernel_size=2) def forward(self, x): print(f"x type:{type(x)}") print(f"weight type:{type(self.conv1.weight)}") return self.conv1(x) class WrapperModel(nn.Module): def __init__(self, count): super(WrapperModel, self).__init__() self.blocks = [] for i in range(count): self.blocks.append(SubModel()) def forward(self, x): for block in self.blocks: x = block(x) return x class MyModel(nn.Module): def __init__(self): super(MyModel, self).__init__() self.conv = nn.Conv1d(in_channels=1, out_channels=1, kernel_size=2) self.wrapper = WrapperModel(2) def forward(self, x): x = self.conv(x) x = self.wrapper(x) return x
报错信息
执行时在SubModel的forward方法调用self.conv1(x)时出现以下错误:
RuntimeError: Input type (torch.cuda.FloatTensor) and weight type (torch.FloatTensor) should be the same
打印输出:
x type:<class 'torch.Tensor'> weight type:<class 'torch.nn.parameter.Parameter'>
尝试单独执行model.wrapper.to(device)无法解决问题,手动将SubModel的conv1移至CUDA可正常运行,但操作繁琐。
问题原因
核心问题出在WrapperModel的实现:你用普通Python列表self.blocks存储子模块,PyTorch的模块追踪机制无法识别列表内的子模块,调用.to(device)时只会处理模型中被注册的模块,列表里的SubModel参数不会被自动迁移到CUDA,导致输入张量在CUDA、权重参数在CPU,引发设备不匹配错误。
解决方案
方案1:使用nn.ModuleList替代普通列表(推荐)
PyTorch提供nn.ModuleList专门用于存储子模块,它会被模型自动追踪,修改WrapperModel的__init__方法即可:
class WrapperModel(nn.Module): def __init__(self, count): super(WrapperModel, self).__init__() self.blocks = nn.ModuleList() # 替换为ModuleList for i in range(count): self.blocks.append(SubModel())
修改后,只需要执行一次整体迁移命令,就能将所有子模块、包装模块的参数全部移至CUDA:
device = torch.device("cuda" if torch.cuda.is_available() else "cpu") model = MyModel().to(device)
方案2:手动注册子模块(不推荐)
如果坚持使用普通列表,需要在WrapperModel中逐个注册子模块:
class WrapperModel(nn.Module): def __init__(self, count): super(WrapperModel, self).__init__() self.blocks = [] for i in range(count): sub_model = SubModel() self.add_module(f"submodel_{i}", sub_model) # 手动注册子模块 self.blocks.append(sub_model)
此方法需要手动维护子模块的命名,不如nn.ModuleList简洁高效。
内容的提问来源于stack exchange,提问作者Hao Wu
相关产品推荐
相关产品推荐

