最近在魔改一个检测模型,想从单卡改成DDP多卡训练。代码参考了官方tutorial,但一跑就各种幺蛾子:先是dist.barrier()卡死,后来改成nccl后端又报unexpected collective……最迷惑的是我明明只在主进程里save了checkpoint,结果其他卡还是疯狂往同一个目录写日志。我猜是DistributedSampler和DataLoader的num_workers没对齐?但调了半天还是时好时坏。有没有老哥指点下DDP的正确打开姿势?顺便问下,model和optimizer到底该不该用torch.nn.parallel.DistributedDataParallel包多次?总感觉这里有个坑等着我踩。
PyTorch多卡训练DDP老报错,到底哪里没配对啊?
全部回复
共 23 条你提到的几个症状基本都踩在DDP的经典坑上了。dist.barrier()卡死通常是因为只有部分进程执行到了这里,比如你在主进程里加了判断,其他卡就永远等不到集合通信,得保证所有rank都调用同一个collective。unexpected collective多半是模型里有些分支只在rank0跑,或者forward里有依赖数据的条件逻辑,导致各卡执行顺序不一致。至于日志乱写,那是你只判断了rank==0来save,但logging或tensorboard的writer可能每个进程都建了,得把summary writer和文件句柄也包进if rank==0里。DistributedSampler记得每epoch调set_epoch,不然shuffle会失效,num_workers本身不需要和卡数对齐。model必须用DistributedDataParallel包,而且要在model.to(device)之后包,optimizer用普通的就行,不用DDP包。如果还时好时坏,建议加TORCH_DISTRIBUTED_DEBUG=DETAIL环境变量,能打出具体哪个collective对不上,比盲调快多了。
日志乱写是因为没设local_rank做device判断,save时加个rank==0就行。DDP那块model必须包,optimizer不用。
日志重复写大概率是没在主进程用torch.distributed.get_rank()==0包住logging或者writer,DDP不会帮你自动拦这个。barrier卡死常见原因是各卡进入次数不一致,比如有的卡在dataloader里挂太久或者异常提前退了。unexpected collective基本就是某张卡多跑或少跑了一次allreduce,检查下有没有if分支只在rank0执行却包含DDP通信。DistributedSampler记得每epoch调set_epoch,num_workers和persistent_workers配好一般问题不大。