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

使用DataParallel时PyTorch张量与多GPU适配问题排查

问题

我开发了一套大型机器学习代码,单GPU运行完全正常,现在尝试用多GPU做数据并行适配,却出现了问题。报错信息如下:

RuntimeError: Expected all tensors to be on the same device, but found at least two devices, cuda:1 and cuda:0! (when checking argument for argument index in method wrapper_CUDA__index_select)

相关代码片段

模型文件

import torch
from torch import nn
from functools import partial
import copy

from ..mlp import MLP
from ..basis import gaussian, bessel
from ..conv import GatedGCN


class Encoder(nn.Module):
    """ALIGNN/ALIGNN-d Encoder.
    The encoder must take a PyG graph object `data` and output the same `data`
    with additional fields `h_atm`, `h_bnd`, and `h_ang` that correspond to the atom, bond, and angle embedding.

    The input `data` must have three fields `x_atm`, `x_bnd`, and `x_ang` that describe the atom type
    (in onehot vectors), the bond lengths, and bond/dihedral angles (in radians).
    """

    def __init__(self, num_species, cutoff, dim=128, dihedral=False):
        super().__init__()
        self.num_species = num_species
        self.cutoff = cutoff
        self.dim = dim
        self.dihedral = dihedral

        self.embed_atm = nn.Sequential(MLP([num_species, dim, dim], act=nn.SiLU()), nn.LayerNorm(dim))
        self.embed_bnd = partial(bessel, start=0, end=cutoff, num_basis=dim)
        self.embed_ang = self.embed_ang_with_dihedral if dihedral else self.embed_ang_without_dihedral

    def embed_ang_with_dihedral(self, x_ang, mask_dih_ang):
        cos_ang = torch.cos(x_ang)
        sin_ang = torch.sin(x_ang)

        h_ang = torch.zeros([len(x_ang), self.dim], device=x_ang.device)
        h_ang[~mask_dih_ang, :self.dim // 2] = gaussian(cos_ang[~mask_dih_ang], start=-1, end=1,
                                                        num_basis=self.dim // 2)

        h_cos_ang = gaussian(cos_ang[mask_dih_ang], start=-1, end=1, num_basis=self.dim // 4)
        h_sin_ang = gaussian(sin_ang[mask_dih_ang], start=-1, end=1, num_basis=self.dim // 4)
        h_ang[mask_dih_ang, self.dim // 2:] = torch.cat([h_cos_ang, h_sin_ang], dim=-1)

        return h_ang

    def embed_ang_without_dihedral(self, x_ang, mask_dih_ang):
        cos_ang = torch.cos(x_ang)
        return gaussian(cos_ang, start=-1, end=1, num_basis=self.dim)

    def forward(self, data):
        # Embed atoms
        data.h_atm = self.embed_atm(data.x_atm)

        # Embed bonds
        data.h_bnd = self.embed_bnd(data.x_bnd)

        # Embed angles
        data.h_ang = self.embed_ang(data.x_ang, data.mask_dih_ang)

        return data


class Processor(nn.Module):
    """ALIGNN Processor.
    The processor updates atom, bond, and angle embeddings.
    """

    def __init__(self, num_convs, dim):
        super().__init__()
        self.num_convs = num_convs
        self.dim = dim

        self.atm_bnd_convs = nn.ModuleList([copy.deepcopy(GatedGCN(dim, dim)) for _ in range(num_convs)])
        self.bnd_ang_convs = nn.ModuleList([copy.deepcopy(GatedGCN(dim, dim)) for _ in range(num_convs)])

    def forward(self, data):
        edge_index_G = data.edge_index_G
        edge_index_A = data.edge_index_A

        for i in range(self.num_convs):
            data.h_bnd, data.h_ang = self.bnd_ang_convs[i](data.h_bnd, edge_index_A, data.h_ang)
            data.h_atm, data.h_bnd = self.atm_bnd_convs[i](data.h_atm, edge_index_G, data.h_bnd)

        return data


class Decoder(nn.Module):
    def __init__(self, node_dim, out_dim):
        super().__init__()
        self.node_dim = node_dim
        self.out_dim = out_dim
        self.decoder = MLP([node_dim, node_dim, out_dim], act=nn.SiLU())

    def forward(self, data):
        return self.decoder(data.h_atm)


class ALIGNN(nn.Module):
    """ALIGNN model.
    Can optinally encode dihedral angles.
    """

    def __init__(self, encoder, processor, decoder):
        super().__init__()
        self.encoder = encoder
        self.processor = processor
        self.decoder = decoder

    def forward(self, data):
        data = self.encoder(data)
        data = self.processor(data)
        return self.decoder(data)

训练文件

from tqdm.notebook import trange
from datetime import datetime
import glob
import sys
import os

def train(loader,model,parameters,PIN_MEMORY=False):
    model.train()
    total_loss = 0.0
    model = nn.DataParallel(model, device_ids=[0, 1]).cuda()
    #model = model.to(parameters['device'])
    optimizer = torch.optim.AdamW(model.module.processor.parameters(), lr=parameters['LEARN_RATE'])

    #model = model.to(parameters['device'])
    loss_fn = torch.nn.MSELoss()
    for i,data in enumerate(loader, 0):
        optimizer.zero_grad(set_to_none=True)
        #data = data.to(parameters['device'], non_blocking=PIN_MEMORY)
        data = data.cuda()
        #encoding = model.encoder(data)
        #proc = model.processor(encoding.module)
        #atom_contrib, bond_contrib, angle_contrib = model.decoder(proc.module)
        atom_contrib, bond_contrib, angle_contrib = model(data)

        all_sum = atom_contrib.sum() + bond_contrib.sum() + angle_contrib.sum()

        loss = loss_fn(all_sum, data.y[0][0])
        loss.backward()
        optimizer.step()
        total_loss += loss.item()
    return total_loss / len(loader)

def run_training(data,parameters,model):
    follow_batch = ['x_atm', 'x_bnd', 'x_ang'] if hasattr(data['training'][0], 'x_ang') else ['x_atm']
    loader_train = DataLoader(data['training'], batch_size=parameters['BATCH_SIZE'], shuffle=True, follow_batch=follow_batch)
    loader_valid = DataLoader(data['validation'], batch_size=parameters['BATCH_SIZE'], shuffle=False)

    L_train, L_valid = [], []
    min_loss_train = 1.0E30
    min_loss_valid = 1.0E30

    stats_file = open(os.path.join(parameters['model_dir'],'loss.data'),'w')
    stats_file.write('Training_loss     Validation loss\n')
    stats_file.close()
    for ep in range(parameters['num_epochs']):
        stats_file = open(os.path.join(parameters['model_dir'], 'loss.data'), 'a')
        print('Epoch ',ep,' of ',parameters['num_epochs'])
        sys.stdout.flush()
        loss_train = train(loader_train, model, parameters);
        L_train.append(loss_train)
        loss_valid = test_non_intepretable(loader_valid, model, parameters)
        L_valid.append(loss_valid)
        stats_file.write(str(loss_train) + '     ' + str(loss_valid) + '\n')
        if loss_train < min_loss_train:
            min_loss_train = loss_train
            if loss_valid < min_loss_valid:
                min_loss_valid = loss_valid
                if parameters['remove_old_model']:
                    model_name = glob.glob(os.path.join(parameters['model_dir'], 'model_*'))
                    if len(model_name) > 0:
                        os.remove(model_name[0])
                now = datetime.now().strftime('%Y-%m-%d_%H-%M-%S')
                print('Min train loss: ', min_loss_train, ' min valid loss: ', min_loss_valid, ' time: ', now)
                torch.save(model.state_dict(), os.path.join(parameters['model_dir'], 'model_' + str(now)))
        stats_file.close()
        if loss_train < parameters['train_tolerance'] and loss_valid < parameters['train_tolerance']:
            print('Validation and training losses satisy set tolerance...exiting training loop...')
            break

数据结构(批量大小为2时)

Graph_DataBatch(atoms=[2], edge_index_G=[2, 89966], edge_index_A=[2, 1479258], x_atm=[5184, 5], x_atm_batch=[5184], x_atm_ptr=[3], x_bnd=[89966], x_bnd_batch=[89966], x_bnd_ptr=[3], x_ang=[1479258], x_ang_batch=[1479258], x_ang_ptr=[3], mask_dih_ang=[1479258], atm_amounts=[6], bnd_amounts=[6], ang_amounts=[6], y=[179932, 1])

核心疑问

  1. 我的理解:已将批处理数据发送至GPU,模型参数也部署在GPU上,DataParallel会自动拆分数据并分发至各个GPU,这个理解是否正确?
  2. 我的代码是否确实在执行上述逻辑?
  3. 该错误是否与此相关?若无关,该错误想传达什么信息?

另外,我排查时发现Encoder的forward函数内data.x_atm始终在cuda:0,即使nvidia-smi显示模型部署在cuda:0和cuda:1上,尝试多种X.to('cuda')或X.cuda()调用都没改变张量设备。


分析与解答

1. 关于DataParallel的核心理解是否正确

这个理解的核心逻辑是对的:nn.DataParallel的工作机制是将模型复制到指定的所有GPU上,把输入数据按batch维度拆分后分发到各个GPU做前向计算,最后将各GPU的计算结果收集到主GPU(默认是device_ids的第一个,这里为cuda:0)完成损失计算、反向传播等后续操作。主GPU负责参数更新,其他GPU仅执行计算任务,所有模型参数会被复制到各个GPU,反向传播时梯度会汇总到主GPU的参数上统一更新。

2. 代码是否执行了正确的逻辑

你的代码没有执行正确的多GPU并行逻辑,核心问题有三点:

  • 重复包装DataParallel:在train函数中,每次调用都会重新执行model = nn.DataParallel(model, device_ids=[0, 1]).cuda(),导致模型被多次包装、参数重复复制到GPU,epoch间模型状态混乱。正确做法是在run_training开始时就完成DataParallel包装,而非每个train循环重复操作。
  • 数据与模型设备绑定时机错误:你用data = data.cuda()将数据默认放到cuda:0,但此时模型尚未被包装成DataParallel,后续分发逻辑会出错。正确流程是先包装模型,再将数据放到主GPU,由DataParallel自动拆分分发到其他GPU。
  • 优化器初始化时机错误:在train函数内每次epoch都重新创建优化器,会丢失动量、学习率调度等状态,且优化器绑定的是每次新包装的model.module参数,导致参数更新逻辑混乱。

3. 错误是否与此相关,错误的含义

这个错误直接由上述逻辑错误导致。错误信息表示执行index_select操作时,参与计算的张量分别位于cuda:0和cuda:1,无法完成跨设备计算。

具体原因:

  • 重复包装DataParallel导致模型参数被多次复制到cuda:0和cuda:1,但数据始终固定在cuda:0,当模型副本在cuda:1尝试处理数据时,部分关联张量(如edge_index)未被正确分发,出现设备不匹配。
  • PyG的GraphBatch对象调用data.cuda()时,可能未将所有内部张量(如edge_index_G、edge_index_A)同步移动到GPU,部分张量仍留在CPU或cuda:0,与cuda:1上的模型参数触发设备冲突。

你排查时发现data.x_atm始终在cuda:0,是因为data.cuda()直接将整个batch放到了主GPU,但由于模型在train函数内才被包装,DataParallel无法正常完成数据拆分分发,最终导致部分张量与模型参数设备不匹配,触发报错。


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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.21 04:58:11