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

FSDP未在多GPU间分片模型的问题排查求助

问题:FSDP未在多GPU间分片模型导致OOM,手动分配模型部分到不同GPU的方法

我有一个大模型,即使batch size设为1也无法在单GPU上容纳,因此选择FSDP解决显存问题。手动封装了部分层并监控GPU内存,发现2块GPU中仅GPU 0被使用,GPU 1仅分配3MB内存,调用transformer_1模块时触发CUDA内存分配错误。想知道FSDP未自动分片的原因,以及是否可以手动将模型不同部分分配到不同GPU处理。

内存监控显示:

  • 调用x = self.transformer_0(x)前:[(16575, 81559), (3, 81559)]
  • 调用x = self.transformer_0(x)后:[(50421, 81559), (3, 81559)]
  • 调用x = self.transformer_1(x)时触发CUDA内存分配错误

模型代码

class LightningModel_distributed_patch_16(pl.LightningModule):
    def __init__(self, *args, **kwargs):
        super(LightningModel_distributed_patch_16, self).__init__(*args, **kwargs)
        total_depths = 10
        self.original_model =  patchedSeg_l(attentions=[attentions.SelfAttention_big, attentions.SelfAttention_big], depths=[total_depths, total_depths], positioned=[True, True], dim_3=False, attentions1_depth=(4,4), img_h=256, img_w=256, heads=12, patches=(16,16))
        self.original_model.pos_embed = Summer(PositionalEncodingPermute3D(3))
        self.pos_embed = self.original_model.pos_embed
        self.transformers1 = self.original_model.transformers1
        self.to_patches = self.original_model.to_patches
        self.conv1 = self.original_model.conv1
        self.conv2 = self.original_model.conv2
        self.bn1 = self.original_model.bn1
        self.bn2 = self.original_model.bn2
        self.bn3 = self.original_model.bn2
        self.relu = self.original_model.relu
        self.transformer_0 = self.original_model.transformer_0

        #all the modules before this on cuda:0
        #all the modules after this on cuda:1
        self.transformer_1 = self.original_model.transformer_1
        self.transformers_out = self.original_model.transformers_out
        self.conv_out = self.original_model.conv_out
        self.out_layer = self.original_model.out_layer
        self.from_patches = attentions.FromPatches(16,16, False)

        self.batch_size = 256
        self.store_losses_printing = []
        self.store_correct_number = []
        self.store_correct_percent = []
        self.store_losses_printing_val = []

        self.s_logits = []

    def configure_sharded_model(self):
        self.trainer.model.transformers1 = wrap(self.trainer.model.transformers1)
        self.trainer.model.transformer_0 = wrap(self.trainer.model.transformer_0)
        self.trainer.model.transformer_1 = wrap(self.trainer.model.transformer_1)
        self.trainer.model.transformers_out = wrap(self.trainer.model.transformers_out)

    def forward(self, parts):
        x, indexes_a, batch_index = parts
        x = torch.squeeze(self.pos_embed(x.transpose(1,2)).transpose(1,2))
        x = x[batch_index]
        x = self.conv1(x)
        x = self.bn1(x)
        x = self.transformers1(x)
        x = x.to(torch.float32)
        x = self.relu(x)
        x = self.to_patches(x)
        N, P, C, X, Y = x.shape
        x = x.view(N*P, C, X, Y)
        # x = torch.index_select(x, 0, indexes_a)
        x = self.conv2(x)
        x = self.bn2(x)
        x = self.relu(x)
        #this is where I print the amount of memory taken up
        print(get_gpu_memory()) 
        x = self.transformer_0(x)
        print(get_gpu_memory()) 
        #transformer_1 is where I get a cuda out of memory error
        x = self.transformer_1(x)
        x = self.bn3(x)
        x = self.original_model.from_patched(x).unsqueeze(0)
        x = self.transformers_out(x)
        x = x.to(torch.float32)
        x = self.conv_out(x)
        x = self.out_layer(x)
        return self.to_patches(x).flatten(0,1).contiguous()

    def configure_optimizers(self):
        lr_list = [1e-3, 1e-4, 1e-2]
        initial_lr = lr_list[0]
        optimizer = torch.optim.SGD(self.trainer.model.parameters(), lr=initial_lr, momentum=0.3)
        scheduler = CustomLRScheduler(optimizer, lr_list)
        return [optimizer], [scheduler]

    def cross_entropy_loss(self, logits, labels):
        loss = torch.nn.functional.cross_entropy(logits.float(), labels)
        return loss

    def training_step(self, train_batch, batch_idx):
        input_batch, output_batch = train_batch
        output_batch = output_batch.squeeze()
        input_batch[1] = input_batch[1].squeeze()
        if(len(output_batch.shape) == 2):
            output_batch = torch.unsqueeze(output_batch, 0)
        total_params = sum(p.numel() for p in self.parameters())
        logits = self.forward(input_batch)
        gpu_memory = get_gpu_memory()
        for idx, (used, total) in enumerate(gpu_memory):
            print('gpu',idx,'used',used,'of memory')
            print('gpu',idx,'used',total,'total memory')
            self.log(f'gpu_{idx}_memory_used', used, on_step=True, on_epoch=False)
            self.log(f'gpu_{idx}_memory_total', total, on_step=True, on_epoch=False)
        print(gpu_memory)

        print(logits.shape, output_batch.shape)
        correct=logits.argmax(dim=1).eq(output_batch).sum().item()
        total=self.batch_size
        loss = self.cross_entropy_loss(logits, output_batch)
        self.log("train_loss", loss, on_step=True, on_epoch=True, prog_bar=True, logger=True, sync_dist=True)
        self.store_losses_printing.append(loss) 
        logs={"train_loss": loss}
        self.store_correct_number.append(correct)
        self.store_correct_percent.append(correct / (output_batch.shape[0] * output_batch.shape[1] * output_batch.shape[2]))
        batch_dictionary={
            #REQUIRED: It ie required for us to return "loss"
            "loss": loss,
            
            #optional for batch logging purposes
            "log": logs,
 
            # info to be used at epoch end 
            "correct": correct,
            "total": total
        }
        return batch_dictionary
    
    def on_train_epoch_end(self):
        loss = sum(output.item() for output in self.store_losses_printing) / len(self.store_losses_printing)
        correct = sum(output for output in self.store_correct_number) / len(self.store_correct_number)
        correct_percent = sum(output for output in self.store_correct_percent) / len(self.store_correct_percent)
        self.logger.experiment.add_scalar("average_correct", correct, global_step=self.current_epoch)
        self.logger.experiment.add_scalar("average_correct_percent", correct_percent, global_step=self.current_epoch)
        print('')
        print('For epoch {}: Average returned loss: {}, Average correct: {}, Correct Percentage Average: {}'.format(self.current_epoch, loss, correct, correct_percent))
        self.store_losses_printing.clear()
        self.store_correct_percent.clear()
        self.store_correct_number.clear()

    def validation_step(self, batch, batch_idx):
        # print(self.trainer.model)
        if next(self.conv1.parameters()).device != self.device:
            self.conv1.to(self.device)
            self.conv2.to(self.device)
            self.bn1.to(self.device)
            self.bn2.to(self.device)
            self.conv_out.to(self.device)

        inputs, output_batch, name = batch
        print(inputs[0].dtype)
        
        output_batch = output_batch.cpu()
        output_batch[torch.where(output_batch == 2)] = 1
        inputs[1] = inputs[1].squeeze()
        output_batch = output_batch.squeeze()
        if(len(output_batch.shape) == 2):
            output_batch = torch.unsqueeze(output_batch, 0)
        
        logits = self.forward(inputs).cpu()
        correct=logits.argmax(dim=1).eq(output_batch).sum().item()
        total=self.batch_size

        loss = self.cross_entropy_loss(logits, output_batch)

        self.log("val_loss", loss, on_step=True, on_epoch=True, prog_bar=True, logger=True, sync_dist=True)
        self.logger.experiment.add_scalar("val_current_loss", loss.item(), global_step=0)
        self.logger.experiment.add_scalar("val_current_correct", correct, global_step=0)
        gpu_memory = get_gpu_memory()
        for idx, (used, total) in enumerate(gpu_memory):
            print('gpu',idx,'used',used,'of memory')
            print('gpu',idx,'used',total,'total memory')
        print(gpu_memory)

        return {'val_loss':loss, 'total':total, 'correct':correct}

