最近尝试用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模型也遇到过一模一样的现象,单卡稳得一批,一上DDP就疯。后来排查发现是不同卡的batch norm统计量没同步好,但你这个13B纯LLaMA应该没BN,那大概率是梯度累积和DDP的梯度平均交互出问题了,试试把梯度累积改成每张卡各自累积完再all-reduce,别用梯度累积跨步同步。另外13B在4卡上等效batch才64,对这么大的模型确实偏小,loss震荡可能跟优化器状态初始化有关,可以试着把Adam的epsilon调大一点到1e-6或1e-5,能缓解前期跳变。torch.compile先关掉排除变量,我这边开了compile后梯度同步时序会变怪,尤其是和DDP的bucket_size设置混在一起的时候。
我之前也遇到过类似的情况,DDP下loss震荡大概率不是torch.compile的锅,建议先把它关掉试试,因为多卡通信和graph capture叠加有时候会出怪问题。13B配4卡每卡bs=2加梯度累积,等效batch才64,确实偏小,尤其是微调大模型,梯度噪声会明显放大,可以试试把学习率再降一个量级或者加大累积步数。另外检查一下不同卡的data shuffle是不是同步的,还有loss打印是不是只取了rank0的,如果没同步的话看起来会更乱。前面几百步震荡剧烈有时候是layer norm和embedding在适应新分布,可以试试冻结前几层只训练后半部分,等稳定了再解冻。
试试关掉torch.compile,2.0的DDP对编译图和梯度同步的兼容性偶尔会抽风。
我之前跑7B也遇到过类似情况,尤其是用torch.compile配合DDP时,loss曲线跟心电图似的。后来排查发现是梯度累积和DDP的梯度同步顺序冲突,导致累积的梯度被部分覆盖了,你可以试试把梯度累积改成手写循环,或者在backward后加个梯度平均再step。另外13B在4卡上每卡2的batch确实偏小,相当于全局才8,Llama这种模型对全局batch挺敏感的,建议要么加大到每卡4,要么用gradient checkpointing腾显存来提batch。还有个坑是warmup步数可能不够,前几百步震荡正常,但你要看下震荡幅度是不是逐渐收敛,如果一直发散就得查下数据sampler是不是没设shuffle,多卡下数据分布不均也会导致这个。
DDP前几百步loss震荡太正常了,试试把梯度累积换成grad clip到1.0,顺便关掉torch.compile看下。
我之前在2.0上跑DDP也遇到过类似的震荡,后来发现torch.compile跟梯度累积一起用容易出问题,尤其是动态shape的时候,建议先关掉compile纯跑DDP对比一下。另外13B这个规模,4卡总共才8的batch,等效batch太小了,V100上显存又紧,可以试试梯度checkpointing把batch提到4,或者直接上DeepSpeed的ZeRO阶段2,loss会稳很多。还有个细节,你warmup步数是不是跟着batch size也放大了?多卡下梯度更新频率变了,warmup没调的话前几百步震荡很正常。
之前用DDP训6.7B也遇到过类似情况,后来发现torch.compile在2.0里跟DDP的梯度桶同步有兼容问题,关掉compile之后loss稳了不少。另外你梯度累积8步但每卡batch只有2,等效全局batch才64,对13B来说确实偏小,建议把梯度累积提到16或者直接加大每卡batch,试试看会不会改善。还有个小细节,检查下DDP的find_unused_parameters是不是设了True,有时候模型里某些层没参与训练会导致梯度同步异常。
我之前也遇到过类似的情况,不过是在多卡跑7B的时候。你那loss跳高又降下来,感觉更像是梯度不同步或者某个batch的噪声太大,13B模型4卡每卡2的batch确实偏小,等效batch才64,大模型对梯度噪声很敏感,建议试试把梯度累积加到16步或者直接上gradient checkpointing换更大batch。torch.compile在这种场景下有时会改变算子融合顺序,导致数值抖动,可以先关掉compile对比一下。另外检查下DDP的bucket_cap_mb,默认值在模型大时可能导致通信和计算重叠不好,调小到25左右有时候能稳定不少。
我之前调7B也遇到过类似的,后来发现torch.compile在DDP下跟梯度累积有兼容问题,你先试试关掉compile纯跑DDP对比下。另外13B这个规模4卡V100确实显存吃紧,有效batch才64,对13B来说偏小了,loss震荡可能跟这个有关。建议把梯度累积提到16步,或者干脆用梯度checkpointing把batch size往上顶一顶,看看曲线会不会稳下来。还有个细节,warmup步数别太少,至少几百步起步,不然前段学习率变化太剧烈也会加剧震荡。
我之前在32卡上跑过7B模型,也遇到一模一样的loss震荡,后来排查下来发现大概率不是torch.compile的锅,而是DDP在梯度累积和all-reduce之间的交互出了问题。你每卡batch size=2,累积8步,相当于有效batch size=64,但DDP默认是在每步(每2个样本)就做一次梯度同步,这样累积8步期间,每步的梯度方向差异很大,尤其是前几百步模型还在快速适应阶段,all-reduce出来的梯度噪声会被放大,看起来就是loss乱跳。你可以试试把梯度累积的粒度改一下,比如用no_sync上下文管理器包住前7个微步,只在第8步触发同步,这样梯度同步次数直接减少8倍,震荡会明显缓解。另外13B模型在V100上跑,梯度精度fp16混合精度下,loss突然跳高也可能是inf/nan导致的,建议开一下gradient clipping,max_norm设1.0左右,同时检查一下是不是某些层在同步后参数更新过大。还有个小细节,warmup步数最好跟着有效batch size走,你改成64后,warmup至少要3000步起步,不然前几百步学习率还在爬升,叠加噪声会更乱。我后来把这几样都改了,loss曲线就平滑多了,你可以试试看。
之前做多卡微调也碰到过类似情况,后来发现大概率是梯度累积和DDP的梯度同步顺序不对,PyTorch 2.0里累积步数没整除卡数的话,某些step的梯度会重复累加,导致loss突然跳高。你可以试试把梯度累积改成8*4=32,或者干脆用zero2/zero3这类分片优化器,同时把torch.compile先关掉对比一下,有时候编译会引入数值微小的变化,在小batch下放大了震荡。另外13B在4卡上每卡2的batch确实偏小,建议先跑个短实验,把lr降到1e-6看曲线形态,如果还是抖,大概率是数据sampler没设shuffle的seed,多卡下每个epoch数据顺序不一样。
我之前在2.0上跑DDP也遇到过类似情况,尤其是用torch.compile的时候,loss曲线会莫名抖动,后来发现是编译优化和梯度累积的交互有点问题,建议你先关掉compile试试纯eager模式对比一下。另外13B模型在4卡上每卡batch size=2,全局才8,梯度累积8步等效batch=64,对13B来说其实偏小了,LLaMA系列对batch size很敏感,尤其微调时小batch容易让loss漂移。还有个坑是DDP的梯度同步时机,如果你在累积梯度时用了no_sync,但某些卡步数没对齐,会导致某几步的梯度实际上是部分更新的,那loss跳高就很正常了。我后来是把梯度累积改成在DDP外面手动做,每步都同步梯度,但用optimizer.step()间隔执行,这样loss稳很多。你还可以检查一下数据加载的shuffle,多卡下如果每个epoch的shuffle种子不一致,不同卡看到的数据分布差异大,也会加剧震荡。最后建议你把warmup步数拉长到总步数的10%左右,3e-6的学习率配warmup按理说够低了,但13B微调时前几百步本来就会有个适应期,不一定全是bug。
13B换4卡V100,显存瓶颈下梯度累积步数太长,等效batch不够大,试试直接砍半累积步数。
我之前也遇到过类似的,13B上DDP前几百步loss不稳定挺正常的,尤其是从单卡切到多卡,effective batch size变了(你4卡×2×8=64,不算小但可能对13B还是偏小)。建议先别急着上torch.compile,那玩意儿对显存和梯度同步有额外开销,V100上不一定划算。可以试试把梯度累积去掉,直接加大per-gpu batch到4或8,看看震荡有没有缓解;另外确认一下你是用的DistributedSampler并且shuffle=True,数据顺序不一致也会导致loss跳变。如果还不行,查一下不同卡的梯度范数是不是差异很大,可能是某些层初始化或者数据分布的问题。
大概率是梯度累积和DDP的梯度同步没配合好,试试关掉torch.compile,或者把梯度累积改成手动accumulate。
之前调7B也遇到过一模一样的现象,后来排查发现是梯度累积和DDP的梯度同步顺序没对上,导致等效batch size其实没变,你试试把梯度累积改成all-reduce之后再做,或者干脆先关掉torch.compile跑几百步对比一下。另外13B在4卡上每卡batch 2确实偏小,BN或者layernorm的统计噪声会被放大,建议梯度累积步数再翻倍,或者用gradient checkpointing腾显存把batch提到4。还有个细节,warmup步数别跟着batch size线性调,按总token数算会更稳,前几百步loss跳高有时候是数据顺序问题,shuffle种子换一下可能就没了。
梯度累积8步配合DDP,实际等效batch才64,13B模型这个规模确实偏小了,试试把累积提到16步。
之前跑多卡也遇到过类似情况,特别是模型一大loss就特别敏感。我觉得大概率不是torch.compile的锅,DDP的梯度同步本身是没问题的,你可以先关掉compile试试排除变量。另外13B确实比1.3B对batch size和lr更敏感,你4卡x2的batch其实等效才8,对13B来说可能偏小了,建议先试下把梯度累积加到16或者直接每卡batch提到4,看看震荡幅度有没有变化。还有个细节,DDP前几百步震荡有时候是BN的统计量在作怪,但LLaMA没有BN,所以更可能是优化器状态初始化的问题,可以试试把AdamW的beta2调低到0.95,或者用bfloat16混合精度跑一下,V100上fp16的loss缩放也可能导致跳变。
之前调7B也遇到过一模一样的现象,后来发现是梯度累积和DDP的bucket通信没对齐导致的,你试试把gradient_as_bucket_view设成True,再把no_sync和accumulation的搭配检查一下。另外13B模型4卡V100的话,单卡batch2确实偏小了,等效batch才64,对13B来说有点极限,建议把梯度累积加到16步试试,或者干脆用DeepSpeed的ZeRO阶段2。还有个坑是torch.compile在DDP下偶尔会改变算子融合顺序,先关掉compile跑几百步对比下loss曲线,能排除不少干扰。
之前跑7B的时候也遇到过类似的震荡,后来发现一个比较隐蔽的点:DDP默认的gradient all-reduce是在每次backward之后同步的,但如果你用了梯度累积,得确保只在累积到第8步的时候才做同步,否则每步都同步的话等效batch size其实没变,loss自然不稳。可以检查一下是否用了NoSync控制或者手动把gradient accumulation和DDP的bucket协调好。另外13B在4卡上每卡batch=2,全局batch=8,对于这个量级确实偏小,尤其是LLaMA这种模型对batch size挺敏感的,建议试试把梯度累积加到16或者32步,等效batch到64-128,学习率可以再往下降一档。还有个坑是torch.compile在V100上对某些算子会触发fallback,导致前向计算有微小差异,叠加多卡后就被放大了,可以先关掉compile纯DDP跑几百步对比下loss曲线,排除这个变量。最后检查一下数据加载,多卡时每个rank的shuffle种子要设一致,不然不同卡喂的数据分布都不同,loss也会抖。