最近尝试用PyTorch 2.0的DDP(DistributedDataParallel)在4卡V100上微调一个13B的LLaMA模型,batch size每卡设了2,梯度累积8步。之前单卡跑小模型(1.3B)loss曲线挺平滑的,但一上多卡,loss就开始疯狂震荡,尤其是前几百步,偶尔还会出现loss突然跳高然后又降下来的情况。我试过调低学习率从1e-5到3e-6,也加了warmup,但效果不明显。怀疑是不是torch.compile或者梯度同步的问题?还是说13B模型本来就需要更大的batch size才能稳?有没有踩过类似坑的朋友分享下经验?先谢谢了!
PyTorch 2.0跑大模型,DDP训练loss震荡严重,求大佬指点
全部回复
共 119 条我之前调7B的时候也遇到过类似的,前几百步loss跟过山车似的,后来发现是DDP的梯度all-reduce和batch norm的统计量在搞鬼,你可以先试试关掉torch.compile看看是不是编译优化带来的数值抖动。另外13B加4卡这个配置等效batch size才64,对LLaMA来说确实偏小,建议把梯度累积提到16步或者直接上更大batch,loss会稳不少。还有个小细节,检查下dataloader的shuffle和seed,多卡下数据分配不均匀也会导致这种跳变。
我之前也遇到过类似情况,13B本身就比1.3B敏感得多,4卡V100上每卡2的batch其实等效才8,对这么大模型来说确实偏小,loss震荡挺正常的。你可以试试把梯度累积提到16或者干脆用gradient checkpointing换更大batch,先排除这个因素。另外torch.compile在DDP下偶尔会跟梯度同步有交互问题,建议先关掉compile跑几百步对比一下,如果明显稳了就是它的问题。还有个细节,warmup步数别太少,13B模型我一般warmup到总步数的5%以上才有效果。
试试把梯度累积改成跨卡同步,或者先关掉compile验证下,13B这规模lr确实得再降点。
之前单卡能稳说明模型本身没问题,多卡loss震荡大概率是梯度同步加累积步数叠加导致的,试试关掉torch.compile再看下。
我上周刚用DDP调过7B的模型,也遇到过loss跳变,后来发现是梯度累积和DDP的bucket通信没对齐导致的,你可以试试把gradient_accumulation_steps放到DDP外面手动做,别用amp和compile同时开。另外13B在4卡上每卡2的batch确实偏小,等效batch才64,对13B来说可能真不够稳,建议先试个500步看看warmup结束后是否回落,如果还震再考虑加大batch或换ZeRO。
我之前在2.0上跑DDP也遇到过类似情况,尤其是用了torch.compile之后更明显。感觉你怀疑的方向挺对的,但问题可能不在梯度同步本身,而是compiler在分布式下的图优化跟DDP的梯度all-reduce时序有交互,导致数值行为跟单卡不完全一致。建议先做个A/B测试,把torch.compile关掉跑同样配置,如果loss震荡明显减轻,那就是编译优化的问题,可以直接不用它或者改用mode=reduce-overhead试试。另外13B配4卡每卡batch 2确实太小了,有效batch才64,对这么大模型来说噪声很大,哪怕梯度累积8步也没法完全抵消数据分布差异。我实际经验是至少每卡batch 4、累积16步,也就是有效batch 256左右,loss曲线才明显稳下来。还有V100上跑13B可能遇到精度问题,试试bf16或者fp16加上动态loss scaling,有时候震荡是数值溢出导致的。最后注意一下dataloader的shuffle,多卡下每个rank的数据分布如果不够随机,也会造成前几百步的剧烈波动,可以试试设个固定seed或者换更均匀的采样器。
我之前也遇到过类似情况,DDP下loss震荡很多时候不是lr的锅,而是梯度同步时不同卡上数据分布差异被放大了,尤其LLaMA这种大模型对batch的随机性更敏感。你可以试试把梯度累积改成在每张卡内先累积完再同步,或者用DistributedSampler的shuffle设置保证每个step数据充分打乱。另一个思路是排查torch.compile,2.0刚出时有些算子在DDP下会产生数值抖动,可以先关掉compile跑几百步对比一下。如果还不行,建议把batch size提到4或8,哪怕是梯度累积,实际有效batch变大后loss会稳很多。
之前踩过一模一样的坑,13B在4卡上确实容易这样。我后来发现是DDP的bucket size默认值对超大模型太不友好了,梯度更新时各卡之间会有微小的不同步,导致loss跳变。你可以试试调大bucket_cap_mb到100甚至200,或者用find_unused_parameters=False显式声明所有参数都会用到,能减少不少同步开销。另外你确认一下是不是用了fp16?loss震荡有时候是混合精度下的梯度缩放因子在DDP里没调好,建议把static_loss_scale固定试试。
我之前微调7B也遇到过类似情况,后来发现是梯度累积和DDP的allreduce交互时,等效batch size在卡间同步上出了偏差,建议先关掉torch.compile试试,它跟DDP在某些版本下有兼容性问题。另外13B用4卡的话,每卡batch=2加8步累积,全局batch其实才64,对这么大模型确实偏小,loss震荡很正常,你可以试着把梯度累积加到16或者直接上gradient checkpointing换更大batch。还有个细节,不同卡的随机种子或数据shuffle顺序不一致也会导致前期震荡,检查下DataLoader的sampler是不是设了drop_last=True。
之前跑7B时也遇到过类似情况,后来发现是DDP的梯度all-reduce和gradient accumulation顺序有冲突。你累积8步的话,得确认一下是不是在每步accumulation里都做了梯度同步,最好只在最后一步同步。另外torch.compile对动态shape敏感,建议先关掉对比试试,有时候它和DDP的bucketize配合会引入额外抖动。
13B这个规模确实对全局batch size更敏感,4卡×2×8=64的等效batch对llama来说偏小,可以试试把gradient accumulation提到16步,或者直接用zero1把优化器状态切分一下,能显著缓解loss毛刺。如果还不行,检查一下数据加载的shuffle种子,多卡下每个rank的数据分布不均匀也会导致这种震荡。
这现象我之前在调多卡的时候也遇到过,尤其是从单卡切到DDP后,前几百步loss跳高其实不一定是优化器的问题,很可能是不同卡上的数据分布差异被放大了。你可以试试把每个batch里的数据shuffle时固定一下随机种子,或者干脆先关掉torch.compile跑几百步对比下,我之前发现compile的图优化在某些场景下会和DDP的梯度同步有点微妙冲突。另外13B每卡batch size 2确实偏小,梯度累积虽然等效增大了batch,但BN或者LayerNorm的统计量更新还是按实际batch算的,可能也是震荡来源之一,可以考虑先不累积,直接拉大每卡batch size看看。
我之前在2.0上踩过类似的坑,主要怀疑是torch.compile和DDP的梯度同步在动态shape下会有隐性bug,你可以先关掉compile试试,纯DDP跑几百步对比下loss曲线。另外13B配4卡确实有点吃紧,等效batch才64,对这么大模型来说偏小,建议把梯度累积提到16步,或者干脆用bfloat16混合精度,V100虽然不支持bf16但可以用fp16加梯度缩放试试。还有个小细节,warmup步数别太短,至少占总数5%以上,不然前期优化器状态没稳定确实容易跳变。
我之前在4卡上跑7B也遇到过一模一样的现象,后来排查发现torch.compile和DDP的梯度桶同步在某些版本下会互相干扰,尤其是当模型里有动态shape的操作时。你可以先试试把torch.compile关掉,只用纯DDP跑几百步看看loss曲线,如果稳了那就是编译优化的问题。另外13B模型每卡batch size只有2确实太小了,梯度累积8步虽然等效batch size是64,但BN(如果有)或者某些layer的统计量在多卡间是不同步的,这会放大震荡。还有个坑是warmup步数,13B模型前几百步loss跳高很可能是因为warmup太短导致lr还没降到合理范围,建议把warmup拉到总步数的5%甚至10%。我后来是把学习率改成余弦退火,并关掉compile,同时把梯度累积改成先本地累积再同步,震荡就明显缓解了。不过你用的V100不支持bf16,如果是fp16混合精度,还得检查一下loss scaling是否在DDP下正确更新,有时候梯度裁剪和scaler的顺序不对也会造成突然跳变。
试试把梯度累积改成跨卡同步,或者先关掉compile跑几步对比下,13B这规模确实容易这样。
我之前在2.0上踩过类似的坑,感觉torch.compile跟DDP一起用的时候,特别是在小batch下,容易出现你说的这种loss陡升陡降。可以先试试把compile关掉,纯用DDP跑几百步对比一下,如果稳定了那就是编译图优化和梯度同步之间的时序问题。另外13B模型在4卡上每卡batch=2,算上梯度累积等效batch才64,对这么大模型确实偏小,loss震荡不一定是bug,更可能是优化器在前期对噪声太敏感了。你可以试试把梯度累积提到16步,等效batch到128,同时把学习率再往下压一点,比如1e-6级别,看震荡幅度会不会明显收窄。还有个细节是warmup步数别太短,13B模型前几百步loss波动大很正常,至少跑个两三千步warmup,让lr慢慢爬上去。如果你用的是bf16或者fp16,也得留意下loss scaling是不是在DDP下被每个rank单独调整了,偶尔也会造成数值抖动。最后建议把每步的梯度norm打印出来,如果某几步norm突然飙到几十上百,那大概率是某张卡的数据特别“难”,可以检查下数据shuffle是不是每个rank都用了不同的种子。
先确认一下dataloader的shuffle和sampler设对了没,多卡下数据重复会导致loss跳变。
我之前在34B模型上也碰到过一模一样的情况,单卡很稳,一上DDP就跟过山车似的。后来排查下来,问题不在torch.compile,而是梯度累积和all-reduce之间的交互——你想想,8步累积本来是为了模拟大batch,但DDP默认是每步都同步梯度的,累积期间梯度被反复平均,等效batch size其实不是单纯乘8那么简单,前几步的噪声会被放大。我建议你先试试把梯度累积关掉,直接拉大单卡batch size到8或16,看看loss曲线平不平,这能帮你区分是优化器问题还是通信问题。另外13B在4卡V100上,每卡2的batch确实偏小了,虽然理论上有gradient checkpointing撑着,但BN统计(如果你用了LayerNorm之外的任何归一化)和梯度噪声都会因为卡间数据分布不一致而恶化。还有个小坑,PyTorch 2.0的DDP在find_unused_parameters=True时,如果模型里有参数没参与loss计算,会导致梯度稀疏同步,loss跳高那一下很可能是这个。你可以开gradient_as_bucket_view=True,顺便用torch.distributed.barrier()手动对齐一下前几个step。最后,warmup建议从0开始,但步数要拉长到总步数的10%以上,3e-6的学习率对13B其实还是偏高,尤其配合AdamW,你可以试试2e-6加cosine decay,把peak放在200步之后。如果还是震荡,建议开gradient clipping,max_norm设1.0,能压住那种突然的尖峰。
13B在4卡上等效batch才64(2×4×8),对这么大模型确实偏小,loss震荡不奇怪。我之前调7B时也遇到过类似情况,后来把梯度累积提到16步,同时把学习率scheduler的warmup步数拉长到总步数的10%才稳住。另外torch.compile在DDP下偶尔会引入数值抖动,建议先关掉compile跑几百步对比下,排除这个变量。你用的什么优化器?AdamW的话beta2调到0.98对稳定性有帮助。
我之前也遇到过类似情况,13B直接上DDP确实容易这样。一个可能是梯度累积和DDP的all-reduce交互导致梯度噪声变大,建议先关掉torch.compile试试,它有时候会改变算子融合顺序影响数值稳定性。
另外前几百步震荡如果幅度在可控范围,其实可以观察下是不是在逃离初始的尖锐局部最小值,我试过把warmup拉长到总步数的10%以上,同时用grad clip(比如1.0)能压住那种突然跳高的尖峰。batch size的话,4卡x2x8累积等效64,理论上对13B不算太小,但如果你用的是fp16,建议检查下loss scale是否频繁溢出,这也会造成诡异波动。
还有个思路:对比一下单卡(哪怕小batch)跑同样步数的loss曲线,如果单卡也抖,那可能就不是DDP的问题,而是模型本身或数据顺序的锅。
13B配4卡确实吃紧,试下关掉torch.compile,loss震荡多半是它跟DDP的通信竞争闹的。
这问题我熟,之前调7B也遇到过,DDP下loss震荡大概率不是torch.compile的锅,你先试试把梯度累积关了直接上大batch,V100 4卡跑13B的话单卡batch2确实偏小,等效batch才64,对13B来说有点不够看。另外注意下warmup步数,别设太短,我建议至少占到总步数的5%以上,还有检查下不同卡的数据shuffle是否一致,有时候这个也会导致前期震荡。