如何在PyTorch Lightning中手动将大模型拆分至多GPU运行?
多GPU部署大Transformer模型的问题与解决方案
问题背景
单GPU运行含Transformer模块列表的大模型时,执行到第17个隐藏层触发CUDA out of memory error;尝试以下方案后仍存在问题:
- 自定义
SplitModel结合DDP策略,所有Transformer层仍固定在cuda:0,OOM问题未解决 - 在
forward中手动切换设备,OOM问题解决但反向传播触发RuntimeError: grad.device() == bucket_view.device()内部断言错误 - 使用FSDP模型分片策略,遇到批归一化张量类型不匹配(
torch.cuda.FloatTensor与torch.cuda.HalfTensor)的错误
核心结论
自定义反向传播层手动切换设备理论可行,但实现复杂且易破坏分布式训练逻辑,不推荐作为优先方案。更高效的解决路径是采用手动模型并行或修复FSDP配置,以下是具体方案:
方案1:手动模型并行(优先推荐)
放弃DDP(数据并行),改用模型并行——在初始化阶段就将模型分段固定到对应GPU,仅在设备间传递张量而非移动参数,彻底避免梯度设备不匹配问题:
import torch import torch.nn as nn import pytorch_lightning as pl class SplitModel(pl.LightningModule): def __init__(self, transformers, segment1, segment2): super().__init__() self.device1 = torch.device('cuda:0') self.device2 = torch.device('cuda:1') # 固定模型段到对应设备 self.segment1 = segment1.to(self.device1) # 拆分Transformer层:前17层到cuda:0,剩余到cuda:1 self.transformers_1 = nn.ModuleList(transformers[:17]).to(self.device1) self.transformers_2 = nn.ModuleList(transformers[17:]).to(self.device2) self.segment2 = segment2.to(self.device2) self.loss_fn = nn.CrossEntropyLoss().to(self.device1) def forward(self, x): # 输入转移到cuda:0执行输入层 x = x.to(self.device1) x = self.segment1(x) # 执行前17层Transformer for transformer in self.transformers_1: attn, ff = transformer x = attn(x) + x x = ff(x) + x # 张量转移到cuda:0执行剩余Transformer层 x = x.to(self.device2) for transformer in self.transformers_2: attn, ff = transformer x = attn(x) + x x = ff(x) + x # 执行输出层后转移回cuda:0计算损失 x = self.segment2(x) return x.to(self.device1) def training_step(self, batch, batch_idx): inputs, labels = batch labels = labels.to(self.device1) outputs = self(inputs) loss = self.loss_fn(outputs, labels) self.log('train_loss', loss) return loss def configure_optimizers(self): return torch.optim.Adam(self.parameters(), lr=1e-3) # 训练配置:无需DDP,直接指定多设备 trainer = pl.Trainer( precision="16-mixed", accelerator="cuda", devices=[0, 1], strategy="auto" ) data_loader = # 你的数据加载器 trainer.fit(SplitModel(transformers, segment1, segment2), data_loader)
关键细节:
- 模型参数在初始化时就绑定到对应GPU,避免
forward中动态移动参数导致的梯度混乱 - 仅在设备间传递中间张量(
x),参数始终固定在各自设备 - 无需DDP,因为这是模型并行(不同层分布在不同GPU)而非数据并行(同层复制到多GPU)
方案2:修复FSDP的批归一化类型问题
FSDP的张量类型不匹配是混合精度下的自动转换冲突,可通过两种方式解决:
- 显式固定BN层 dtype
# 自定义固定dtype的BN层 class FP32BatchNorm(nn.BatchNorm1d): def __init__(self, num_features): super().__init__(num_features) self.weight.data = self.weight.data.to(torch.float32) self.bias.data = self.bias.data.to(torch.float32) self.running_mean = self.running_mean.to(torch.float32) self.running_var = self.running_var.to(torch.float32) # 递归替换模型中所有BN层 def replace_bn(model): for name, module in model.named_children(): if isinstance(module, (nn.BatchNorm1d, nn.BatchNorm2d, nn.BatchNorm3d)): setattr(model, name, FP32BatchNorm(module.num_features)) else: replace_bn(module) # 初始化模型后调用 replace_bn(model)
- 关闭FSDP自动混合精度
trainer = pl.Trainer( precision="16-mixed", accelerator="cuda", devices=[0,1], strategy="fsdp", fsdp_config={"mixed_precision": False} )
方案3:自定义反向传播层(备选)
仅当上述方案均不可行时考虑,需手动处理梯度的设备转移,示例如下:
class DeviceTransferLayer(nn.Module): def __init__(self, target_device): super().__init__() self.target_device = target_device def forward(self, x): return x.to(self.target_device) def backward(self, grad_output): # 反向传播时将梯度转移回原设备 return grad_output.to(self.target_device) # 在forward中使用示例 layer_to_device2 = DeviceTransferLayer(torch.device('cuda:1')) # 正向传递:x = layer_to_device2(x) # 反向时梯度自动转移回原设备
注意:该方式需与模型并行逻辑严格对齐,极易出错且无法兼容DDP,仅作为极端场景的备选方案。
内容的提问来源于stack exchange,提问作者Sam Moldenha
相关产品推荐
相关产品推荐

