最近尝试用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震荡严重,求大佬指点
全部回复
共 13 条我最近也碰到过类似问题,感觉13B模型在4卡上梯度累积8步其实等效batch size才64,对这么大模型来说可能确实偏小了,loss震荡很常见。你试过把梯度累积加到16或者直接增大每卡batch size吗?另外torch.compile有时候在DDP下会有诡异的数值抖动,我关掉之后稳定了不少。建议先排查下是不是梯度同步时机的问题,可以打印下各卡梯度norm看看有没有异常。
可能是梯度累积和DDP的BN同步冲突了,试试关掉torch.compile看看稳不稳。
试试关掉torch.compile,13B用DDP时编译优化反而容易导致梯度不同步。
DDP下13B模型loss震荡确实挺常见的,我猜不完全是torch.compile的锅,更多是有效batch size不够大。你4卡每卡2条加8步累积,实际batch size才64,对13B来说确实偏小了,梯度噪声会比较大。可以试试把学习率再降到1e-6,或者干脆把梯度累积提到16步,先让前几百步稳下来再说。另外检查下不同卡的梯度是否一致,有时候数据加载不均匀也会导致这种突然跳变。
这个现象我调4卡DDP时也遇到过,13B模型对显存和通信压力都大,loss震荡很可能跟梯度同步的噪声有关。建议你试试把梯度累积步数翻倍到16,或者把每卡batch size提到4(如果显存够的话),这样能缓解小batch带来的方差问题。另外torch.compile有时候在DDP下会引入额外的不稳定性,可以先用eager模式跑跑看对比一下。
这问题我最近也遇到了,13B模型上了多卡之后loss震荡真的挺头疼的。你说怀疑梯度同步,我觉得大概率不是torch.compile的问题,毕竟2.0的DDP后端对同步机制已经优化过了,反而更可能是你的有效batch size太小了——每卡2条×4卡×8步累积,其实才64条,对于13B这种大模型来说,梯度噪声太大了,前几百步模型还在找方向,震荡是正常的。我之前试过把梯度累积加到12步,同时把学习率降到2e-6,warmup从100步拉到500步,震荡明显收敛了一些,但偶尔还是会有loss spike,后来发现是数据加载里有个shuffle没处理好,导致某些batch里特殊样本过于集中。另外,你用的是bf16还是fp16?如果是fp16,可能会因为梯度溢出丢精度导致跳变,可以试试设置grad_scaler的growth_interval大一点,或者干脆换bf16。还有个野路子,我试过在DDP初始化时把find_unused_parameters设为True,虽然官方说没必要,但对某些带条件分支的模型确实有奇效。你要是方便的话,可以跑一版纯单卡(只改batch size凑同等有效量)对比下,如果单卡也震荡,那就是模型本身对学习率敏感,跟多卡无关了。
试试把梯度累积改成4步,batch size调大点,13B模型对batch size敏感,太小容易震荡。
这个我也有类似经历,13B用DDP开compile容易出幺蛾子,尤其是torch.compile的inductor后端在某些算子同步上会搞出诡异的梯度噪声。建议你先关掉compile跑跑看,如果loss稳了那八成就是它的问题。另外你梯度累积8步但每卡batch只有2,等效全局batch才64,对13B来说确实偏小了,试试把梯度累积提到16或者32,学习率再降到1e-6左右,前500步让loss先降下来再调高。还有检查下DDP的bucket_cap_mb,默认值太大可能导致小batch时梯度更新不同步,改成25试试。
这情况我跑7B模型也遇到过,前几百步loss跳得跟心电图似的,后来排查发现大概率是梯度同步的问题。你每卡batch size 2加梯度累积8步,实际上有效batch size才64,对于13B模型来说确实偏小了,模型参数空间大但有效数据量不够,优化器容易在局部震荡。建议你先试试把梯度累积提到16步,或者每卡batch size加到4(如果显存够),让有效batch size到128以上再看看震荡幅度。另外torch.compile在DDP下对某些操作会有数值抖动,尤其是flash attention或者fused kernel的版本差异,你可以先关掉compile用纯eager模式跑几百步对比一下loss曲线。还有个小细节:warmup步数建议至少占训练总步数的5%-10%,你如果只加了几十步warmup,对13B这种大参数收敛稳定性帮助有限。最后检查一下dataloader的shuffle和分布式采样器是不是配置正确,有时候数据分布不均也会导致loss忽高忽低。
试过关掉torch.compile再跑吗?有时候它跟DDP的梯度同步会打架。
之前也遇到过类似情况,13B模型在DDP下batch size偏小确实容易导致梯度噪声大,尤其是前几步。建议试试把梯度累积步数再翻倍,或者用zero2代替DDP,能缓解同步时的震荡。另外torch.compile对动态图支持有时会引入数值波动,可以先关掉compile跑几百步对比一下,排除这个干扰。
我也遇到过类似情况,13B模型在DDP下loss震荡挺常见的,尤其是梯度累积配合小batch时,不同卡上的梯度方差会被放大。建议试试把梯度累积去掉,直接增大每卡batch size到4或8,虽然显存压力大但梯度更稳。另外torch.compile在DDP下有时会引入数值波动,可以先关掉compile跑跑看,排除这个因素。
我最近也遇到过类似的情况,13B模型在DDP下loss震荡确实挺常见的。个人感觉不光是batch size的问题,torch.compile有时候反而会引入一些数值不稳定,可以试试先关掉compile跑跑看。另外检查一下梯度同步的all-reduce是不是正常,有时候卡间通信延迟会导致梯度不一致,尤其是前几步随机初始化时更明显。建议你先把梯度累积步数翻倍试试,或者给不同卡设不同的随机种子做对比实验。