最近尝试用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 条看到你这个情况我第一反应是——大概率不是torch.compile的锅,我之前试过在DDP下用compile,虽然偶尔有玄学bug,但不会导致loss这么剧烈的震荡。你试试把梯度累积改成单步看看,我怀疑是累积8步配合DDP的allreduce时机出了问题,尤其是前几百步模型还没稳定的时候,梯度累积会让不同卡的梯度在时间上错位更严重。另外13B模型在4卡上每卡batch size=2其实偏小了,虽然梯度累积等效batch=64,但DDP的BN同步(如果有)和参数更新的频率还是受影响,你可以试试把每卡batch提到4,梯度累积减到4步,看看震荡会不会收敛。还有个小细节——检查一下你的学习率是否按卡数做了线性缩放,虽然warmup能缓解,但1e-5对4卡来说如果没缩放可能还是太大了。我调7B模型时遇到过类似跳高又回落的情况,最后发现是数据加载顺序在不同卡上不一致,导致某些step的梯度方向冲突,建议你看看dataloader的shuffle是不是没设固定seed。最后如果还是不行,可以试试先把torch.compile关掉,只用DDP跑个step看看,排除编译优化带来的数值噪声。
试过关掉torch.compile用原生DDP跑吗?感觉可能是图编译跟梯度同步的兼容问题。
看到这个情况我第一反应是梯度累积和DDP的同步问题,13B模型用4卡V100每卡batch size才2,全局batch size其实很小,但梯度累积8步之后有效batch size是64,理论上不算太小。不过你提到前几百步震荡厉害,我怀疑是torch.compile在DDP下对某些算子做了图优化,导致梯度同步时机变了,我之前用2.0的时候也遇到过类似问题,关掉compile之后反而稳定很多。另外可以检查一下是否用了梯度裁剪,13B模型训练时梯度范数容易炸,设置max_grad_norm=1.0能抑制那种突然跳高的loss。还有个小细节是warmup步数最好跟总训练步数匹配,比如总步数5000的话warmup设个200-300步,太短了效果不明显。至于batch size,你现在的有效batch 64对13B模型来说确实偏小,如果能加大到128或256震荡应该会缓解,但受限于显存可以试试梯度累积再翻倍到16步。最后建议你把DDP的find_unused_parameters设为False,如果模型里有dropout或者layer dropout,参数不同步也可能导致loss异常。
我之前调13B也遇到过类似问题,后来发现是梯度累积和DDP的all_reduce时机没对齐导致的,试试把梯度累积改成手动实现,或者检查下no_sync上下文管理器用得对不对。还有可能是torch.compile的图优化和DDP的梯度桶有冲突,可以先关掉torch.compile跑几步对比下。另外13B确实对全局batch size敏感,4卡每卡2步累积8步等效batch才64,偏小了,有条件可以试试增大到128或256。
我之前用Deepspeed ZeRO-3跑类似规模的模型也遇到过loss震荡,后来发现是梯度all-reduce的精度问题,尤其是混合精度训练下fp16梯度容易出异常。你可以试试把DDP的bucket_cap_mb调小到25或50,强制更频繁的同步,或者检查下torch.compile是不是对某些算子做了非确定性优化。另外13B模型每卡batch size 2确实有点极限,梯度累积8步等效全局batch才64,建议至少提到128以上,学习率再降到1e-6量级试试,前200步慢慢预热可能比warmup更稳。
这个坑我去年也踩过,13B模型在4卡上batch size每卡2确实太极限了,梯度累积8步等效batch才64,对大模型来说全局batch太小会导致梯度噪声爆炸,loss震荡几乎是必然的。你可以试试把每卡batch size提到4,梯度累积降到4步,虽然显存压力大点,但V100 32G版应该能撑住。另外torch.compile在DDP下确实有坑,我遇到过它和梯度同步的交互问题,尤其是用inductor后端时,前几百步会有诡异的loss spike,建议先关掉compile跑几轮对比下。还有个细节是DDP的broadcast参数初始化在不同卡上可能有微小差异,你可以检查下是否用了seed统一。如果显存实在不够,可以考虑把lr降到1e-6并配合cosine schedule,让模型在前10%步数慢热适应。对了,你用的模型是否做了梯度裁剪?对大模型训练来说,max_grad_norm设到1.0对抑制突然跳高很管用。
这问题我最近也遇到过,13B模型上DDP确实容易炸,特别是前几百步loss跳高那个现象,我怀疑跟梯度噪声和同步延迟有关。你每卡batch size才2,全局有效batch其实只有64(4卡28累积),对于13B来说确实偏小了,模型参数多,梯度方差大,多卡同步时会放大这种不稳定性。建议试试把每卡batch size提到4,梯度累积降到4步,这样全局batch能到128,可能会稳一些。另外torch.compile在2.0版本跟DDP配合偶尔会有诡异的数值问题,我关掉compile之后loss震荡明显减轻了,你可以先不加compile跑跑看。还有warmup步数可以拉长到总步数的10%左右,学习率3e-6其实不算太低,但可以试试用cosine调度加一个很低的最终学习率。最后检查一下数据加载有没有shuffle不一致的问题,有些情况下不同卡的数据顺序差异太大也会导致初期loss抖动。
之前用DDP训7B也遇到过类似情况,后来发现是梯度累积和分布式通信的交互出了问题。你试试把梯度累积改成用DDP的no_sync上下文手动控制,等累积完再同步,这样能减少通信频率,可能会稳一些。
另外torch.compile在大模型下有时会改变算子融合顺序,导致数值行为不一致,可以先用纯eager模式跑几百步对比一下。13B配4卡确实batch偏小,有效batch才64,建议要么增大每卡batch,要么用gradient checkpointing换显存来提升单卡batch。
还有个坑是不同卡上的数据分布不均,特别是LLaMA的padding策略,检查下DataLoader里是否用了distributed sampler,并且shuffle时seed要固定。如果前几百步震荡剧烈,可能是初始loss在下降过程中经过了一些陡峭区域,也可以试试更大warmup步数,比如到总步数的10%。
看到你说torch.compile我第一反应也是这个,2.0的编译模式在多卡下有时候会跟DDP的梯度同步打架,尤其是前几百步loss跳高然后又降,挺像梯度累积和同步时机不对导致的。13B这个规模确实对有效batch size更敏感,你4卡乘2乘8算下来其实才64,对13B来说还是偏小,可以考虑试试把梯度累积提到16或者每卡batch再往上挤一点。另外建议先关掉compile跑个几百步对比一下,如果稳定了那就基本锁定是它的问题,别急着动学习率。
我之前也遇到过类似情况,DDP下loss震荡有时候是梯度同步的延迟问题,特别是模型大了以后,通信开销会放大噪声。你可以试试关掉torch.compile看有没有改善,它虽然提速但有时会改变数值行为。另外13B配4卡,每卡2的batch确实偏小了,等效batch才64,建议把梯度累积加到16步试试,或者干脆用梯度裁剪,能压住那种突然跳高的尖峰。还有个思路是检查一下不同卡的随机种子,数据shuffle不一致也会导致前几百步特别不稳。
之前用DDP跑7B也遇到过类似的loss震荡,后来排查发现大概率不是torch.compile的锅,而是梯度累积和DDP的梯度同步在交互时出了问题。你梯度累积8步的话,注意要在最后一步才做all-reduce,如果用的是PyTorch 2.0的GradScaler或者NoScaler,加上compile的图优化,有时候会把累积逻辑搞乱,建议先关掉compile纯DDP跑一下对比看看。另外13B在4卡上每卡batch=2,全局有效batch=64(2×4×8),这个size对13B来说其实偏小,尤其是微调阶段,loss震荡很可能就是有效batch不够导致的,可以试试把梯度累积提到16步,或者干脆每卡batch加到3-4(如果显存够)。还有一个容易忽略的点:warmup步数要按全局batch算,你原来是按单卡还是多卡算的?如果warmup只覆盖了前几百步,那正好和你震荡区间重合,建议把warmup拉长到总步数的5%-10%。最后建议把loss打印改成滑动平均,别直接看原始值,前几百步原始loss跳高有时候只是个别样本的极端梯度,不代表整体发散,你可以用EMA平滑一下再判断。
我之前微调7B模型的时候也碰到过类似的震荡,后来发现罪魁祸首是梯度累积和DDP的梯度同步顺序在torch 2.0下变了。你累积8步再同步,但DDP默认是每步都做all-reduce,如果没设no_sync,那累积的梯度其实已经被平均过好几次了,等效于学习率被放大,前期肯定不稳。建议试试把梯度累积改成用no_sync包住前7步,只在最后一步同步,或者干脆减小累积步数。另外13B在4卡V100上每卡batch=2,全局batch才8,对这么大模型确实偏小,loss跳高可能是某些batch里出现异常样本,梯度clipping设个1.0能压住。torch.compile的话,我建议先关掉试试,它有时候会改变算子融合顺序,和DDP的bucket通信重叠,导致数值行为和单卡不一致。还有个细节,你用的是LLaMA的话,注意attention的scale有没有因为序列长度变化而调整,多卡下每个rank的padding可能不一样,也会引入噪声。总之先跑个50步,把梯度范数打印出来看看,如果震荡时梯度范数突然爆炸,那就是累积和同步的问题,而不是优化器参数的事。
我之前调7B模型的时候也遇到过类似情况,后来发现是torch.compile和DDP的梯度all-reduce之间有个隐藏的buffer问题,关掉compile或者把gradient_as_bucket_view打开会稳很多。另外13B配4卡确实有点吃紧,batch size等效才64,建议试试gradient checkpointing把batch再顶上去,或者直接上ZeRO stage 2,能明显缓解loss尖刺。还有个小细节,确认下dataloader的shuffle是不是每轮都重置了,多卡下数据顺序不一致也会导致前几百步乱跳。
之前微调7B的时候也遇到过类似情况,后来发现是梯度累积和DDP的梯度同步顺序没对齐,尤其是用了torch.compile之后,计算图和梯度归并的时机变了,建议先关掉compile试试。另外13B在4卡上每卡batch=2确实偏小,BN或者LayerNorm的统计量在多卡间会抖,可以试试把梯度累积改成每卡batch=4、累积4步,或者干脆上gradient checkpointing换显存来加大真实batch。还有个坑是warmup步数要跟着总batch数重新算,你从1.3B直接迁移到13B,warmup可能相对变短了,前几百步震荡大概率跟这个有关。
我之前调7B的时候也遇到过类似情况,后来发现是不同卡上的数据分布不均导致的,尤其是文本长度差异大的时候,DDP的梯度同步会特别敏感。你试试把每个batch的数据按长度做个排序,或者用DistributedSampler加个shuffle seed固定,loss震荡会小很多。另外torch.compile在这种规模下确实可能引入数值抖动,可以先关掉它跑几百步对比一下。
还有一点,13B用4卡其实显存挺吃紧的,梯度累积8步加上DDP的all-reduce,等效batch size才64,对于这么大模型来说确实偏小。我后来把学习率调到2e-6,同时把warmup步数拉到总步数的10%,震荡才慢慢收敛。你也可以观察下是不是特定某个step出现跳高,如果是,大概率是某张卡上出现了异常样本。
说实话,前几百步震荡有时候是正常的,尤其是从随机初始化或者预训练权重过渡的时候。你可以试着把log间隔调大一点,比如每50步打一次loss,别被短期的波动吓到。如果持续到1000步以后还这样,再考虑调模型架构或者优化器参数。
这batch size确实有点小,13B模型梯度噪声大,试试把梯度累积提到16步或者直接上gradient checkpointing。
之前也遇到过类似情况,后来发现是不同卡上数据分布不均,建议检查下sampler的shuffle逻辑,加个随机种子试试。
我之前在调多卡的时候也碰到过类似情况,尤其是刚上DDP那会儿,loss跟心电图似的。后来我发现一个坑:torch.compile在2.0里跟DDP的梯度all-reduce顺序有时候会打架,特别是用了find_unused_parameters=True的时候,会导致某些层的梯度没同步干净,表现出来就是loss突然跳一下。你可以先试下把torch.compile关掉,纯DDP跑几百步看看曲线是不是稳了,如果稳了那就是编译和图优化的问题。
另外13B这个规模,4卡每卡batch=2,算上梯度累积8步,其实等效batch也就64,对13B来说确实偏小了。我体感上13B微调,等效batch至少得128起步才比较稳,不然前几百步loss很容易在高原期来回横跳。你可以试试把梯度累积提到16步,或者干脆把每卡batch提到4(如果显存顶得住的话)。
还有个容易被忽略的点:warmup的步数。你说加了warmup但效果不明显,我猜你可能只加了200步左右?13B这种大模型,前几百步loss震荡有时候是正常的参数空间重排,warmup要拉到总步数的5%-10%才够,比如总共训练3000步,warmup至少150-300步,不然等于没加。另外建议你把梯度裁剪设到1.0,有时候个别batch的异常梯度会通过all-reduce放大,clip一下能压住那种突然跳高的尖峰。
我之前在DDP上调7B也遇到过类似情况,后来发现大概率是梯度累积和allreduce的交互出了问题。你梯度累积8步的话,默认行为是每步都同步梯度,等效batch size其实没变大,建议用no_sync包裹累积区间,只在最后一步同步。另外13B模型在4卡上单卡batch=2确实偏小了,BN统计和梯度噪声都会被放大,试试把gradient accumulation提到16步,或者干脆用gradient checkpointing换更大batch。torch.compile在DDP下偶尔会引入数值抖动,可以先关掉对比一下。
我之前微调7B的时候也碰到过类似情况,单卡稳如老狗,一上DDP就开始抽风。后来排查了一圈,发现torch.compile跟DDP的梯度同步在某些版本下有兼容性问题,尤其是2.0刚出那会儿,建议你先试试把compile关掉跑几百步对比一下,排除这个因素。另外你gradient accumulation设了8,但DDP里每个step的梯度本来就是all-reduce过的,accumulation会放大不同rank之间的噪声差异,尤其是batch size小的时候,13B模型本身梯度方差就大,前几百步震荡未必是bug,可能是模型在找比较陡峭的loss landscape里的局部结构。我当初把实际有效batch size从16提到64(通过加卡或者换大batch)之后,震荡幅度明显小了很多,感觉13B这种规模至少得保持全局batch在64以上才比较稳。还有个容易忽略的坑是不同卡的数据shuffle顺序不一致,导致每个step各个rank的梯度方向差异大,建议检查一下dataloader的seed设置,确保每个rank的shuffle是独立的但分布均匀。最后你warmup确实加了,但3e-6的学习率对13B微调可能还是偏高,试过降到1e-6甚至更低吗?我最后是配合梯度裁剪(max_norm=1.0)才把跳高现象压下去的,你可以试试看。
我之前跑7B也遇到过类似情况,尤其是前几百步loss跳高特别明显,后来发现是梯度累积和DDP的bucket通信没对齐导致的,你可以试试把gradient_accumulation_steps放到DDP外面手动做,别依赖模型内部的累积逻辑。另外torch.compile在V100上可能对某些算子有兼容性问题,建议先关掉compile纯DDP跑几百步对比一下,能快速排除这个变量。13B配4卡每卡batch=2确实偏小,BN或者LayerNorm的统计量在卡间不同步会加剧震荡,可以考虑把batch size提到每卡4或者用gradient checkpointing省显存换更大batch。还有个细节是warmup步数要跟着总batch size走,你调低lr但没提warmup步数,如果是按单卡算的,多卡下有效batch变大,warmup也得相应加长。