GPU内存监控代码

import subprocess

def get_gpu_memory():
    """Get the current GPU memory usage."""
    try:
        result = subprocess.check_output(
            ['nvidia-smi', '--query-gpu=memory.used,memory.total', '--format=csv,nounits,noheader'],
            encoding='utf-8'
        )
        # The result is a string, split it into a list of strings
        gpu_memory = [line.split(', ') for line in result.strip().split('\n')]
        return [(int(used), int(total)) for used, total in gpu_memory]
    except subprocess.CalledProcessError as e:
        print("Error fetching GPU memory usage:", e)
        return []

训练器配置代码

bregmas = torch.load("dataset/training/bregmas_blurred_n.pt")
output_batch = torch.load("dataset/training/right_masks_blurred.pt")
input_batch = torch.load("dataset/training/right_inputs_blurred.pt")

names_val = torch.load("dataset/validation/names.pt")
bregmas_val = torch.load("dataset/validation/bregmas.pt")
output_val = torch.load("dataset/validation/right_masks.pt")
input_val = torch.load("dataset/validation/right_inputs.pt")

r_model = LightningModel_distributed_patch_16()

dataset = MyDataset_model_num(input_batch, output_batch, idx_num, indexes=bregmas, batch_size=1024)
train_dataloader = DataLoader(dataset, batch_size=1, num_workers=4, shuffle=True)
logger = pl.loggers.TensorBoardLogger('tb_logs', name='segmentation_log_right_blurred_v2_'+version+'_'+str(idx_num), version="right_idx_"+str(idx_num))
checkpoint_callback = ModelCheckpoint(
        dirpath=directory,  # Directory to save the checkpoints
        filename='{epoch:02d}-{val_loss:.2f}',  # Filename template
        save_top_k=-1,  # Save all checkpoints
        save_weights_only=False,  # Save the full model
        every_n_epochs=1  # Save every epoch
    )
