最近在搞多卡训练,用PyTorch的DistributedDataParallel(DDP)跑一个BERT微调任务。单卡没问题,但换成4卡后,发现loss下降得特别慢,而且每个卡的loss值都不一样。我确认了数据加载用的是DistributedSampler,模型参数也是broadcast过的,但就是梯度不同步?检查了一下,发现reduce操作好像没生效,梯度还是各自为政。有人说是没设对init_method,我试了tcp://localhost:12355和env://都不行。求大佬指点一下,到底是什么环节漏了?我用的PyTorch 1.12,CUDA 11.6,机器是4张V100。
MCP里用PyTorch做分布式训练,DDP报错梯度不同步怎么办?
全部回复
共 45 条遇到过类似的情况,后来发现是batch size太小导致各卡计算出的梯度差异太大,DDP的allreduce反而放大了这种不稳定。可以试试把batch size调大或者用梯度累积,另外检查下是不是用了SyncBN,BERT里有些BN层默认不同步也会造成梯度不一致。还有个小坑,PyTorch 1.12的DDP在init_method上确实容易踩雷,可以试试把world_size和rank显式传进去,别全靠环境变量。
我之前用1.12也踩过这个坑,建议先检查一下world_size和rank是不是真的传对了,很多情况是初始化时进程组没完全建立。另外可以试试在backward之后手动加一步torch.distributed.all_reduce看看梯度是否同步,能快速定位是DDP本身的问题还是别的环节。还有你确认一下模型内部有没有BatchNorm之类的层,DDP下BN的同步模式默认是关闭的,得显式设置。
我也遇到过类似的问题,最后发现是没在模型forward之后手动调用reduce,DDP默认只同步梯度但不做all-reduce,得确认一下你是不是用了model.no_sync或者梯度累积没处理好。另外init_method用env://的话,记得检查环境变量RANK和WORLD_SIZE有没有正确设置,有时候是启动命令漏了这些参数。
遇到过类似情况,排查下来往往是DDP初始化时漏了设置环境变量,比如RANK和WORLD_SIZE没传对,或者在代码里手动调用了all_reduce导致的冲突。建议你检查一下启动命令是不是用torchrun或者torch.distributed.launch,以及模型里有没有自己写梯度累加逻辑。另外PyTorch 1.12在DDP的梯度同步上有个小bug,建议先升级到1.13或2.0试试看。
遇到过类似的问题,除了init_method,你检查一下DDP包装模型的时候有没有设置find_unused_parameters=True,BERT里有些层可能不会被所有卡用到,这个参数不设会导致梯度同步出问题。另外建议把torch.distributed.barrier()放在每个epoch开始前,确保所有进程同步到同一个起点。还有个小细节,确认一下你每个进程的batch_size是不是一样,这个也会影响reduce结果。