替换预训练ResNet50后出现CUDA/CPU设备不匹配RuntimeError的解决咨询
设备不匹配问题:预训练ResNet50在联邦学习get()方法中触发RuntimeError
问题背景
自定义PyTorch模型结合联邦学习send/get流程可在指定设备(cuda/cpu)正常运行,替换为预训练ResNet50并执行model.to(device)后,调用model.get()时触发错误:
RuntimeError: Expected object of device type cuda but got device type cpu for argument #1 'self' in call to th_set
原自定义模型及训练代码
use_cuda = not args.no_cuda and torch.cuda.is_available() device = torch.device("cuda" if use_cuda else "cpu") class Net(nn.Module): def __init__(self): super(Net, self).__init__() self.conv1 = nn.Conv2d(1, 20, 5, 1) self.conv2 = nn.Conv2d(20, 50, 5, 1) self.fc1 = nn.Linear(4*4*50, 500) self.fc2 = nn.Linear(500, 10) def forward(self, x): x = F.relu(self.conv1(x)) ... x = self.fc2(x) return F.log_softmax(x, dim=1) model = Net().to(device) def train(args, model, device, train_loader, optimizer, epoch): model.train() for batch_idx, (data, target) in enumerate(federated_train_loader): model.send(data.location) data, target = data.to(device), target.to(device) output = model(data) model.get() if batch_idx % args.log_interval == 0: loss = loss.get() print('Train Epoch: {} [{}/{} ({:.0f}%)]\tLoss: {:.6f}'.format( epoch, batch_idx * args.batch_size, len(train_loader) * args.batch_size, 100. * batch_idx / len(train_loader), loss.item()))
替换为预训练ResNet50的代码
model = models.resnet50(pretrained=True) model.to(device)
get方法实现片段
def module_get_(nn_self): for element_iter in tensor_iterator(nn_self): for p in element_iter(): p.get_() if isinstance(nn_self.forward, Plan): nn_self.forward.get() return nn_self self.torch.nn.Module.get_ = module_get_ self.torch.nn.Module.get = module_get_
错误原因分析
预训练ResNet50包含多层嵌套子模块(如layer1、layer2等),当前get()方法仅遍历顶层模块参数,未递归处理嵌套子模块,导致部分子模块参数仍停留在CPU;同时联邦学习send/get流程会改变参数设备状态,get()时未同步所有层级参数的设备信息,最终触发设备不匹配错误。
解决方法
1. 修改get()方法,递归遍历所有子模块
更新module_get_函数,确保递归处理所有嵌套子模块的参数:
def module_get_(nn_self): # 处理当前模块参数 for element_iter in tensor_iterator(nn_self): for p in element_iter(): p.get_() # 递归处理所有子模块 for child in nn_self.children(): module_get_(child) if isinstance(nn_self.forward, Plan): nn_self.forward.get() return nn_self self.torch.nn.Module.get_ = module_get_ self.torch.nn.Module.get = module_get_
2. 确保预训练模型全量加载到目标设备
初始化时直接将模型移动到目标设备,并校验参数设备:
model = models.resnet50(pretrained=True).to(device) # 验证所有参数是否已移动到目标设备 for param in model.parameters(): assert param.device == device, "部分参数未正确迁移到目标设备"
3. 训练流程中同步数据与模型设备
训练时将数据移动到模型当前所在设备(而非固定本地device),适配send/get后的设备变化:
def train(args, model, device, train_loader, optimizer, epoch): model.train() for batch_idx, (data, target) in enumerate(federated_train_loader): model.send(data.location) # 数据移动到模型当前所在设备 data, target = data.to(next(model.parameters()).device), target.to(next(model.parameters()).device) output = model(data) model.get() if batch_idx % args.log_interval == 0: loss = loss.get() print('Train Epoch: {} [{}/{} ({:.0f}%)]\tLoss: {:.6f}'.format( epoch, batch_idx * args.batch_size, len(train_loader) * args.batch_size, 100. * batch_idx / len(train_loader), loss.item()))
4. 修复tensor_iterator的遍历逻辑
如果tensor_iterator仅遍历顶层模块,修改为递归遍历所有子模块:
def tensor_iterator(module): yield module.parameters for child in module.children(): yield from tensor_iterator(child)
内容的提问来源于stack exchange,提问作者seni
相关产品推荐
相关产品推荐

