PyTorch中DataParallel多GPU内存不均衡致OOM问题的解决方法咨询
解决PyTorch DataParallel多GPU内存不均衡问题(NMT场景)
遇到DataParallel下单GPU内存爆掉、其他GPU闲置的情况太常见了,尤其是NMT这种序列模型,本身就对内存敏感。结合你的场景,我整理了几个实用的解决思路:
一、调整DataParallel的主设备分配
DataParallel默认把主设备设为cuda:0,它不仅要存模型副本,还要负责数据分发、结果收集,额外开销比其他GPU大很多。你可以手动指定主设备到其他GPU,分摊压力:
# 先把模型移到非0GPU(比如cuda:1) model = model.to('cuda:1') # 指定device_ids包含所有要用到的GPU,output_device设为刚才的主设备 model = nn.DataParallel(model, device_ids=[1, 0, 2], output_device=1)
这样原来的cuda:0就不用承担主设备的额外工作,内存占用会明显降下来。
二、优化数据的batch分配(针对NMT序列特性)
NMT的样本序列长度差异大,DataParallel只是简单按样本数均分batch,但如果某块GPU分到的样本全是长序列,总token数远超其他GPU,内存自然会爆。你可以:
- 在DataLoader的
collate_fn里做动态padding分组:先按序列长度把样本分组,让每个batch内的样本长度尽量接近,再分配给GPU,这样每个GPU的总token数更均衡。 - 不要只看样本数量的batch size,改用token数限制的batch size,比如每个GPU的batch总token数不超过固定值,避免长序列占满内存。
三、降低主设备的额外内存开销
- 把不需要的变量(比如验证集数据、日志记录、中间计算的非张量数据)都放到CPU上,别占GPU内存。
- 定期在训练循环的合适时机调用
torch.cuda.empty_cache()清理无用显存,但别太频繁(比如每几个epoch一次),不然会拖慢训练速度。
四、启用半精度训练(大幅减内存)
NMT模型用半精度(FP16)训练几乎不会影响最终精度,但能把内存占用砍半左右,这是最立竿见影的方法:
from torch.cuda.amp import GradScaler, autocast scaler = GradScaler() optimizer = torch.optim.Adam(model.parameters()) for epoch in range(epochs): for src, tgt in train_dataloader: src, tgt = src.to('cuda'), tgt.to('cuda') optimizer.zero_grad() # 启用自动混合精度 with autocast(): outputs = model(src, tgt[:, :-1]) loss = criterion(outputs.reshape(-1, outputs.size(-1)), tgt[:, 1:].reshape(-1)) # 用scaler处理梯度,避免半精度下的梯度溢出 scaler.scale(loss).backward() scaler.step(optimizer) scaler.update()
五、切换到DistributedDataParallel(推荐长期方案)
DataParallel的设计天生就有主设备瓶颈,而**DistributedDataParallel(DDP)**是每个GPU都有独立的模型副本和数据采样,完全并行,内存分配会均匀很多,训练速度也更快。针对NMT的基本配置示例:
import os import torch import torch.distributed as dist from torch.nn.parallel import DistributedDataParallel as DDP from torch.utils.data import DistributedSampler def setup(): dist.init_process_group("nccl") def cleanup(): dist.destroy_process_group() def main(): setup() local_rank = int(os.environ["LOCAL_RANK"]) torch.cuda.set_device(local_rank) # 初始化模型并移到当前GPU model = YourNMTModel().to(local_rank) model = DDP(model, device_ids=[local_rank]) # 用DistributedSampler分配数据,保证每个GPU拿到不重复的样本 train_dataset = YourNMTDataset() train_sampler = DistributedSampler(train_dataset) train_dataloader = torch.utils.data.DataLoader( train_dataset, batch_size=your_batch_size, sampler=train_sampler ) # 训练循环和正常流程一致,注意用sampler.set_epoch(epoch)保证每个epoch数据打乱 optimizer = torch.optim.Adam(model.parameters()) for epoch in range(epochs): train_sampler.set_epoch(epoch) for src, tgt in train_dataloader: src, tgt = src.to(local_rank), tgt.to(local_rank) optimizer.zero_grad() outputs = model(src, tgt[:, :-1]) loss = criterion(outputs.reshape(-1, outputs.size(-1)), tgt[:, 1:].reshape(-1)) loss.backward() optimizer.step() cleanup() if __name__ == "__main__": main()
启动时用torchrun命令:
torchrun --nproc_per_node=3 your_train_script.py
先试试前面几个轻量的调整,如果还是有问题,直接切DDP会彻底解决内存不均衡的问题。
内容的提问来源于stack exchange,提问作者Wasi Ahmad
相关产品推荐
相关产品推荐

