PySyft联邦训练ResNet50时model.get()触发CUDA设备不匹配错误
PySyft训练ResNet50时model.get()报错RuntimeError的解决方法
问题场景
使用PySyft联邦学习框架训练自定义封装的ResNet50模型时,在model.get()步骤触发设备类型不匹配错误,报错信息如下:
RuntimeError: Expected object of device type cuda but got device type cpu for argument #1 'self' in call to th_set
相关代码
import torch import torch.nn as nn import torch.nn.functional as F from torchvision import models import syft as sy class Net(nn.Module): def __init__(self): super(Net, self).__init__() self.model = models.resnet50(pretrained=False) self.fc1 = nn.Linear(2048,2048) self.fc2 = nn.Linear(2048, 3) self.dropout = nn.Dropout(0.3) def forward(self, x): x = self.model.conv1(x) x = self.model.bn1(x) x = self.model.relu(x) x = self.model.maxpool(x) # 省略中间ResNet层的调用 x = self.model.avgpool(x) x = x.view(-1,2048*1*1) x = nn.functional.relu(self.fc1(x)) x = self.dropout(x) x = nn.functional.softmax(self.fc2(x), dim=1) return x 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) optimizer.zero_grad() output = model(data) loss = F.nll_loss(output, target) loss.backward() optimizer.step() 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()))
问题原因
报错核心是模型参数的设备分布不一致:
- 初始化时,
models.resnet50()默认在CPU上创建参数,即使外部Net实例被移到CUDA,子模块self.model的参数可能仍留在CPU; model.send()将模型发送到远程worker后,调用model.get()时默认会把模型取回CPU,但原模型的部分参数仍绑定在CUDA设备上,导致设备类型冲突。
解决方法
1. 确保模型完全加载到目标设备
在创建Net实例后,立即将整个模型(包括子模块)移到指定设备:
model = Net().to(device) # device为cuda设备,比如torch.device("cuda:0")
2. 调用model.get()时指定目标设备
修改train函数中的model.get(),显式指定取回后的设备为训练用的CUDA设备:
model.get(device=device)
3. 修正损失函数与输出的不匹配(额外优化)
代码中用F.nll_loss计算损失,但模型输出是softmax结果,而nll_loss要求输入是log_softmax输出,这会导致损失计算异常。建议修改forward函数的输出层:
# 替换原softmax行 x = nn.functional.log_softmax(self.fc2(x), dim=1)
验证步骤
- 重新初始化模型并确认所有参数都在CUDA上:
print(next(model.parameters()).device) # 应输出cuda:x
- 运行训练代码,检查
model.get()是否再触发设备错误。
内容的提问来源于stack exchange,提问作者seni
相关产品推荐
相关产品推荐

