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

如何在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的张量类型不匹配是混合精度下的自动转换冲突,可通过两种方式解决:

  1. 显式固定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)
  1. 关闭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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.09 20:57:08