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

多GPU环境下PyTorch训练SNN时设备不匹配错误求助

问题描述

在多GPU环境下基于PyTorch运行使用snntorch库snn.Leaky模块构建的类RNN结构脉冲神经网络(SNN)训练任务时,出现设备不匹配错误。

模型代码

net_1 = nn.Sequential(nn.Conv2d(2, 12, 5),
                nn.MaxPool2d(2),
                snn.Leaky(beta=beta, spike_grad=spike_grad, init_hidden=True),
                nn.Conv2d(12, 32, 5),
                nn.MaxPool2d(2),
                snn.Leaky(beta=beta, spike_grad=spike_grad, init_hidden=True),
                nn.Flatten(),
                nn.Linear(32*5*5, 10),
                snn.Leaky(beta=beta, spike_grad=spike_grad, init_hidden=True, output=True)
                )
net_1.cuda()
net = nn.DataParallel(net_1)

训练代码

def forward_pass(net, data):
    spk_rec = []
    utils.reset(net)  # resets hidden states for all LIF neurons in net
    for step in range(data.size(1)):  # data.size(0) = number of time steps
        datas = data[:,step,:,:,:].cuda()
        net = net.to(device)
        spk_out, mem_out = net(datas)

        spk_rec.append(spk_out)

    return torch.stack(spk_rec)

optimizer = torch.optim.Adam(net.parameters(), lr=2e-2, betas=(0.9, 0.999))
loss_fn = SF.mse_count_loss(correct_rate=0.8, incorrect_rate=0.2)
num_epochs = 5
num_iters = 50

loss_hist = []
acc_hist = []
t_spk_rec_sum = []
start = time.time()

net.train()
# training loop
for epoch in range(num_epochs):
    for i, (data, targets) in enumerate(iter(trainloader)):
        data = data.to(device)
        targets = targets.to(device)


        spk_rec = forward_pass(net, data)
        loss_val = loss_fn(spk_rec, targets)

        # Gradient calculation + weight update
        optimizer.zero_grad()
        loss_val.backward()
        optimizer.step()
        # Store loss history for future plotting
        loss_hist.append(loss_val.item())
        print("time :", time.time() - start,"sec")
        print(f"Epoch {epoch}, Iteration {i} \nTrain Loss: {loss_val.item():.2f}")
        acc = SF.accuracy_rate(spk_rec, targets)
        acc_hist.append(acc)
        print(f"Train Accuracy: {acc * 100:.2f}%\n")

错误信息

Traceback (most recent call last):
  File "/home/hubo1024/PycharmProjects/snntorch/multi_gpu_train.py", line 87, in <module>
    spk_rec = forward_pass(net, data)
  File "/home/hubo1024/PycharmProjects/snntorch/multi_gpu_train.py", line 63, in forward_pass
    spk_out, mem_out = net(datas)
  File "/home/hubo1024/anaconda3/envs/spyketorchproject/lib/python3.10/site-packages/torch/nn/modules/module.py", line 1130, in _call_impl
    return forward_call(*input, **kwargs)
  File "/home/hubo1024/anaconda3/envs/spyketorchproject/lib/python3.10/site-packages/torch/nn/parallel/data_parallel.py", line 168, in forward
    outputs = self.parallel_apply(replicas, inputs, kwargs)
  File "/home/hubo1024/anaconda3/envs/spyketorchproject/lib/python3.10/site-packages/torch/nn/parallel/data_parallel.py", line 178, in parallel_apply
    return parallel_apply(replicas, inputs, kwargs, self.device_ids[:len(replicas)])
  File "/home/hubo1024/anaconda3/envs/spyketorchproject/lib/python3.10/site-packages/torch/nn/parallel/parallel_apply.py", line 86, in parallel_apply
    output.reraise()
  File "/home/hubo1024/anaconda3/envs/spyketorchproject/lib/python3.10/site-packages/torch/_utils.py", line 461, in reraise
    raise exception
