最近在公司用单机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震荡不收敛是什么情况?
全部回复
共 18 条大概率是梯度不同步那批大loss把优化器状态带崩了,试试梯度裁剪调小点加梯度累积看看。
感觉跟LoRA关系不大,先检查下DDP里bn层或者梯度规约是不是有精度问题,换all_reduce试试。
试试把梯度裁剪开到1.0,还有确认下是不是用了gradient accumulation导致DDP梯度没同步干净。
我之前也踩过类似的坑,不过不是LoRA,是full finetune。你单卡正常但DDP震荡,感觉大概率不是同步时机的问题,PyTorch的DDP梯度all-reduce是后端自动做的,时机上基本靠谱。我当时的排查方向是loss scaling和混合精度,A100上如果开了AMP,DDP下每个rank的grad scale可能没同步好,尤其是你用bf16的话,偶尔爆一下到几十特别像这个。你可以试试把amp关掉纯fp32跑几百步,如果稳了,那就是精度策略的问题。
另外你提到总batch从4涨到32,这个变化太大了,即便按线性缩放调lr,7B模型上LoRA的适应速度也可能跟不上。我试过把lr调成原来的1/4甚至1/8,配合更长的warmup,震荡会缓和很多。还有个细节,DistributedSampler虽然shuffle=True,但默认每个epoch的shuffle种子是一样的,如果不同rank的shuffle顺序相关性太强,也会导致梯度方向不一致,你可以显式给sampler设个不同的seed试试。
至于梯度裁剪,DDP下clip确实容易忽略,但一般只影响收敛速度,不太会造成前几百步完全降不下去。我怀疑还有个隐藏点:你的LoRA权重初始化是不是所有rank都一样?如果每张卡加载的base model权重有细微差异,DDP的梯度同步会把噪声放大。建议你确认下每卡加载的checkpoint完全一致,最好打印下第一层梯度norm对比下各rank。另外你用的是单机8卡,可以试试把NCCL的ring-allreduce改成tree模式,有时候网络拓扑会影响同步稳定性。
我之前也踩过类似的坑,单卡正常但DDP一上就崩,后来发现是梯度累积和all_reduce的时机没对齐,特别是LoRA这种只更新部分参数的情况,建议先关掉梯度裁剪试试,或者把clip值设大一点,比如从1.0提到5.0,看loss爆炸是否缓解。另外你检查过不同卡上的数据分布吗,DistributedSampler虽然shuffle了,但如果每个rank的batch内样本相关性太强,也会导致梯度方向不一致,可以试试把每卡的batch size减小到2,总batch保持16,对比一下震荡幅度。还有一个思路是暂时把BN换成LN或者固定住,虽然LoRA一般不动backbone,但某些实现里还是会更新LayerNorm的参数,这种细微差异在多卡下会被放大。
遇到过类似情况,最后发现是gradient clipping没设对,DDP下梯度norm是跨卡全局的,你单卡能用的阈值在32卡batch下会被放大,建议把clip值按卡数缩放试试。另外LoRA本身没问题,但要注意all_reduce的时机,如果用了gradient accumulation,得确保accumulation step和DDP的hook对齐,不然梯度会叠出问题。还有个小坑,DistributedSampler虽然shuffle了,但每个epoch的seed要手动设一致,否则不同卡的数据顺序不一致也可能导致震荡。
检查下不同卡的loss是否一致,大概率是某个rank数据没打乱或梯度没同步,先排除bug再调超参。
查一下梯度累积有没有开,DDP默认allreduce是同步的,LoRA一般没坑,但总batch变大后lr可能还得再降。
试下把梯度裁剪开到1.0再配合梯度累积,DDP下梯度噪声比单卡大很多,LoRA这种低秩更新更容易被带偏。
遇到过类似情况,当时查了半天发现是DDP默认的梯度all-reduce是在每个step反向传播后同步的,但如果你用了梯度累积或者某些优化器状态没正确广播,会导致各卡梯度不一致。另外LoRA的适配器参数初始化的scale很小,DDP下梯度噪声会被放大,建议试试把梯度裁剪的max_norm调到0.5或者更低,同时确认一下是不是用了find_unused_parameters=True,有时候这参数会导致某些层梯度同步异常。
大概率是梯度不同步或累积异常,试试把梯度裁剪设小一点,或者用all_reduce手动验证下梯度值。
DDP默认梯度是异步通信的,LoRA参数少但同步开销大,建议关掉bucket_cap_mb或改成gradient_as_bucket_view=True。
遇到过类似的坑,先别急着怀疑梯度同步,DDP的梯度allreduce本身在正常情况下是没问题的。你单卡batch size=4,DDP总batch=32,这相当于把有效学习率放大了8倍,虽然你按线性缩放调了lr,但warmup和优化器状态(比如Adam的动量)在分布式下的行为其实和单卡不完全一致,尤其是LoRA这种只训少量参数的情况,容易对噪声更敏感。建议你先试试把每卡batch size降到2,保持总batch不变,看震荡是否缓解,这能帮你区分是优化器问题还是梯度同步问题。另外,检查一下你是不是用了梯度累积,如果累积步数和DDP的梯度同步混在一起,很容易出现loss突然爆高的情况。
试试把梯度裁剪开到1.0,DDP下不同卡的梯度范数差异会放大,尤其LoRA层数深的时候。
建议查一下不同卡上的数据分布是否均匀,DDP的梯度是全局平均,单卡正常但多卡震荡很可能是采样器或数据shuffle导致个别卡batch差异过大。
LoRA本身不影响DDP,问题大概率出在总batch变大后学习率没配合适,试试把lr降到原来的1/4再跑几百步看看。
我之前跑DDP也遇到过类似情况,后来发现是梯度裁剪的时机问题——DDP里梯度同步是在backward之后自动做的,但如果你在optimizer.step之前手动裁剪,要确保用的是all_reduce之后的梯度,否则每卡裁剪的scale不一致就会震荡。另外LoRA这边有个坑,就是只对部分参数更新,DDP默认会对所有梯度做同步,虽然不影响正确性但可能引入额外噪声,试试把DDP的gradient_as_bucket_view打开,或者干脆把未更新的参数设成requires_grad=False看下。还有,你确认下是不是用了混合精度,AMP在DDP下loss缩放因子同步不对也会爆,可以先纯FP32跑几百步排除这个因素。
遇到过类似的坑,不过我们是在多机训练的时候。你单卡正常但DDP震荡,我感觉大概率不是LoRA本身的问题,而是梯度同步或者loss计算方式在分布式下变了味。一个很容易忽略的点是,DDP默认会把梯度做all-reduce平均,但如果你在loss里已经手动做了mean,那相当于梯度被额外缩放了一个batch size倍数,学习率再线性缩放就叠buff了,前期很容易炸。你可以先确认下loss里是不是用了reduction='mean',如果是,DDP下改成sum或者在外面统一除以总batch数试试。
另外你说用DistributedSampler且shuffle=True,这个我踩过坑——每个epoch开始前,sampler的set_epoch必须调用,否则所有卡在每个epoch拿到的shuffle顺序完全一样,等于变相放大了局部数据分布差异,loss震荡会更明显。可以在train loop里检查下有没有在每个epoch开头对sampler做set_epoch(epoch)。
还有个思路,你提到梯度裁剪,但没具体说怎么设的。DDP下如果clip_norm设置太小,特别是前期LoRA参数初始化和主模型scale不匹配,梯度范数本身就波动大,clip反而会加剧震荡。建议先不裁剪,或者把clip值调大比如1.0以上,观察前几百步的梯度范数曲线,看是不是有周期性尖峰。
最后,如果你用的是BF16混合精度,可以试试纯FP32跑几百步对比下,有时候AMP的loss scaling在DDP下同步时机不对也会导致诡异震荡。以前我们排查过,最后发现是梯度累积和DDP的bucket划分冲突,你如果有梯度累积,记得确保累积完再同步。先试试这几个点,大概率能定位到。
大概率是梯度不同步或BN统计量问题,试试设find_unused_parameters=False加梯度裁剪,LoRA本身没分布式坑。
我之前也踩过类似的坑,单卡正常DDP就崩,最后发现是梯度累积和allreduce的交互出了问题。你试试把梯度裁剪关掉,或者把clip值调大点,有时候DDP下梯度范数本身就会比单卡大不少。另外LoRA在分布式下有个小坑,就是lora的A矩阵初始化是0,B矩阵随机,如果不同卡的初始化不一致(虽然理论上seed一致),但保险起见检查下每卡的seed设置。还有,你确认下是不是用的find_unused_parameters=True,有时候某些层没参与更新会引发奇怪的梯度行为。最后实在不行,可以先固定总batch size不变,用梯度累积模拟DDP,对比下是不是纯数据并行导致的差异。
这问题我熟,之前跑13B也踩过坑。你试过把总batch size固定成32,但每卡batch size改成8(也就是只开4卡)验证下吗?我怀疑是梯度同步前梯度范数差异太大,DDP allreduce之后等效batch变大,但LoRA的秩低导致参数更新对batch敏感,试试梯度裁剪设个1.0,或者干脆用fp16混合精度跑,能缓解不少。
另外你确认下是不是用了梯度累积,DDP里梯度累积要手动同步,不然会重复累加。还有,LoRA的A矩阵初始化是高斯,B矩阵是零,分布式下每卡初始化一样吗?最好把seed在每卡设成不同的,不然同步后梯度方向会特别诡异。