8张RTX3080预训练BART遇CUDA OOM,疑似未启用多GPU分布式训练求解
解决BART预训练CUDA内存不足及分布式训练失效问题
1. 先确认分布式训练是否真的生效
在你的pretrain_bart.py开头加几行验证代码,跑一次脚本看输出:
import torch.distributed as dist import torch # 初始化分布式环境 dist.init_process_group(backend='nccl') rank = dist.get_rank() world_size = dist.get_world_size() print(f"当前进程rank: {rank}, 总进程数: {world_size}") torch.cuda.set_device(rank) print(f"当前进程绑定GPU: {torch.cuda.current_device()}")
如果8个进程分别输出rank 0到7,说明分布式启动没问题;要是所有进程rank都是0,那就是启动配置出了问题。
2. 修正分布式启动脚本
- 优先用
torchrun替代torch.distributed.launch(PyTorch 1.10+推荐,自动处理节点配置,减少手动出错):
把启动命令改成:torchrun --nproc_per_node=8 pretrain_bart.py \ --num-layers 12 \ --hidden-size 1024 \ --num-attention-heads 16 \ --micro-batch-size 1 \ --global-batch-size 8 - 要是坚持用
torch.distributed.launch,必须在pretrain_bart.py里接收--local_rank参数:
在参数解析部分加:
然后用这个参数绑定GPU并初始化分布式:import argparse parser = argparse.ArgumentParser() parser.add_argument('--local_rank', type=int, default=-1) args = parser.parse_args()torch.cuda.set_device(args.local_rank) dist.init_process_group(backend='nccl')
3. 模型与数据的分布式适配
- 模型必须用DDP包装:别用旧的
DataParallel,换成DistributedDataParallel(DDP),确保每个GPU只加载模型的一份副本:from torch.nn.parallel import DistributedDataParallel as DDP model = BartForPretraining(config).to(args.local_rank) model = DDP(model, device_ids=[args.local_rank]) - 数据必须用DistributedSampler拆分:保证每个进程只处理部分训练数据,避免重复加载全量数据:
训练时每个epoch开始前要调用from torch.utils.data import DistributedSampler train_dataset = YourCustomDataset(...) train_sampler = DistributedSampler(train_dataset) train_loader = DataLoader( train_dataset, batch_size=args.micro_batch_size, sampler=train_sampler, num_workers=4 )train_sampler.set_epoch(epoch),保证数据打乱的随机性。
4. 内存优化补充
- 开启自动混合精度训练,能直接砍掉近一半内存占用:
from torch.cuda.amp import autocast, GradScaler scaler = GradScaler() for batch in train_loader: with autocast(): outputs = model(**batch) loss = outputs.loss scaler.scale(loss).backward() scaler.step(optimizer) scaler.update() - 加环境变量优化内存碎片:
在启动脚本前加上:
减少小内存块碎片化,避免明明有剩余内存却无法分配的情况。export PYTORCH_CUDA_ALLOC_CONF=max_split_size_mb:128
内容的提问来源于stack exchange,提问作者messi
相关产品推荐
相关产品推荐