RuntimeError: Caught RuntimeError in replica 0 on device 0.
Original Traceback (most recent call last):
  File "/home/hubo1024/anaconda3/envs/spyketorchproject/lib/python3.10/site-packages/torch/nn/parallel/parallel_apply.py", line 61, in _worker
    output = module(*input, **kwargs)
  File "/home/hubo1024/anaconda3/envs/spyketorchproject/lib/python3.10/site-packages/torch/nn/modules/module.py", line 1130, in _call_impl
    return forward_call(*input, **kwargs)
  File "/home/hubo1024/anaconda3/envs/spyketorchproject/lib/python3.10/site-packages/torch/nn/modules/container.py", line 139, in forward
    input = module(input)
  File "/home/hubo1024/anaconda3/envs/spyketorchproject/lib/python3.10/site-packages/torch/nn/modules/module.py", line 1130, in _call_impl
    return forward_call(*input, **kwargs)
  File "/home/hubo1024/anaconda3/envs/spyketorchproject/lib/python3.10/site-packages/snntorch/_neurons/leaky.py", line 162, in forward
    self.mem = self.state_fn(input_)
  File "/home/hubo1024/anaconda3/envs/spyketorchproject/lib/python3.10/site-packages/snntorch/_neurons/leaky.py", line 201, in _build_state_function_hidden
    self._base_state_function_hidden(input_) - self.reset * self.threshold
  File "/home/hubo1024/anaconda3/envs/spyketorchproject/lib/python3.10/site-packages/snntorch/_neurons/leaky.py", line 195, in _base_state_function_hidden
    base_fn = self.beta.clamp(0, 1) * self.mem + input_
  File "/home/hubo1024/anaconda3/envs/spyketorchproject/lib/python3.10/site-packages/torch/_tensor.py", line 1121, in __torch_function__
    ret = func(*args, **kwargs)
RuntimeError: Expected all tensors to be on the same device, but found at least two devices, cuda:0 and cpu!

用户疑问:已确认无CPU张量定义代码,单GPU运行正常,当前使用torch.utils.data.DataLoader构建训练加载器,怀疑是问题根源,是否需要更换多GPU训练专用的DataLoader?若需要,哪里能找到相关参考资料?


问题分析与解决

不需要更换torch.utils.data.DataLoader,该加载器完全支持多GPU训练,问题根源在于snn.Leaky模块的隐藏状态(如mem)和参数(如beta)在多GPU复制时未同步到对应设备,具体解决步骤如下:

1. 修正模型初始化逻辑

不要提前调用net_1.cuda(),让nn.DataParallel自动处理模型副本的设备分配,确保所有副本的参数和隐藏状态都在对应GPU上:

# 移除net_1.cuda(),直接初始化后用DataParallel并移到目标设备
net_1 = nn.Sequential(nn.Conv2d(2, 12, 5),
                nn.MaxPool2d(2),
                snn.Leaky(beta=beta, spike_grad=spike_grad, init_hidden=True),
                nn.Conv2d(12, 32, 5),
                nn.MaxPool2d(2),
                snn.Leaky(beta=beta, spike_grad=spike_grad, init_hidden=True),
                nn.Flatten(),
                nn.Linear(32*5*5, 10),
                snn.Leaky(beta=beta, spike_grad=spike_grad, init_hidden=True, output=True)
                )
net = nn.DataParallel(net_1).to(device)

2. 清理forward_pass中的冗余设备操作

nn.DataParallel会自动将数据分发到对应GPU,无需在循环内重复调用net.to(device),移除该行代码:

def forward_pass(net, data):
    spk_rec = []
    utils.reset(net)
    for step in range(data.size(1)):
        datas = data[:,step,:,:,:].to(device)  # 统一用to(device)替代cuda(),适配多设备
        spk_out, mem_out = net(datas)
        spk_rec.append(spk_out)
    return torch.stack(spk_rec)

3. 确保所有参数在GPU上

检查beta等参数是否为GPU张量,若beta是CPU张量,会导致计算时设备不匹配:

# 初始化beta时直接移到目标设备
beta = torch.tensor(0.9).to(device)

4. 修复隐藏状态重置逻辑

utils.reset(net)可能只重置了主GPU上的隐藏状态,需确保所有GPU副本的隐藏状态都被重置。若snntorch的utils.reset不支持多GPU,可手动遍历副本:

def reset_multi_gpu_net(net):
    if isinstance(net, nn.DataParallel):
        for module in net.module.modules():
            if isinstance(module, snn.Leaky):
                module.reset()
    else:
        utils.reset(net)

# 在forward_pass中替换原utils.reset(net)
reset_multi_gpu_net(net)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.17 09:45:35