You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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)

验证步骤

  1. 重新初始化模型并确认所有参数都在CUDA上:
print(next(model.parameters()).device)  # 应输出cuda:x
  1. 运行训练代码,检查model.get()是否再触发设备错误。

内容的提问来源于stack exchange,提问作者seni

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.29 04:53:27