最近在魔改一个检测模型,想从单卡改成DDP多卡训练。代码参考了官方tutorial,但一跑就各种幺蛾子:先是dist.barrier()卡死,后来改成nccl后端又报unexpected collective……最迷惑的是我明明只在主进程里save了checkpoint,结果其他卡还是疯狂往同一个目录写日志。我猜是DistributedSampler和DataLoader的num_workers没对齐?但调了半天还是时好时坏。有没有老哥指点下DDP的正确打开姿势?顺便问下,model和optimizer到底该不该用torch.nn.parallel.DistributedDataParallel包多次?总感觉这里有个坑等着我踩。
PyTorch多卡训练DDP老报错,到底哪里没配对啊?
全部回复
共 10 条踩过一样的坑,多半是sampler没跟dataloader绑对,还有save得用rank判断加进程锁。
讲真,你遇到的这几个坑我基本都踩过一遍,dist.barrier()卡死八成是卡在dataloader的worker数量不一致上,试试把num_workers设成0先排除问题。unexpected collective大概率是某个分支里不小心调用了collective操作但没所有进程都执行,检查下代码里有没有提前return或者条件判断。保存checkpoint的话,我习惯用torch.distributed.get_rank()判断然后只让rank0写,日志路径也带上rank就不会互相覆盖了。至于model和optimizer,我建议还是用DistributedDataParallel包一下模型,optimizer保持原样就行,但记得在load_state_dict前先module访问。另外强烈建议把torch.cuda.set_device(local_rank)放在最前面,这个漏了也会出各种玄学问题。
看到你说save只在主进程但日志还是乱写,大概率是logger初始化的时候没做rank判断,跟DDP本身关系不大。dist.barrier()卡死常见原因是你没把init_method里的rank和world_size跟实际启动的进程数对上,或者用了spawn后没传对rank参数。unexpected collective多半是某个分支里有的卡进了collective有的没进,检查下所有return和continue路径是不是都走到了barrier。DistributedSampler和num_workers其实不用严格对齐,但记得每个epoch调sampler.set_epoch(),不然shuffle失效,而且dataloader里pin_memory在多卡下最好开着。model和optimizer用DistributedDataParallel包model就行,optimizer保持原样,但注意torch.nn.parallel.DistributedDataParallel第一个参数要直接传model,别套Sequential啥的。
八成是sampler没给每个rank分好数据,试试把num_workers设成0排除干扰,还有logger记得按rank隔离。
说到DDP的坑我太有感触了,上个月调一个分割模型也是被这些collective通信搞到怀疑人生。你那个dist.barrier卡死的问题,八成是某个进程提前return了或者数据量不一致,导致所有rank没走到同一个同步点,这个得检查下代码里有没有隐形的分支。nccl报unexpected collective的话,大概率是模型里有sparse tensor或者自定义的通信操作,跟nccl的集合通信语义对不上,建议先全换成dense tensor试试。至于save checkpoint那个,虽然你只在主进程save了,但如果其他rank的DataLoader还在跑,日志写入可能是通过print或者logging触发的,跟save没啥关系,得把日志输出也限定在rank0。DistributedSampler和num_workers确实容易出问题,特别是每个rank的worker数量不一样的时候,会导致每个epoch的shuffle顺序错乱,进而引发后续的同步崩溃,建议统一num_workers并且设置seed。model和optimizer的话,DDP会自己处理梯度同步,所以model直接包一下就行,不用手动做别的,但optimizer必须在DDP包裹之后再创建,不然state_dict会乱。最后说个玄学,如果你用了混合精度或者梯度累积,记得把no_sync上下文管理器用在正确的地方,不然很容易出现梯度没同步完就开始下次前向的报错。
说到DDP这个坑我太有共鸣了,刚入坑时也是被dist.barrier()卡到怀疑人生,后来发现多半是进程数跟init_process_group里world_size没对上,或者有进程提前崩了没参与同步。你那个unexpected collective八成是代码里某个分支只在部分进程执行了集合通信操作,比如loss计算里带了个all_reduce但被if包住了,这种不对称最容易炸。关于日志那个问题,我猜你用了logging.FileHandler或者print重定向,但没检查rank,其实最稳的做法是给每个进程单独分配一个日志文件名后缀,或者干脆只在rank==0时初始化logger。DistributedSampler和num_workers其实没直接关系,但有个容易被忽略的点:DataLoader里shuffle必须设False,否则会和Sampler冲突,导致每个epoch数据分布错乱。至于model和optimizer,DistributedDataParallel只包model就行,optimizer保持原样,但save的时候要把model.module.state_dict()和optimizer.state_dict()一起存,加载时也得注意先load_state_dict再包DDP。我现在的习惯是写一个setup函数统一处理rank、device、seed,然后所有跟进程相关的操作都走一个if dist.get_rank() == 0的封装,基本能避免八成问题。另外你提到魔改检测模型,如果里面有自定义的forward里用了torch.where或者mask操作,记得检查这些张量是不是在所有进程都同步了,不然梯度会不一致。最后想问下,你nccl后端是用的gloo做备份了吗?有时候环境变量NCCL_DEBUG=INFO能直接告诉你卡在哪一步。
说实话你贴的这几个报错我基本都踩过,最后发现八成不是num_workers的问题,而是dist.barrier()之前有某个进程提前return了,或者数据集长度在各卡上不一致导致sampler算出来的索引对不上。nccl那个unexpected collective大概率是代码里某个分支只有部分进程执行了通信操作,比如在验证集上忘了加if dist.get_rank() == 0这种判断。
关于checkpoint保存,你只在主进程save是对的,但日志目录那个问题通常是因为你没在子进程里重新设置logging的file handler,所有rank共享了同一个文件描述符,疯狂写同一个文件。我一般会按rank拆目录,或者每个进程单独开一个log文件。
model和optimizer的包装方式,官方推荐是先把model放到gpu上,再包DistributedDataParallel,optimizer不需要额外处理,但要注意torch.load的时候得用map_location指定到对应rank的device,否则加载到cpu再搬回gpu容易出隐性bug。
另外DistributedSampler有个坑,就是每个epoch要手动调用set_epoch(epoch),不然每个epoch的shuffle结果一样,模型会过拟合到固定的batch顺序上。我之前漏了这一步,loss曲线看着正常,但验证集上一直不涨,查了半天。
最后建议你先把dist.barrier()全去掉试试,很多场景下其实不需要,反而容易卡死。真正需要同步的地方用all_reduce或者all_gather更稳妥。如果还不行,把torch.distributed.init_process_group的timeout参数调大点,默认30秒在小数据集上可能不够。
你这几个问题我基本都踩过,最坑的其实是dist.barrier()卡死,多半是某个进程提前退出了,比如DataLoader的worker崩了但主进程还在等。nccl报unexpected collective的话,检查下是不是有if dist.get_rank()==0包裹了不该包的通信操作,比如barrier或者all_reduce必须所有rank都执行。save checkpoint那个事儿,光靠主进程判断不够,得确保其他rank的logger和writer也只在主进程初始化,或者干脆把输出路径按rank分目录。另外DistributedSampler记得在每个epoch开头调set_epoch,不然shuffle会失效但不会报错,容易让你误判。model和optimizer不需要手动包Distributed,直接用DDP包model就行,optimizer保持原样,但注意DDP的gradient同步是自动的,别自己再去all_reduce梯度了。
看到这个我简直梦回上周,一模一样的问题,最后发现是sampler没传进dataloader,导致每个epoch的shuffle根本没生效。你那个barrier卡死八成是rank和world_size没设对,试试看环境变量LOCAL_RANK和全局RANK是不是搞混了。checkpoint那块建议只在rank0上save,但log可以先写到各自的临时目录最后再合并,不然多卡同时写一个文件必炸。model和optimizer就用DistributedDataParallel包一下就行,但注意BN层要转成SyncBN,不然精度会掉。
这种问题八成不是sampler的锅,你先检查下是不是所有进程的batch size没按总卡数翻倍,DDP里每个进程拿到的其实是单卡batch。日志乱写那个简单,直接把logging输出重定向到带rank的文件就行,别纠结是不是主进程的问题。还有model和optimizer都不用包DistributedDataParallel,optimizer只留主进程step就行,其他卡forward/backward完就等着。最后建议你把dist.barrier()全删了,DDP本身就带隐式同步,手动加反而容易卡死。