最近在公司用单机8卡A100跑一个7B的LoRA微调任务,用的是PyTorch自带的DistributedDataParallel。单卡(batch size=4)跑同样的数据和超参,loss能正常下降,但一上DDP(每卡batch size=4,总batch 32)loss就剧烈震荡,前几百步完全降不下去,偶尔还会爆一下到几十。我已经排查了数据加载(用的DistributedSampler,shuffle=True)、学习率也按线性缩放规则调了,甚至试过不同warmup步数,都没啥改善。想问问各位大佬,是不是DDP下梯度同步的时机或者梯度裁剪设置有问题?还是说LoRA本身在分布式下有什么特殊坑?另外,我看有些代码里会设置find_unused_parameters=True,这个会影响吗?恳请指点方向,谢谢!
PyTorch多卡DDP训练大模型,loss震荡不收敛是什么情况?
全部回复
共 29 条试试把梯度裁剪设到1.0,DDP下不同卡的loss波动会被放大,单卡不明显而已。
检查下是不是每张卡的batch size太小导致BN统计不稳,LoRA层多卡同步下梯度噪声也容易炸。
单卡能降DDP震荡,先别急着怀疑LoRA,我遇到过类似情况最后是卡间数据分布不均导致的——虽然用了DistributedSampler,但如果你tokenizer或者采样器没设seed,不同卡拿到的数据顺序其实不一样,极端batch下梯度方向冲突会特别大。另外你试过把总batchsize固定成8(每卡1)对比一下吗?如果这样不震荡,那大概率是lr scaling那里还要按sqrt调而不是线性。梯度裁剪我建议先设个1.0看看,DDP下梯度范数本来就比单卡大不少,不裁的话偶尔爆一下很正常。
我之前跑类似任务也踩过这个坑,DDP下loss震荡大概率不是LoRA本身的问题,而是梯度同步时allreduce把不同卡上的gradient noise也平均进去了,小batch时单卡梯度方向比较稳,总batch变大后反而容易在局部震荡区打转。你试过把每卡batch size调大点、同时减少卡数吗,比如4卡每卡batch 8,总batch还是32,这样每卡梯度更平滑些。另外梯度裁剪的阈值在DDP下可能需要按总batch的倍数放宽,不然个别卡的异常梯度会把全局拉爆。还有个细节,DistributedSampler虽然shuffle了,但如果你没在每次epoch开始前调用set_epoch,不同卡的样本顺序可能一直固定,导致局部数据分布差异被放大。
之前跑类似任务也踩过这个坑,排查了一圈发现根因不在DDP本身,而是总batch变大后BN层或者LayerNorm的行为变了。LoRA虽然只训adaptor,但如果你冻结的主干里有BN,DDP下每卡独立统计就没问题,可一旦用了SyncBN,统计噪声会被放大,loss就容易抽风。你单卡能收敛但DDP不行,建议先确认一下模型里有没有BN,有的话换掉或者干脆别开SyncBN试试。
另外你提的梯度同步时机,其实DDP是每个step反向传播完就做all-reduce,理论上没问题,但有个容易忽略的点是梯度裁剪的阈值是不是也跟着总batch缩放了。你单卡剪裁阈值如果是按单卡loss量级设的,DDP下梯度范数会随batch增大而变大,固定阈值可能直接裁没了有效更新,反而震荡更凶。可以试着把梯度裁剪关掉跑几百步看趋势,或者把阈值按总batch的平方根倍率放大。
还有个更阴间的可能,就是DistributedSampler加shuffle=True时,如果每个epoch没有正确调用set_epoch,数据顺序在每个rank上是固定的,但不同rank间会重复采样,导致模型反复看同一批子集,loss自然不收敛。你检查下训练循环里有没有在每个epoch开始前重新设置sampler的epoch,这个细节特别容易漏。
LoRA在分布式下的特殊点我倒是觉得不大,因为adaptor本身参数少,通信量低,反而是你总batch从4到32后,学习率线性缩放虽然做了,但warmup步数可能不够。7B模型用AdamW的话,warmup一般要占总步数的1%到3%,你如果总步数本来就不多,几百步warmup根本不够稳定优化器状态。可以先试试把warmup拉长到总步数的5%,或者干脆用warmup加cosine decay,别用固定学习率衰减。
如果上面都试了还不行,那就查一下DDP的gradient accumulation有没有跟DistributedSampler冲突。你是不是实际用了累积梯度,但sampler的drop_last没设对,导致不同rank的batch数量不一致,最后几步梯度同步时出现nan或者异常值。这问题我遇到过,用单卡看不出来,多卡必炸。
之前跑多卡DDP也遇到过类似情况,后来发现是梯度累积和同步的交互问题——如果用了梯度累积,要确保只在累积到指定步数后才做all-reduce,否则梯度会部分更新导致震荡。另外LoRA的rank如果设得比较小,分布式下不同卡上的低秩矩阵初始化差异会被放大,可以试试固定种子并检查一下每卡初始化是否一致。还有个小点:你确认一下DDP里用的是不是同一个学习率调度器实例,有些写法会让每张卡各自调度,步数错位也会爆loss。如果还不行,先跑个几百步用梯度日志对比单卡和DDP的梯度范数,看是不是同步前就已经有outlier了。
DDP下loss震荡我也踩过,大概率不是梯度同步的锅,而是总batch从4跳到32后优化动态变了。你按线性缩放调了lr,但LoRA的A/B矩阵初始化方差很小,大batch下早期梯度信噪比反而更差,warmup不够就容易炸。建议试试把每卡batch降到2、gradient accumulation补回来,同时确认下DDP的gradient_as_bucket_view和find_unused_parameters设置,LoRA冻结层多,这个坑挺常见的。
试试把梯度裁剪放到backward之后、optimizer.step之前,DDP下各卡梯度是平均的,裁剪阈值得跟着总batch重新调。
总batch大了32倍,学习率不能只按线性缩,LoRA的A/B矩阵初始化对同步很敏感,试试调小点再加梯度裁剪。
你提到单卡能正常降、DDP一上就炸,这基本可以排除数据和超参本身的问题,方向还是在分布式同步这块。有个容易被忽略的点是DDP默认在每次backward时做all-reduce,如果你用了梯度累积或者在LoRA这种只有部分参数参与反传的场景下,没设置find_unused_parameters=True,梯度同步可能会出问题,loss震荡就很正常了。还有个坑是学习率缩放,你从单卡batch4变成总batch32,如果按线性缩放把lr放大8倍,对LoRA这种低秩适配其实很容易过大,建议先别缩放,保持原lr跑一下看曲线。梯度裁剪的时机也值得检查,DDP下每个rank算的是局部梯度,裁剪要放在all-reduce之后才等价于全局裁剪,否则各卡裁的各的,效果会打折扣。另外A100上可以试试把gradient_as_bucket_view打开,减少显存拷贝带来的数值抖动,虽然这个更多是性能问题。爆到几十那种大概率是某个step梯度爆炸,可以加个torch.nn.utils.clip_grad_norm_配合监控每步grad norm看看是哪层在作妖。