mixed_precision_policy = MixedPrecision(
    param_dtype=torch.float16,
    reduce_dtype=torch.float16,
    buffer_dtype=torch.float16
)

fsdp_strategy = FSDPStrategy(accelerator='cuda', 
                             cpu_offload=True, 
                             auto_wrap_policy=None, 
                             mixed_precision=mixed_precision_policy) #offloads model to cpu when able, there is no made wrap policy (manually done in model), and train model using 16 bit floats
Trainer = pl.Trainer(precision=32,
                     default_root_dir="segmentation_models/", 
                     strategy = fsdp_strategy,
                     devices=[0,1], 
                     logger=logger, 
                     callbacks=[checkpoint_callback])

val_dataset = MyDataset_model_num_validation(input_val, output_val, idx_num, bregmas_val, names_val, batch_size=1024)
val_dataloader = DataLoader(val_dataset, batch_size=1, num_workers=1)
Trainer.fit(r_model, train_dataloader, val_dataloader)

解决方案

一、FSDP未自动分片的原因

  1. 手动封装逻辑错误:configure_sharded_model中通过self.trainer.model访问模块是错误的,此时模型未完成FSDP初始化,无法正确触发分片。
  2. 精度配置冲突:Trainer设置precision=32,但FSDP启用了float16混合精度,两者冲突导致混合精度未生效,显存占用翻倍。
  3. 未wrap全部大模块:仅wrap了transformer系列模块,conv1/conv2等显存占用高的模块未被分片,全部堆积在GPU0。
  4. CPU Offload干扰:cpu_offload=True会将未分片模块移到CPU,但手动wrap的模块未被正确识别,导致大模块仍留在GPU0。

二、修复FSDP自动分片

1. 修正configure_sharded_model方法

直接操作当前实例的模块属性,避免通过self.trainer.model访问:

def configure_sharded_model(self):
    self.transformers1 = wrap(self.transformers1)
    self.transformer_0 = wrap(self.transformer_0)
    self.transformer_1 = wrap(self.transformer_1)
    self.transformers_out = wrap(self.transformers_out)
    self.conv1 = wrap(self.conv1)
    self.conv2 = wrap(self.conv2)

2. 统一精度配置

将Trainer的precision改为"16-mixed",和FSDP的混合策略匹配:

fsdp_strategy = FSDPStrategy(
    accelerator='cuda', 
    # 先关闭CPU Offload,确认分片生效后再考虑启用
    auto_wrap_policy=None, 
    mixed_precision=mixed_precision_policy
)

Trainer = pl.Trainer(
    precision="16-mixed",
    default_root_dir="segmentation_models/", 
    strategy=fsdp_strategy,
    devices=[0,1], 
    logger=logger, 
    callbacks=[checkpoint_callback]
)

三、手动分配模型到不同GPU

若需手动控制模块部署(不依赖FSDP自动分片),可采用以下方式:

1. 模块级设备绑定

在setup方法中直接将模块分配到指定GPU:

def setup(self, stage=None):
    # 前半部分模块部署到cuda:0
    self.pos_embed.to('cuda:0')
    self.transformers1.to('cuda:0')
    self.conv1.to('cuda:0')
    self.bn1.to('cuda:0')
    self.transformer_0.to('cuda:0')
    
    # 后半部分模块部署到cuda:1
    self.transformer_1.to('cuda:1')
    self.transformers_out.to('cuda:1')
    self.conv_out.to('cuda:1')

2. 前向传播手动转移张量

在forward中切换张量设备,适配不同GPU上的模块:

def forward(self, parts):
    x, indexes_a, batch_index = parts
    # 前半部分在cuda:0执行
    x = x.to('cuda:0')
    x = torch.squeeze(self.pos_embed(x.transpose(1,2)).transpose(1,2))
    x = x[batch_index]
    x = self.conv1(x)
    x = self.bn1(x)
    x = self.transformers1(x)
    x = x.to(torch.float32)
    x = self.relu(x)
    x = self.to_patches(x)
    N, P, C, X, Y = x.shape
    x = x.view(N*P, C, X, Y)
    x = self.conv2(x)
    x = self.bn2(x)
    x = self.relu(x)
    x = self.transformer_0(x)
    
    # 转移到cuda:1执行后半部分
    x = x.to('cuda:1')
    x = self.transformer_1(x)
    x = self.bn3(x)
    x = self.original_model.from_patched(x).unsqueeze(0)
    x = self.transformers_out(x)
    x = x.to(torch.float32)
    x = self.conv_out(x)
    x = self
相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.23 12:02:18