最近在公司用单机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,这个会影响吗?恳请指点方向,谢谢!