如何优化SLURM脚本与PyTorch代码以多GPU训练大预训练模型
多GPU多模态情感分析训练优化方案
针对你用3个A100 GPU训练RoBERTa+ResNet/DenseNet多模态情感分析模型时出现的CUDA OutOfMemoryError,以及单GPU训练耗时过长的问题,以下是SLURM脚本和代码的优化方案:
一、SLURM脚本优化
优化核心是确保资源分配匹配GPU需求,消除通信瓶颈:
- 明确指定A100 GPU数量:添加
--gres=gpu:a100:3,避免集群分配非目标GPU类型。 - 绑定足够CPU核心:设置
--cpus-per-task=12(1个A100对应4核,3个共12核),匹配数据加载的CPU需求,减少CPU-GPU传输延迟。 - 加载兼容环境模块:指定CUDA和PyTorch版本,比如
module load cuda/11.7 pytorch/2.0.1,确保NCCL多GPU通信正常工作。 - 规范日志输出:用
--output=train_%j.out --error=train_%j.err记录训练日志,方便后续排查问题。
修改后的SLURM脚本示例:
#!/bin/bash #SBATCH --job-name=multimodal_sentiment #SBATCH --nodes=1 #SBATCH --ntasks=1 #SBATCH --gres=gpu:a100:3 #SBATCH --cpus-per-task=12 #SBATCH --mem=64G #SBATCH --time=24:00:00 #SBATCH --output=train_%j.out #SBATCH --error=train_%j.err # 加载依赖环境 module load cuda/11.7 module load pytorch/2.0.1 # 用torchrun启动DDP训练,指定单节点3个进程 torchrun --nproc_per_node=3 train.py --per_gpu_batch 32 --accumulation_steps 2
二、代码层面优化
1. 解决显存不足问题
- 拆分全局batch到单GPU:不要直接设置全局batch为128,改为每个GPU分配32,3个GPU总batch为96,再通过梯度累积(设置
accumulation_steps=2)等效于全局batch 192,既满足训练需求又降低单GPU显存占用。 - 启用混合精度训练:使用PyTorch的
torch.cuda.amp模块,可减少约50%的显存占用,且几乎不损失精度。 - 分层冻结预训练参数:先冻结RoBERTa、ResNet/DenseNet的底层参数(比如前8层),只训练顶层融合层;待模型收敛后,再微调部分底层参数,大幅降低显存开销。
- 优化数据加载:设置
pin_memory=True和num_workers=4(每个GPU对应4个worker,总12个匹配SLURM的CPU核心数),加速数据传输,避免内存泄漏。 - 定期清理显存:在验证阶段结束后调用
torch.cuda.empty_cache(),释放无用张量占用的显存。
2. 充分利用多GPU资源
放弃低效的DataParallel,改用DistributedDataParallel(DDP),这是PyTorch多GPU训练的最优方案,能避免单GPU瓶颈,提升训练效率。核心代码修改如下:
import os import torch import torch.distributed as dist from torch.nn.parallel import DistributedDataParallel as DDP from torch.utils.data.distributed import DistributedSampler from transformers import RobertaModel from torchvision.models import resnet50 class FusionLayer(torch.nn.Module): # 你的模态融合层实现 def __init__(self, text_dim=768, img_dim=2048, num_classes=2): super().__init__() self.fc = torch.nn.Linear(text_dim + img_dim, num_classes) def forward(self, text_feat, img_feat): concat_feat = torch.cat([text_feat[:, 0, :], img_feat], dim=1) return self.fc(concat_feat) def main(): # 初始化DDP进程组 dist.init_process_group(backend='nccl') local_rank = int(os.environ['LOCAL_RANK']) torch.cuda.set_device(local_rank) # 加载预训练模型并转移到GPU text_model = RobertaModel.from_pretrained('roberta-base').to(local_rank) img_model = resnet50(pretrained=True).to(local_rank) fusion_model = FusionLayer().to(local_rank) # 组合模型并封装为DDP model = torch.nn.ModuleDict({ 'text': text_model, 'img': img_model, 'fusion': fusion_model }) model = DDP(model, device_ids=[local_rank]) # 配置分布式数据加载器 from datasets import MultimodalDataset # 你的自定义数据集类 train_dataset = MultimodalDataset(train_data_path) train_sampler = DistributedSampler(train_dataset) train_loader = torch.utils.data.DataLoader( train_dataset, batch_size=int(os.environ['PER_GPU_BATCH']), sampler=train_sampler, num_workers=4, pin_memory=True ) # 混合精度训练配置 scaler = torch.cuda.amp.GradScaler() optimizer = torch.optim.AdamW(model.parameters(), lr=1e-5) accumulation_steps = int(os.environ['ACCUMULATION_STEPS']) epochs = 10 for epoch in range(epochs): train_sampler.set_epoch(epoch) # 确保每个epoch数据打乱一致 model.train() total_loss = 0.0 for batch_idx, (text_input, img_input, label) in enumerate(train_loader): # 转移数据到当前GPU text_input = {k: v.to(local_rank) for k, v in text_input.items()} img_input = img_input.to(local_rank) label = label.to(local_rank) # 混合精度前向传播 with torch.cuda.amp.autocast(): text_feat = model.module.text(**text_input) img_feat = model.module.img(img_input) logits = model.module.fusion(text_feat, img_feat) loss = torch.nn.CrossEntropyLoss()(logits, label) # 梯度累积 loss = loss / accumulation_steps scaler.scale(loss).backward() # 累积到指定步数后更新参数 if (batch_idx + 1) % accumulation_steps == 0: scaler.step(optimizer) scaler.update() optimizer.zero_grad() total_loss += loss.item() * accumulation_steps if local_rank == 0: # 仅主进程打印日志 print(f"Epoch [{epoch+1}/{epochs}], Loss: {total_loss/len(train_loader):.4f}") # 销毁进程组 dist.destroy_process_group() if __name__ == '__main__': import argparse parser = argparse.ArgumentParser() parser.add_argument('--per_gpu_batch', type=int, default=32) parser.add_argument('--accumulation_steps', type=int, default=2) args = parser.parse_args() os.environ['PER_GPU_BATCH'] = str(args.per_gpu_batch) os.environ['ACCUMULATION_STEPS'] = str(args.accumulation_steps) main()
3. 额外显存优化细节
- 启用梯度裁剪:添加
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0),防止梯度爆炸,同时稳定显存占用。 - 验证阶段切换eval模式:
model.eval()关闭dropout、批量归一化的随机操作,减少显存消耗。 - 关闭预训练模型冗余日志:加载模型时设置
logging_level=logging.WARNING,避免IO占用影响训练速度。
三、显存瓶颈排查
如果仍出现OOM,可在关键节点添加显存监控代码,定位瓶颈阶段:
print(f"Rank {local_rank} - Before forward: {torch.cuda.memory_allocated()/1024**3:.2f} GB") with torch.cuda.amp.autocast(): # 前向传播代码 print(f"Rank {local_rank} - After forward: {torch.cuda.memory_allocated()/1024**3:.2f} GB") loss.backward() print(f"Rank {local_rank} - After backward: {torch.cuda.memory_allocated()/1024**3:.2f} GB")
内容的提问来源于stack exchange,提问作者ririto
相关产品推荐
相关产品推荐

