最近尝试用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下loss震荡大概率不是torch.compile的锅,先试试把compile关掉跑几百步对比下。另外13B配这个有效batch size其实不算大,4卡×2×8=64,对13B来说确实偏小,loss跳高可能是某些batch的梯度方向太极端。建议你先把梯度裁剪加上,max_grad_norm设个1.0,再把warmup步数拉长到总步数的10%,看看前500步能不能稳住。如果还不行,试试用zero冗余优化器替代DDP,显存换稳定有时候挺值的。
调成梯度累积16步或者直接上bs=4试试,13B对噪声比小模型敏感得多,我这边是这么稳住的。
试试关掉torch.compile,2.0的DDP和编译叠加偶尔会出玄学问题,纯DDP跑跑看对比下。
大概率是梯度累积和DDP的梯度同步没配合好,试试关掉compile或者把累积步数换成梯度检查点。
之前遇到过类似问题,把lr降到1e-6加长warmup能稳点,但13B多卡确实容易这样,batch再大点试试。
看到13B直接上4卡DDP,前几百步loss震荡其实挺常见的,尤其你梯度累积8步,等效batch也就64,对13B来说偏小了。我之前调7B的时候发现,torch.compile跟DDP的梯度同步在某些版本下会有交互问题,建议先关掉compile试试。另外可以查一下不同卡的loss是不是本身就不一致,有时候数据shuffle没设好种子也会导致这种跳变。
我之前也遇到过类似问题,后来发现大概率不是DDP本身,而是torch.compile在2.0里跟DDP的梯度归并逻辑偶尔会打架。你可以先关掉compile试试,纯DDP跑个几百步对比下loss曲线,如果稳了那就基本实锤了。另外13B模型每卡batch=2确实偏小,梯度累积8步虽然等效batch=64,但BN或者某些层的统计量在多卡下还是会有差异,建议把累积步数加到16甚至32,同时把lr再降个量级,比如5e-7起步,warmup步数拉长到总步数的10%再观察下。我之前用8卡跑7B,前200步loss也会跳,后来发现是数据采样顺序不一致导致的,你可以检查下每个rank的shuffle种子是不是没设好。
我之前微调7B也遇到过类似情况,后来发现是DDP的bucket划分和梯度累积交互的问题,试试把gradient_as_bucket_view设成True,或者干脆关掉torch.compile看下,感觉compile对动态loss曲线影响挺大的。另外13B在4卡上每卡才2的batch确实有点小,等效batch才64,对13B来说梯度噪声偏大,建议至少凑到128以上,或者试试用AdamW的betas调低一点,比如0.9和0.95,能稍微压一下震荡。你检查过不同卡上的数据shuffle顺序吗,如果各卡数据分布不均也会导致loss跳变。
先确认下torch.compile关了没,DDP下它跟梯度累积容易冲突。另外13B这规模lr确实得再砍半试试。
13B这个规模4卡确实吃紧,试试关掉torch.compile或者把梯度累积改成同步bn,说不定能稳下来。
13B这规模DDP前几百步loss乱跳太正常了,建议先关掉torch.compile跑跑看,大概率是图编译和DDP通信撞了。
梯度累积8步配4卡,等效batch才64,对13B来说确实偏小,试试把每卡batch提到4或者累积步数翻倍。
试试把梯度累积改成跨卡all-reduce后再累积,或者直接上梯度裁剪,我之前碰到过类似情况。
我之前在2.0上跑DDP也碰到过类似情况,后来发现是torch.compile和DDP的梯度钩子配合有问题,关掉compile或者用ddp_find_unused_parameters=False试试,loss会稳不少。另外13B配4卡每卡batch 2确实偏小,等效batch才64,LLaMA这种模型建议至少等效128起步,不然前几百步loss跳高可能是某些层梯度没同步好。你可以先不compile跑个500步对比下,如果还不稳再查下数据shuffle是不是每卡重复了,这个坑也挺隐蔽的。
同款配置踩过坑,13B这个规模DDP前几百步loss跳高基本不是lr的锅,大概率是LayerNorm和embedding的梯度在不同卡上同步延迟导致的。你可以试试把torch.compile关掉,然后梯度累积改成先all-reduce再更新,或者直接上FSDP,显存占用和稳定性都会好很多。另外每卡bs=2对13B确实偏小,effecitve batch要凑到64以上才稳,你累积8步其实才16,试试把累积加到16步看看。
试试关掉torch.compile,2.0的DDP和编译图叠加有时会出同步问题,loss跳高多半是梯度没规整好。
我之前微调7B的时候也遇到过一模一样的现象,单卡稳如老狗,一上DDP就开始抽风。后来排查下来,发现大概率不是torch.compile的锅,而是梯度累积和DDP的梯度同步逻辑在打架——你想想,梯度累积8步意味着每8个micro-batch才做一次all-reduce,但DDP默认是每次backward都触发同步的,如果没配合好,等于是把没累积完的梯度给平均了,loss不震荡才怪。建议你查一下是不是用了NoSync或者梯度累积的官方写法,确保只在累积到第8步时才同步。另外13B在4卡上每卡batch=2确实偏小,等效全局batch才8,对LLaMA这种模型来说,BN或者layer norm的统计量估计会很不稳定,我自己是把全局batch提到16以上才明显改善,你可以试试梯度累积改成16步。还有个细节,warmup步数别只加几百步,13B模型前几千步都在适应分布式优化器的状态,我那时候warmup直接拉到总步数的5%才压住尖峰。你要是开了amp混合精度,也可以检查下loss scaling是不是在频繁调整,那个也会导致loss突然跳变。
建议看下不同卡间数据shuffle是否一致,DDP下没设seed的话每个进程数据顺序不同会加剧震荡。
我之前微调7B的时候也遇到过类似情况,尤其是梯度累积配合DDP,其实很容易出问题。你可以检查一下是不是不同卡的loss本身就不同步,因为梯度累积8步但DDP的all-reduce是每步都做的,这样等效batch其实没到你想的那么大,反而放大了噪声。另外torch.compile在2.0初期对DDP支持有bug,建议先关掉compile纯跑一次对比,如果loss稳了那就基本定位了。还有个小细节,warmup步数要跟总训练步数匹配,前几百步震荡有时候是优化器状态没预热好,试着把warmup拉长到总步数的10%看看。
13B模型用4卡V100、每卡batch size才2确实有点小,梯度累积虽然能补等效batch,但DDP的梯度同步是在累积之前做的,每步all-reduce的噪声其实没被平均掉。前几百步震荡大概率跟这个有关,再加上V100不支持bf16,混合精度用fp16的话loss scale抖动也会放大不稳定。建议试试把per-device batch提到4或8,哪怕seqlen砍一点,或者换用梯度累积后再同步的方案(比如zero或fsdp的梯度分片)。torch.compile在DDP下偶尔会有图捕获和同步的坑,可以先关掉compile单独验证一下。
梯度累积加DDP,loss scale没同步好就容易炸,建议先关掉torch.compile跑几十步看看。
梯度累积加DDP确实容易震,试试把gradient accumulation去掉改成真实大batch看看还抖不抖。