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未自动分片的原因
- 手动封装逻辑错误:
configure_sharded_model中通过self.trainer.model访问模块是错误的,此时模型未完成FSDP初始化,无法正确触发分片。 - 精度配置冲突:Trainer设置
precision=32,但FSDP启用了float16混合精度,两者冲突导致混合精度未生效,显存占用翻倍。 - 未wrap全部大模块:仅wrap了transformer系列模块,
conv1/conv2等显存占用高的模块未被分片,全部堆积在GPU0。 - 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
相关产品推荐
相关产品推荐

