PyTorch多GPU训练Transformer遇设备不匹配RuntimeError求助
错误根源
报错RuntimeError: Expected all tensors to be on the same device, but found at least two devices, cuda:1 and cuda:0!的核心原因是:Transformer模型中动态生成的位置编码张量(如pos)没有和当前GPU上的模型参数处于同一设备。
当使用nn.DataParallel时,模型会被复制到所有可用GPU,每个GPU处理部分batch数据。如果你的Encoder/Decoder的forward函数中,位置编码张量是基于全局device变量创建的,或者默认在CPU生成,就会出现部分张量在cuda:0、模型参数在cuda:1的设备不匹配情况。
解决方案
1. 修正位置编码张量的设备绑定
在Encoder和Decoder的forward函数中,生成位置编码张量时,不要依赖全局device变量,而是绑定到模型已有参数的设备。
假设你的Encoder forward中有类似生成位置张量的代码,修改如下:
# 错误示例:直接用全局device或默认CPU生成 pos = torch.arange(0, src.shape[1]).unsqueeze(0) # 修正后:使用模型参数所在设备 pos = torch.arange(0, src.shape[1]).unsqueeze(0).to(self.tok_embedding.weight.device)
这样无论模型被复制到哪个GPU,位置张量都会自动适配当前设备,避免不匹配。
2. 规范模型初始化与DataParallel使用
不需要给Encoder/Decoder单独套nn.DataParallel,只需给顶层Transformer模型做并行处理,同时确保所有子模块正确移到设备:
# 先创建编码器和解码器(无需提前传入device,由后续to(device)统一处理) enc = Encoder(INPUT_DIM, HIDDEN_DIM, ENC_LAYERS, ENC_HEADS, ENC_PF_DIM, ENC_DROPOUT) dec = Decoder(OUTPUT_DIM, HIDDEN_DIM, DEC_LAYERS, DEC_HEADS, DEC_PF_DIM, DEC_DROPOUT) # 先将Transformer模型移到device,再套DataParallel transformer = Transformer(enc, dec, SRC_PAD_IDX, TRG_PAD_IDX).to(device) model = nn.DataParallel(transformer)
如果你的Encoder/Decoder的__init__方法中有预创建的张量(如固定位置编码表),也要确保这些张量在初始化时就移到正确设备,或者在forward时动态绑定到当前设备。
3. 确保输入数据设备正确
训练循环中,必须将输入的src和trg张量移到指定设备:
for src, trg in train_iterator: src = src.to(device) trg = trg.to(device) # 后续训练步骤...
4. 推荐:改用DistributedDataParallel
nn.DataParallel是单进程多线程的并行方式,性能和稳定性不如官方推荐的DistributedDataParallel(DDP),尤其适合多GPU场景。使用示例:
import torch.distributed as dist from torch.nn.parallel import DistributedDataParallel as DDP from torch.utils.data.distributed import DistributedSampler # 初始化分布式进程组 dist.init_process_group(backend='nccl') local_rank = dist.get_rank() device = torch.device(f'cuda:{local_rank}') # 创建并移动模型到当前GPU enc = Encoder(...).to(device) dec = Decoder(...).to(device) transformer = Transformer(enc, dec, SRC_PAD_IDX, TRG_PAD_IDX).to(device) model = DDP(transformer, device_ids=[local_rank]) # 使用DistributedSampler加载数据(自动分配batch到各进程) train_sampler = DistributedSampler(train_dataset) train_iterator = DataLoader( train_dataset, sampler=train_sampler, batch_size=BATCH_SIZE, num_workers=NUM_WORKERS )
启动脚本时使用torchrun命令:
torchrun --nproc_per_node=2 your_training_script.py
内容的提问来源于stack exchange,提问作者kyouichi

