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

如何在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,能更高效处理参数同步和输入拆分,性能更优。切换步骤:
    1. 初始化分布式环境:
      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
      
    2. 使用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
      )
      
    3. 用DDP包装模型:
      local_rank = setup_distributed()
      transformer = transformer.to(local_rank)
      transformer = DDP(transformer, device_ids=[local_rank])
      
    4. 用torchrun启动训练脚本(指定8张GPU):
      torchrun --nproc_per_node=8 your_training_script.py
      

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.20 14:12:41