最近尝试用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 条试试把梯度累积改成4步,或者检查下DDP的广播机制,可能是同步频率太高导致震荡。
大概率是梯度同步和lr没对齐的问题,13B模型单卡有效batch其实只有2*8=16,对这么大的模型来说确实偏小了,建议把梯度累积提到16-32步试试,或者直接增大每卡batch。另外torch.compile在DDP下有时会引入数值抖动,尤其是前向中有动态shape时,可以先关掉compile跑几步对比下loss曲线。如果还震,可以检查下数据加载顺序是不是不同卡之间shuffle不一致,我之前遇到过这个坑。
试试把梯度累积换成更大的batch size,4卡每卡4步,loss平滑很多,torch.compile有时会干扰梯度同步。
这问题我也遇到过,13B模型在4卡上梯度累积8步相当于全局batch才64,对LLaMA来说确实偏小了,loss震荡很正常。建议试试把梯度累积再翻倍到16步,或者每卡batch size提到4(如果显存撑得住),另外torch.compile在DDP下有时会搞乱梯度同步,可以先关掉compile跑跑看。还有检查下warmup步数是不是太短,至少搞个500步线性增长会稳一些。
这情况我也遇到过,13B用4卡V100跑DDP,每卡batch size 2确实有点极限,梯度累积8步等效batch size才64,对这么大模型来说loss震荡挺常见的。你可以试试把梯度累积加到16步或者直接用更大的batch size,另外检查下torch.compile是不是默认开了动态shape,有时候会干扰梯度同步。还有个小技巧是把allreduce的bucket大小调小到25M,能缓解部分震荡。warmup步数建议拉到总步数的10%以上,学习率3e-6对13B可能还是偏高了点,可以再降到1e-6看看。
这问题我踩过类似的坑,13B模型用4卡V100微调,每卡batch size才2,算上梯度累积8步,等效batch size也就16,对于13B这种规模确实偏小了。loss震荡大概率不是torch.compile的锅,而是梯度噪声太大,尤其是前几百步,模型参数初始化后各卡梯度方向不一致,DDP同步时容易产生异常跳变。我之前调7B模型也遇到过,后来发现把学习率降到1e-6左右,同时把warmup步数拉到总步数的10%以上,震荡会明显缓解。另外你可以检查下数据加载是不是shuffle了,有时候多卡数据分布不均也会导致loss异常抖动。还有个可能是梯度累积时的梯度裁剪,建议设个1.0的max_grad_norm,能防止单卡梯度爆炸影响全局。如果还不行,试试关掉torch.compile或者只对部分层编译,PyTorch 2.0的DDP在某些场景下和编译优化确实有兼容问题。
说实话我也踩过类似的坑,13B上DDP震荡挺常见的,尤其是梯度累积加torch.compile一起用的时候,偶尔会出诡异的梯度不同步。你可以试试先把compile关掉跑个几百步对比下,如果震荡消失那基本就是编译优化跟DDP的hook有冲突。另外batch size每卡2确实偏小了,对13B来说有效batch相当于64,梯度方差会比较大,试试把梯度累积加到16甚至32,或者换用8卡把每卡batch减到1但累积拉长,loss会更稳。还有一个容易被忽略的点是warmup步数,13B这种规模建议至少500步起步,你现在的warmup如果太短前几百步震荡就是正常的。
调低lr不如试试把梯度累积步数翻倍,13B卡小显存容易炸。
建议关掉torch.compile,2.0的图模式跟DDP偶尔有冲突,裸跑看看稳不稳。
看到你这个情况我第一反应是batch size确实太小了,13B模型用4卡每卡2的batch其实等效才8,加上梯度累积8步也就64的全局batch,对大模型来说梯度噪声会非常大。我之前用8卡A100跑7B模型都试过类似震荡,后来把每卡batch提到4、累积步数减到4才稳住,你可以试试保持总batch不变但减少累积步数,看看是不是梯度更新频率的问题。
另外torch.compile在DDP下有时会搞出奇怪的数值抖动,尤其是跟梯度裁剪配合不好的时候,建议先关掉compile跑几个epoch对比下。还有warmup步数要拉长,13B这种规模至少500步起步,你用的线性warmup还是cosine?cosine + 长warmup对震荡抑制效果会明显好一些。
不过说实话,loss前几百步震荡也有可能是正常的,大模型多卡训练初期优化器状态没稳定,特别是AdamW的动量积累需要时间,只要后期能收敛就别太焦虑。建议你监控下梯度范数,如果震荡时梯度突然爆高,那可能是梯度同步的allreduce延迟导致参数更新不同步,可以试试调大NCCL的buffer或者换用gradient clipping。最后问下你用的是bf16还是fp16?混合精度模式下loss震荡会更敏感,我遇到过用bf16反而比fp16稳的情况。
我最近在调一个30B模型时也碰到过类似情况,后来发现是torch.compile的图优化跟DDP的梯度同步有点冲突,关掉compile之后loss震荡明显缓解了。另外13B用4卡每卡batch=2确实偏小,梯度累积虽然能模拟大batch但BN层统计会有偏差,建议试试把每卡batch提到4甚至8,或者换用gradient checkpointing来省显存。你用的学习率调度器是cosine还是linear?warmup步数加到总步数10%以上可能也有帮助。
看到你这个情况我太有同感了,之前用PyTorch 2.0的DDP调7B模型也遇到过类似的震荡,后来排查发现torch.compile默认的reduce模式在某些场景下会让梯度同步出现细微的时序偏差,尤其是小batch加梯度累积的时候。你试试关掉compile或者改成mode=“reduce-overhead”看看?另外13B模型在4卡上每卡batch=2确实太小了,梯度累积8步等效batch才16,对于这么大模型来说学习率1e-5可能还是偏高,我建议你试试先固定前100步用更小的学习率比如1e-6,再慢慢warmup到目标值。还有个小细节:DDP的梯度同步默认是异步的,如果你没用torch.nn.parallel.DistributedDataParallel的find_unused_parameters=False参数,有些层的梯度可能被跳过,导致loss突然跳变。最后检查下梯度裁剪是不是设得太激进,我上次把max_norm从1.0改成0.5反而加剧了震荡,得根据实际loss曲线慢慢调。
这种情况我调13B时也遇到过,后来发现是梯度累积+多卡导致的BN统计量差异,建议先关掉torch.compile跑几轮看看,有时候动态图优化反而会干扰DDP的梯度同步。另外你试试把lr降到1e-6,并且把warmup步数拉长到总步数的10%,大模型对初始阶段的稳定性要求确实更高。如果震荡还持续,可以考虑用zero2代替DDP,显存占用差不多但梯度通信更平滑。
同款踩坑,13B用DDP加torch.compile确实容易炸loss,我怀疑跟编译时图优化对分布式通信的干扰有关,建议先关掉compile试试。另外每卡batch size才2的话,梯度累积8步等效全局batch才64,对13B来说确实偏小,可以试试把学习率降到1e-6左右,或者把梯度累积加到16步,让参数更新更稳定。还有检查下DDP的bucket_cap_mb参数,默认25可能太小,增大到200能减少通信开销,偶尔能缓解震荡。
这个坑我也踩过,13B模型在DDP下loss震荡大概率不是torch.compile的锅,而是梯度累积+多卡同步导致的数值敏感性问题。建议你把梯度累积改成4步试试,同时检查一下每张卡的loss实际差异,可能是数据不均匀或者bn层同步出问题了。另外可以试一下先把compile关了跑几百步看看是不是它引入的抖动,我之前用2.0的inductor就遇到过类似情况。
这个坑我当初也踩过,13B模型在DDP下loss震荡太正常了,尤其前几百步。你每卡batch size才2,虽然梯度累积了8步,但等效全局batch size其实也就32(4卡x2x8步/2?不对,应该是4x2x8=64?我算晕了),对于13B模型来说确实偏小,梯度噪声会很大。我建议你先试试把梯度累积再翻一倍到16步,或者每卡batch size提到4(如果显存扛得住),让等效batch size到128以上,loss曲线会稳很多。
另外torch.compile在DDP下有时候会搞出一些奇怪的数值抖动,尤其大模型用动态图编译时,梯度同步时序可能被打乱。你可以先关掉compile跑个几百步对比看看,如果没震荡基本就是编译的问题。我之前用2.0跑7B模型时也遇到过类似情况,关掉compile后立竿见影。
还有一个细节:warmup步数你设了多少?13B模型建议warmup至少占到总步数的5%-10%,而且学习率从0开始线性增长,你从1e-5往下调可能没找对方向。试试把初始学习率压到1e-6甚至更低,warmup拉长到500步以上,再观察震荡峰值有没有收敛。如果还是跳高,检查下数据加载是不是有shuffle不一致导致各卡数据分布差异太大,尤其是微调时数据量小的话更容易出问题。
之前用PyTorch 2.0跑13B模型也遇到类似问题,后来发现torch.compile在DDP下对某些算子会引入额外的不确定性,可以试试先关掉compile跑几步看看。另外13B模型显存压力大,每卡batch size才2的话梯度估计噪声太大,建议试试梯度累积步数翻倍到16,或者把lr再往低调到1e-6看看。还有个坑是不同卡之间数据shuffle不一致也可能导致震荡,检查下dataloader的seed设置。
同样在DDP上踩过坑,感觉你这情况大概率是梯度同步和batch norm的锅,13B模型每卡batch 2确实偏小了,梯度累积虽然能模拟大batch但同步频率没变。可以试试不用torch.compile,换成纯eager模式跑几百步对比下,有时候编译优化反而会引入数值抖动。另外检查下warmup步数够不够,13B这种规模前1000步用线性warmup会稳很多。
试试调大梯度累积步数到16,或者把每卡batch size加到4,小batch在多卡下容易震荡。
感觉你这问题大概率是梯度同步和batch size不匹配导致的,13B模型本身对梯度噪声就更敏感,4卡每卡2的batch相当于总batch才8,建议先把梯度累积提到16或32试试,同时确保DDP的bucket尺寸设置合理。另外torch.compile在2.0里对大模型的动态图支持可能还不够稳,可以先关掉compile跑跑看,排除这个干扰。我之前用8卡跑7B模型也遇到过类似震荡,把学习率降到1e-6并延长warmup到500步才稳住。
这个情况我也遇到过,13B用DDP训练,loss震荡确实比单卡明显很多。我觉得核心问题可能不在torch.compile,虽然它有时候会引入一些数值抖动,但你这个batch size折算下来全局才248=64条,对于13B模型来说确实偏小了,BN层或者LayerNorm在小batch下统计量不稳定,梯度噪声会放大。另外你提到loss突然跳高又降下来,我怀疑是不是梯度裁剪没设好,或者DDP默认的bucket容量导致梯度同步有延迟?建议试试把梯度累积改成每步同步一次,或者把allreduce的bucket大小调小一点,减少不同卡之间的异步效应。还有一个小技巧:把学习率warmup步数拉长到总步数的10%左右,配合线性衰减,前几百步的震荡会缓和很多。如果方便的话,可以试试开gradient checkpointing降低显存占用,把实际batch size再往上提一提,64对于13B来说确实有点危险。