最近在调一个语义分割模型(DeepLabV3+,backbone是ResNet50),单卡batch size只能开到6,想着用两张3090试试多卡。先试了nn.DataParallel,发现第一张卡显存占用直接飙到23G,第二张卡只有15G左右,感觉负载完全没均衡。后来换了DistributedDataParallel,虽然显存均衡了一些,但遇到一个问题:我代码里有几个BN层,DDP默认会同步BN,结果训练速度反而变慢了。想问下各位,是不是我哪里设置有问题?DDP做同步BN是不是一定比DP快?还有没有其他因为多卡引起的坑(比如数据加载、loss计算)需要注意的?感谢!
PyTorch多卡训练时显存不均衡怎么办?DataParallel还是Distributed?
全部回复
共 11 条DP的显存不均衡太正常了,因为主卡要额外存output和梯度,你那个23G vs 15G基本就是这原因,换DDP方向是对的。不过同步BN那个坑我也踩过,如果batch size本身不大,其实可以试试把sync_bn关掉,或者用torch.nn.SyncBatchNorm.convert_sync_batchnorm手动转一下,速度能回来不少。另外提醒下,DDP里loss要记得做all_reduce平均,不然每个卡算的梯度尺度不一样,收敛会出问题。数据加载那边,多卡时num_workers最好按卡数乘2配,不然有时候数据供不上,训练反而更慢。
DDP显存均衡是正常的,因为DP本身就是把主卡当汇聚节点用,负载不均那是设计缺陷。你BN同步慢大概率是没设dist.barrier或者没把sync_bn=True配好,其实3090单卡batch=6的话,两张卡加起来12,BN统计量差异也不大,完全可以不同步。另外数据加载那边记得把num_workers拉高,不然多卡反而瓶颈在CPU,我踩过这个坑。
DDP那个显存均衡是正常的,因为DP本身就是把输入按batch切到每张卡上,但forward和backward都在主卡做梯度汇总,主卡还得额外存一份全量参数和优化器状态,所以显存天然就比从卡高一截,你这差的8G基本就是模型参数加BN统计量。DDP同步BN确实会慢,因为每步训练都要等所有卡的BN统计量做全局同步,这通信开销在两张卡上可能比计算还耗时,尤其是小batch训练时特别明显。你要是图像分割这种任务,不同卡上的特征分布其实差别不大,完全可以把sync_bn关掉,直接设ddp里find_unused_parameters=False,然后每个进程自己维护BN,速度能提回来不少。还有个坑是loss计算,如果你用的是DP,每张卡的loss是独立算的,但DDP里如果loss里有跨卡tensor操作,比如把logits全收集起来再算,那梯度会重复累加,得手动除一下卡数。数据加载那边记得每个进程设不同的sampler,不然每张卡喂的数据都一样,等于白训,还有num_workers最好设成卡数乘2,不然数据准备跟不上。另外3090的话建议开amp混合精度,显存能再省一半,batch能开到12以上,说不定单卡就够用了。
DDP的同步BN确实拖慢速度,非必要就关掉,负载均衡比BN同步重要多了。
DDP显存均衡是正常的,但同步BN慢的话,你检查下是不是用了torch.nn.SyncBatchNorm转换,其实两张卡没必要同步BN,直接关掉就行,把batch size调大点效果差不多。另外DataParallel那个负载不均是因为它每轮都要把loss和梯度汇总到主卡,本质设计就这样,别纠结了。还有个小坑是DDP要确保每个进程的dataloader shuffle时用不同seed,不然数据重复,你可以在每个rank里给sampler设个固定seed。
DP就这样,负载不均正常,DDP换来的均衡更划算。BN同步慢的话试试关掉syncBN,数据量够大影响不大。
DDP的显存均衡其实是因为它本身就按rank切数据,DP那种前向复制backward归约的模式注定了第一卡要做额外的loss聚合和梯度汇总,所以显存不均很正常,你换几卡都一样。至于BN同步变慢,那是必然的,因为SyncBN要跨卡通信统计量,对比DP每卡自己算BN,开销确实大,你要是显存没吃满或者batch size够大,其实没必要开SyncBN,DDP默认不开启的,你是不是代码里显式设了sync_bn=True?我一般做法是直接在DDP里保持每个卡独立BN,等训练稳定后再用SyncBN微调几轮,效果差不多但速度能快不少。另外你说的坑,我踩过的是数据加载那一块,DDP每个进程要单独设sampler,不然每个epoch都会重复采样同样的数据,还有loss如果涉及跨卡交互,比如对比损失或者全局池化,记得要自己做all_gather,不然算出来的梯度是错的。还有个小建议,3090是24G,单卡batch6的话两卡理论上能到12,但如果你模型里有大的中间特征,记得开cudnn.benchmark=True和torch.cuda.amp混合精度,可能单卡就能开到10以上,根本不用多卡折腾。你那个15G和23G的差距,其实也可以用torch.utils.data.distributed.DistributedSampler配合shuffle=True时设drop_last=True,不然最后一批不完整会让某张卡多算一点。
其实DDP的同步BN不一定要全局开,如果你的batch size已经不小了,可以试试关掉syncBN,让每张卡自己算BN统计量,速度能提上来不少。另外DataParallel那个负载不均本质是它的机制问题,它每step都要把主卡算好的梯度广播回去,通信开销也大,所以除非模型特别小,不然真不建议用了。你那个显存差8G左右其实挺典型的,DDP下还可以检查下是不是输入数据没按卡数均匀切分,有时候最后一个batch不均匀也会导致某张卡多算一点。顺便说个坑,loss如果是在主卡上算的,记得用reduce=True或者手动做all_reduce,不然梯度没问题但loss曲线会看起来很奇怪。
DDP同步BN确实会拖速度,试试换成SyncBatchNorm但只加在需要同步的层上,数据加载记得设num_workers跟pin_memory。
DP显存不均衡是经典问题了,第一张卡要负责汇总所有卡的梯度再回传,多占的那部分基本就是output和梯度聚合的开销,跟batch size关系不大,所以23G和15G这个差距挺正常的。换DDP方向肯定是对的,DP本身就是单进程多线程,GIL加上通信效率低,卡越多越吃亏。至于同步BN拖慢速度,这个得看你的瓶颈在哪,SyncBN跨卡通信确实有开销,但如果单卡batch只有6,BN统计量本身噪声就大,同步之后精度往往更稳,速度慢一点换来收敛稳定我觉得是值得的,实在嫌慢可以试试换GroupNorm或者FrozenBN。DDP不一定绝对比DP快,卡少、模型小的时候差距不明显,但卡多之后DDP优势才出来。另外你说的数据加载这块,记得给每个进程设好DistributedSampler,不然会重复采样;loss那块如果用了带ignore_index的交叉熵,DDP下梯度是自动平均的,别自己再除一遍卡数。还有个容易踩的坑是模型里如果有随机的数据增强或者dropout,DDP每个进程的随机种子记得错开,不然等于变相缩小了batch的多样性。
DP第一张卡显存高是因为它要收集所有卡的输出算loss,主卡天然吃亏,换DDP后每张卡各自算loss就均衡了。同步BN变慢挺正常,跨卡通信有开销,卡少或者BN层不多的话可以试试SyncBN关掉,或者把BN换成GroupNorm,分割任务里GN效果也不差。另外DDP记得把DataLoader的sampler设成DistributedSampler,不然每张卡读重复数据等于白搭,loss那块也要注意别在主卡上单独reduce。