最近在调一个语义分割模型,单卡跑mIoU能到68左右,换到DistributedDataParallel开4卡训练,同样的数据和超参,结果掉到61,而且训练过程loss下降明显变慢。我确认了batch size是等效放大的(每卡8,总batch 32),学习率也按线性缩放调了,但就是复现不了单卡效果。查了下怀疑是BN的统计量在多卡间同步的问题,但我用的是PyTorch默认的BN,不是SyncBN,按理说每卡独立算BN应该没问题?还是说DDP下BN的running mean/var更新方式和单卡本来就不一样?有没有大佬遇到过类似情况,或者能指点一下排查方向?谢谢!
PyTorch多卡训练时BN层和单卡结果差很多,是不是我哪里搞错了?
全部回复
共 113 条这个现象太典型了,我当初也被坑过一轮。你提到“每卡独立算BN应该没问题”,这里其实有个隐含的坑:DDP默认会在每个step结束时做梯度同步,但BN的running mean/var是各自在前向过程中更新的,这跟单卡在同一个step里看到完整batch的统计量完全是两码事。单卡batch=32,BN统计的是32张图的分布,而DDP每卡只算8张,虽然梯度平均了,但BN的滑动平均相当于是在4个不同的小batch上分别更新,再各自维护一份参数,这会让训练后期统计量变得很不稳定,尤其语义分割这种对空间分布敏感的任务,掉几个点太正常了。
我当时的排查路径是:先加SyncBN试试,如果效果回升,基本就实锤了。不过SyncBN会多不少通信开销,尤其卡多的时候。还有个替代方案,你可以在DDP里把BN的momentum调大一点,比如从0.1调到0.3甚至0.5,让running stats更偏向近期数据,有时候能缓解不同卡上统计量漂移的问题。另外检查一下你的DDP是不是包了模型之后再用的,如果模型里有任何自定义的前向逻辑,或者用了frozen BN层,那个同步行为又会变。
最后我建议你做个对照实验:单卡跑的时候,把batch size也改成8,看看mIoU是不是也会掉到61左右。如果掉,说明问题根本不在DDP,而是单一batch的统计量本身就不够稳;如果单卡8还能保持68,那才是DDP的锅。这步能帮你快速定位方向,别急着调SyncBN。
大概率不是BN同步的问题,DDP默认就是每卡独立算BN,running mean/var也是各自更新,等效于把数据切成4份分别训练,所以统计量会有偏差。你batch size虽然等效32,但每卡8的BN统计量方差比单卡32大不少,尤其语义分割这种类别不均衡的任务,影响会更明显。建议先试试单卡batch size直接设32对比一下,如果也有掉点那就是数据分布问题;要是单卡32没问题,再考虑换SyncBN或者调高BN的momentum,比如从0.1改成0.3,让统计量更新更快跟上。另外确认下DDP里shuffle和seed是否一致,数据顺序变了也可能导致收敛路径不同。
DDP默认BN确实是每卡独立算的,running mean/var也是各自更新,但问题在于你总batch从8变成32后,每卡看到的样本分布其实没变,可BN的统计量在每卡上是用8张图算的,这跟单卡用32张图算出来的分布差异挺大,尤其语义分割这种像素级任务对BN统计量很敏感。你可以试试把BN换成SyncBN,或者干脆把每卡batch size调大点,比如单卡8改成单卡16但总batch保持32,看mIoU能不能回来。另外确认下DDP里BN的momentum是不是默认值,有时候多卡下这个值需要调小一点。
巧了,我之前跑检测也踩过类似的坑,最后发现是DDP里BN的running stats更新时机跟单卡不完全一致导致的,虽然每卡独立算,但梯度all-reduce之后BN的统计量更新会受卡间数据分布差异影响。你可以先试试把BN换成SyncBN对比一下,如果差距缩小那基本就实锤了。另外检查下DDP的broadcast_buffers参数,默认True会把主卡的running mean/var广播到其他卡,这可能会干扰每卡自己的统计更新。还有个排查方向是看下不同卡上的数据分布是不是差太多,比如类别不平衡严重的话,单卡batch的统计量方差会很大,等效batch size虽然大了但每卡看到的样本还是局部的。我之前是把BN的momentum调大一点(比如0.1到0.2),让统计量更快适应新分布,虽然不完美但能缓解一些。
我之前调检测模型也踩过这个坑,单卡和DDP的BN行为确实不完全一样。你用的虽然是默认BN,但DDP下每张卡只看到自己的子batch,running mean和var是在各自卡上独立更新的,然后梯度同步时这些统计量并不会跟着梯度做all-reduce,所以随着训练推进,四张卡的BN统计量会逐渐漂移,跟单卡看到的全局分布越差越远。尤其是你每卡batch只有8,这在小batch下BN本身就很不稳定,统计量噪声会更大,直接反映在loss和精度上。我后来试了两个方向:一是把BN换成SyncBN,让统计量全局同步,效果基本能拉回单卡水平,但会牺牲一点训练速度;二是不换BN,但把每卡batch加大到16以上,同时把BN的momentum调小一点,比如默认0.1改成0.01,让统计量更新更平滑,也能缓解差距。另外你确认一下DDP的broadcast_buffers参数,默认是True,会在每次前向传播前把rank 0的buffer广播给其他卡,这个其实会强制统一BN的running stats,但如果你没关掉它,理论上不该差这么多,可以打印一下checkpoint里的running_mean看看四卡是否一致。还有一个容易被忽略的点是数据加载顺序,DDP下每张卡的sampler是分片的,如果你用了shuffle,不同卡看到的数据分布可能差异很大,这在语义分割这种类别不均衡的数据集上特别致命,建议先固定seed并检查一下每卡的数据分布是否接近。总之先别急着怀疑代码逻辑,把BN的同步机制和数据分片都排查一遍,大概率是这两个因素叠加的结果。
默认BN在DDP下每卡独立算统计量,等效于把全局batch切成了4份,分布估计偏差大很正常,可以试试调低lr或增大warmup。
我遇到过类似问题,后来把BN换成SyncBN就稳了,虽然慢点但指标能对齐单卡。
我之前也踩过类似的坑,而且比你更隐蔽。DDP下每个卡确实独立算BN的统计量,但running mean/var的更新是发生在每张卡自己的forward里的,然后梯度同步的时候并不会去同步这两个buffer,所以本质上你跑4卡等于在拿4个“不同世界”的batch norm在训练,虽然梯度是平均的,但每张卡看到的输入分布已经被自己的BN统计量扭曲了。你单卡batch size是8对吧,这个数字对很多语义分割模型来说本来就偏小,BN在batch size=8时统计量噪声很大,4卡独立算相当于每步用了4个不同的噪声估计去更新同一个共享的权重,模型自然会被拉扯得训练不稳。
我建议你先做个小实验:把单卡batch size也改成32(如果显存够),看看mIoU是不是也会掉到61附近,如果掉,那就说明根本不是DDP的锅,纯粹是你的模型对batch size敏感,小batch的BN统计量本来就不准,等效放大总batch但每卡BN独立,其实等于你同时用了4个不同的小batch统计量,比单卡32的BN更差。如果单卡32能回到68,那就是DDP下BN更新方式的问题,你可以试试把BN换成SyncBN,虽然慢一点,但统计量是全卡同步的,等效于一个大batch的BN,通常能解决这种掉点。另外也检查一下你数据加载的shuffle,多卡下每个epoch的样本分配顺序变了,模型见过的数据组合完全不同,有时候单纯是运气问题,多跑几个seed看看方差。最后提醒一下,DDP里每个卡要设置不同的随机种子,不然数据增强和dropout在每张卡上完全一样,那等效batch size其实没变,也会影响BN行为。
遇到过,DDP下每卡独立算BN确实和单卡不完全等价,问题大概率出在running mean/var的更新频率上。单卡是每步用全batch统计量更新,而DDP里每卡只看到自己那8张图,更新方向会更嘈杂,尤其batch不大时统计量漂移明显。建议先试试把每卡batch size调大,或者干脆换SyncBN对比一下,如果SyncBN能回到68左右就实锤了。另外你学习率线性缩放用的是全局batch还是单卡batch?有时候这里差个系数也会让loss曲线变化。
每卡batch太小了,8的BN统计量噪声大,试试开SyncBN或者把单卡batch调大对比下。
你这个现象还挺典型的,多卡掉点不一定全是BN的锅。DDP下每卡BN确实是独立统计的,但等效batch从8变成32后,BN的running stats更新频率和单卡不一样,而且每卡看到的样本分布差异会被放大。另外学习率线性缩放只是经验法则,warmup和weight decay的缩放也得跟着调,不然前期很容易训崩。建议先把4卡的总batch设成和单卡一致(每卡2)跑一遍,确认是不是batch相关,再考虑换SyncBN对比。
DDP下BN还是各卡独立算的,问题多半出在等效学习率没调对,试试warmup或者把lr再降一点看看。
BN统计量每卡独立算,但等效batch变了,实际每卡样本少,running stats本身就飘。
DDP下BN的running stats确实是各卡独立更新的,但问题在于每卡只看到自己的8个样本,统计量噪声比单卡batch 32大不少,尤其分割任务小batch本身就不友好。我之前也踩过类似的坑,换成SyncBN之后mIoU基本能拉回来,你可以先试试这个。另外DDP的梯度平均和BN的momentum交互也会影响收敛节奏,建议把momentum调小一点观察下。如果还不行就看看数据shuffle有没有问题,有时候sampler的seed没对齐也会导致这种诡异差距。