最近尝试用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 条大概率是梯度累积和DDP的allreduce没配合好,试试把梯度累积改成手写并确保只在最后一步同步。13B这规模对lr本来就敏感,3e-6还抖就再砍半。
之前跑7B也遇到过类似情况,后来发现是梯度累积和DDP的梯度平均叠加了,导致有效batch size比预想的大不少,loss就容易抽风。你可以试试把梯度累积改成用DDP的gradient_as_bucket_view配合no_sync,或者干脆把累积步数降下来,让每步的梯度更纯粹。另外13B在4卡上batch确实偏小,bn这种对batch敏感的结构还好,但transformer的loss对lr和batch的匹配挺挑剔的,建议先固定一个有效batch(比如32),再用线性缩放法则调lr试试。torch.compile我这边开了反而更不稳,可以先关掉对比下。
试试关掉torch.compile,DDP下跟gradient accumulation容易有同步时序问题,我之前也踩过这坑。
13B这规模用4卡本来就偏紧,lr可以再降个量级,或者试试bf16混合精度稳一波。
遇到过类似的,13B+多卡小batch确实容易这样,尤其前几百步loss跳高可能是某些layer的梯度异常,跟torch.compile关系不大。我当时是把gradient accumulation改成先本地累积再同步,同时把DDP的bucket_cap_mb调小到20,震荡明显缓解。另外建议你检查下不同卡的loss是否同步,有时候数据加载顺序不一致也会导致这种问题。纯个人经验,不一定对,但你可以试试看。
试过把梯度累积改成跨卡同步没?DDP下loss震荡大概率是bn or grad clip没对齐,跟模型大小关系不大。
这个思路不错,收藏了。
之前用DDP训7B也撞到过一模一样的墙,loss震荡得跟心电图似的。后来排查发现torch.compile在2.0里跟DDP的梯度钩子有兼容性问题,尤其是静态图模式下,每个rank的梯度归约时机不一致,导致优化器看到的全局梯度是“脏”的。你可以先试试把compile关掉,纯DDP跑个几百步对比下曲线,这能直接排除是不是编译优化引入的异步问题。
另外13B配4卡V100,每卡batch size才2,全局有效batch size算上累积其实只有64,对13B这种参数规模来说确实偏小。loss跳高又回落的现象,大概率是某些batch里出现了极端长尾样本,梯度范数瞬间爆表,而你的warmup只调了学习率,没做梯度裁剪。建议把max_grad_norm设到1.0甚至0.5,同时用gradient accumulation的loss缩放再核对一遍,有时候累积步数多了,loss归一化很容易出错。
还有个坑是数据加载的seed没固定,多卡下每个rank的shuffle顺序不一样,前几百步模型还在适应不同数据分布,震荡会更明显。你试试在DistributedSampler里设一个全局seed,并把drop_last设为True,避免最后几个不完整batch干扰。如果这些都没解决,可以考虑把学习率再降到1e-6级别,同时把warmup步数拉长到总步数的10%以上,大模型对lr的敏感度远比小模型高,我后来就是用这个组合稳下来的。
我之前也遇到过类似情况,尤其小模型到大模型切换时loss波动会明显放大。你试试把gradient accumulation改成先本地累积再同步,或者干脆关掉torch.compile看看,这俩在DDP下偶尔会搞出奇怪的数值抖动。另外13B在4卡上等效batch才64,对LLaMA来说确实偏小,我当初加到等效128才稳下来,你可以先拿小学习率跑个几百步看看趋势再调。还有个坑是V100对bf16支持不好,如果你用了混合精度,检查下是不是转成fp32了,不然梯度更新容易出问题。
调大有效batch size试试,13B这个规模4卡x2x8确实偏小,loss震荡大概率是全局batch不够稳。
梯度同步查下DDP的bucket_size,torch.compile有时会改变通信模式,关掉对比下最直接。
我最近也碰到过类似情况,13B模型在4卡上梯度累积8步其实等效batch也就64,对这么大模型来说确实偏小,loss震荡可能跟norm和梯度噪声都有关。你可以试试把梯度累积提到16步,或者改用AdamW加更激进的clip,比如max_grad_norm设到0.5,我这么调之后明显稳了不少。另外torch.compile在DDP下有时候会跟梯度同步打架,建议先关掉compile跑个几百步对比下,排除这个变量。还有个小细节,13B这种量级最好用bf16而不是fp16,V100不支持bf16的话就得小心loss scaling,可以看看是不是精度问题导致偶发loss跳高。
我之前在2.0上跑DDP也碰到过类似的情况,尤其是前几百步loss像过山车一样。后来发现多半不是梯度同步的问题,而是torch.compile在动态shape下会重新编译,导致前向计算不稳定,可以先试试把compile关掉或者用mode=reduce-overhead看看。另外13B配4卡确实偏激进,每卡batch2加梯度累积8步等效batch才64,对这么大模型来说可能真的不够稳,建议把梯度累积提到16步或者直接上更大batch,损失曲线会平滑很多。还有个小细节,检查下不同卡的数据加载顺序是不是一致,我之前就是DataLoader没设固定seed导致每个进程数据分布不一样,loss也是乱跳。
试试关掉torch.compile,2.0的DDP和编译叠加容易出诡异的数值问题,之前我也被这个坑过。
建议查一下梯度累积的同步逻辑,DDP默认是每个微批次都同步的,累积8步可能把噪声放大了。
我之前在2.0上跑DDP也遇到过类似的loss震荡,后来发现torch.compile跟DDP的梯度同步在某些版本下会有兼容性问题,建议你先试试关掉compile跑几百步对比一下。另外13B模型4卡每卡batch2确实有点小,等效batch才64,对于这个规模来说梯度噪声偏大,可以考虑把梯度累积加到16步或者直接上gradient checkpointing换更大batch。还有个细节,你确认一下DDP的bucket_cap_mb设置,默认值在模型大了以后可能会导致梯度更新不均匀,调成25左右试试。
建议先关掉torch.compile试试,DDP+编译在13B上梯度同步容易出问题。另外4卡V100跑13B本身显存就紧,batch太小loss不稳很正常。
我之前调7B模型的时候也遇到过一模一样的现象,后来排查下来发现大概率不是DDP本身的问题,而是数据加载顺序变了。多卡时每个rank拿到的数据子集不一样,如果你的数据集本身有顺序性或者类别分布不均匀,loss震荡就会特别明显,建议先shuffle一下再切分看看。另外你说的torch.compile,我建议先关掉试试,它在2.0里和DDP的梯度钩子偶尔会有交互问题,尤其是和梯度累积一起用的时候,容易产生数值抖动。关于batch size,13B在4卡上每卡2确实偏小了,虽然有效batch size是64(248),但BN或者某些layer对局部batch的敏感度还是会影响稳定性,可以试试把梯度累积改成16步,同时把学习率再降一个量级,比如1e-6起步。还有个小细节,warmup步数别只设几百步,13B这种规模建议至少跑1000步以上的线性warmup,不然前几百步的震荡是压不下去的。最后你提到loss跳高又降下来,我怀疑可能是某个特定样本触发了异常大的梯度,可以开一下grad clip,设到1.0左右,能有效抑制这种尖峰。如果还不行,试试把混合精度从bf16换成fp16,V100对bf16支持不好,精度问题也会放大loss波动。
你这loss震荡大概率是梯度累积和DDP不同步搞的,试试关掉torch.compile或者把累积步数设成1看看。
我之前跑7B的时候也遇到过类似情况,后来发现大概率不是DDP本身的问题,而是梯度累积和loss scaling之间的交互在搞鬼。你每卡batch size只有2,累积8步,等效全局batch才64,对13B模型来说确实偏小,尤其前几百步优化器还在预热期,梯度噪声会特别大。我建议你先试试把梯度累积关掉,直接加大每卡batch size到4或8,哪怕总batch小一点,梯度方差反而会降,loss会稳很多。另外torch.compile在V100上对某些算子会走fallback路径,可能引入数值抖动,你可以先用dynamo的mode='reduce-overhead'或者干脆关掉compile对比一下。还有个容易忽略的点:DDP的broadcast_buffers默认是True,如果你的模型里有BN或者某些自定义buffer更新频繁,多卡之间同步也会造成loss毛刺,可以显式设成False试试。最后,13B微调建议用AdamW加权重衰减,但beta2调到0.98左右,能显著抑制后期震荡,你可以顺手调一下。如果还不行,看看是不是数据加载每个rank的shuffle顺序不一致导致的,我之前就栽在这上面。
我之前微调7B也碰到过类似情况,后来发现是梯度累积和DDP的bucket通信没对齐导致的,你可以试试把梯度累积改成在DDP外面做,或者用NoSync模式手动控制同步时机。另外torch.compile在V100上可能触发某些算子重排,先关掉跑个baseline对比一下。13B这个规模确实对全局batch size敏感,但4卡x2x8=64的等效batch理论上不该这么抖,建议把gradient clip加上,并且确认下warmup步数是不是太短了。
这现象我遇到过,当时用DDP训7B也是前几百步loss跟过山车似的,后来排查发现是梯度累积和DDP的all-reduce交互出了问题。你累积8步但每步都在同步梯度的话,等效batch其实没变大,建议把累积逻辑改成梯度裁剪后再同步。另外13B配4卡V100确实吃紧,可以试试把lr再降到1e-6,或者关掉torch.compile看是否稳定。
跟你配置差不多,之前调7B也遇到过这问题,后来发现大概率不是DDP的锅,是torch.compile在动态shape下会搞出诡异的梯度更新。你试试把compile关掉跑几百步对比下loss曲线,如果稳了那基本就实锤了。另外13B配4卡确实梯度噪声大,你每卡batch2累积8步等效batch才64,对13B来说偏小,建议把梯度累积加到16或者直接上每卡4,看看震荡幅度会不会明显下降。还有个细节,warmup步数别太短,至少搞个总步数的5%以上,不然前期优化器状态没稳定也容易跳loss。