如何在8张GPU上并行化Transformer机器翻译模型并解决AttnMask错误
问题描述
参照原论文实现Transformer机器翻译模型,效果达标但算力需求较高,使用配备8张GPU的设备运行模型,尝试通过以下代码做并行化:
transformer = nn.DataParallel(transformer) transformer = transformer.to(DEVICE)
运行后触发注意力掩码维度不匹配的错误:
File "C:\Projects\MT005\.venv\Lib\site-packages\torch\nn\functional.py", line 5382, in multi_head_attention_forward raise RuntimeError(f"The shape of the 2D attn_mask is {attn_mask.shape}, but should be {correct_2d_size}.") RuntimeError: The shape of the 2D attn_mask is torch.Size([8, 64]), but should be (4, 4).
解决方案
错误根源是nn.DataParallel会自动将输入batch拆分到各个GPU,但你的注意力掩码(attn_mask)未同步做对应拆分,导致每个GPU上的掩码维度和当前子batch不匹配。以下是具体解决方法:
- 动态生成注意力掩码:不要从外部传入固定形状的掩码,而是在模型的
forward函数内,基于当前GPU上的输入序列长度动态生成。比如自注意力的上三角掩码可以这样生成:def forward(self, src, tgt): # 基于当前tgt的序列长度生成自注意力掩码 tgt_seq_len = tgt.size(1) tgt_mask = torch.triu(torch.ones(tgt_seq_len, tgt_seq_len, device=tgt.device), diagonal=1).bool() # 后续注意力计算使用该动态生成的掩码 output = self.decoder(tgt, src, tgt_mask=tgt_mask) return output - 调整掩码的拆分逻辑:如果必须从外部传入掩码,要确保
nn.DataParallel能正确拆分掩码。可以在传入前将掩码调整为[batch_size, seq_len]的形状,并且在模型forward中不对掩码做额外的维度扩展,让DataParallel自动按batch维度拆分。 - 切换为DistributedDataParallel(推荐):
nn.DataParallel是单进程多GPU架构,对Transformer这类含注意力机制的模型兼容性较差。DistributedDataParallel(DDP)采用多进程多GPU,能更高效处理参数同步和输入拆分,性能更优。切换步骤:- 初始化分布式环境:
import torch.distributed as dist from torch.nn.parallel import DistributedDataParallel as DDP from torch.utils.data.distributed import DistributedSampler def setup_distributed(): dist.init_process_group(backend='nccl') local_rank = dist.get_rank() torch.cuda.set_device(local_rank) return local_rank - 使用
DistributedSampler拆分数据集,保证每个GPU获取独立的子数据集:train_sampler = DistributedSampler(train_dataset) train_loader = torch.utils.data.DataLoader( train_dataset, batch_size=batch_size, sampler=train_sampler, num_workers=4 ) - 用DDP包装模型:
local_rank = setup_distributed() transformer = transformer.to(local_rank) transformer = DDP(transformer, device_ids=[local_rank]) - 用
torchrun启动训练脚本(指定8张GPU):torchrun --nproc_per_node=8 your_training_script.py
- 初始化分布式环境:
内容的提问来源于stack exchange,提问作者dsb
相关产品推荐
相关产品推荐

