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

替换预训练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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.29 18:13:33