最近在搞一个多智能体协作的AI Agent项目,环境是PettingZoo的MAMujoco,策略用的是MAPPO。单机单卡跑小规模(2个agent)还没啥问题,但一上4个agent并用torch.distributed做多进程训练,就疯狂报“NCCL通信超时”,有时候还莫名其妙内存溢出。我试过调大timeout和减小batch size,但跑个几百步就卡死。有没有老哥踩过类似的坑?是环境同步的问题还是PyTorch分布式策略没配好?求指点,孩子快调吐了。
用PyTorch搭多智能体强化学习,分布式训练总是报错,求老哥支招
全部回复
共 179 条先查一下PettingZoo版本和gym是不是不兼容,我之前就是这俩版本冲突卡到怀疑人生。
试试把每个agent的obs和action统一pad到固定长度再进分布式,环境返回变长tensor经常把NCCL搞挂。
这问题太典型了,多半不是NCCL本身的事儿,而是PettingZoo环境在子进程里没做好同步。我当初搞MAMujoco也这样,后来发现是环境reset和step里用了全局随机数,多进程下状态就乱了,导致某个agent的observation维度对不上,通信就卡死。建议你先检查下每个worker是不是独立的环境实例,别共享任何全局变量,另外试试把torch.distributed的backend换成gloo跑几步,能跑通就说明是NCCL和PettingZoo的兼容性在搞鬼。
MAMujoco这环境本身就容易卡在同步上,尤其是agent多了之后,PettingZoo的并行环境跟torch.distributed的进程组配合经常出幺蛾子。我之前也遇到过NCCL超时,后来是把每个agent的rollout单独放一个进程,再用共享内存传数据,绕开了多进程直接跑env,稳定不少。你试试把环境实例化和模型更新彻底分开,别让NCCL等env的step,大概率能缓解。另外内存溢出看看是不是每个进程都复制了完整的环境池,用fork启动方式有时候会省很多。
NCCL超时八成跟MAMujoco的步进同步有关,PettingZoo老版本在多进程下环境状态复制容易卡死,先检查一下是不是每个worker独立reset了环境。内存溢出那个,试试把replay buffer放到共享内存或者干脆关掉,MAPPO用不着存那么多transition。另外torch.distributed别用默认的nccl后端,换gloo跑跑看能不能排除通信问题。之前我踩坑是发现子进程里重复加载了模型权重,导致显存爆掉,你查查是不是有隐式clone。
碰到过类似的,NCCL超时大概率不是单一原因,你先试试把MAMujoco的vector env改成单进程串行采集,再用shared memory传数据,能排除掉环境步调不一致的坑。内存溢出那个,检查下是不是每个agent的obs拼接时没释放旧tensor,用del和torch.cuda.empty_cache()手动清一下。另外4个agent的分布式建议直接用torchrun--nproc_per_node=4启动,别自己写mp.spawn,之前我这么改完稳定很多。你用的PyTorch版本是2.1以上吗?不是的话先升个级,老版本NCCL在PettingZoo这种动态action space下容易出问题。
这问题太典型了,MAMujoco本身步长就不短,4个agent一上,NCCL通信量翻倍,timeout调再大也容易撞上同步瓶颈。我之前也卡在这,后来发现是PettingZoo的env.reset和step里藏着全局状态同步,多进程下每个worker的obs维度偶尔不一致,直接导致通信卡死。建议你先在每步step后打印一下各agent的obs shape和reward,确认不是环境侧数据不同步,再去看torch.distributed的backend配置。另外,内存溢出多半是replay buffer或gradient accumulation没按进程数切分,试试把每个进程的batch再减半,同时用torch.cuda.empty_cache()手动清一下显存。要是还不行,可以考虑换成ray的rllib,它内置的MAPPO对多agent支持更省心,至少不用自己手搓分布式细节。
PyTorch的NCCL超时在MAPPO这种需要频繁同步的负载下太经典了,尤其PettingZoo的MAMujoco每个agent的observation和action维度还不一样,多进程下数据打包和梯度通信很容易卡在某个barrier上。我之前遇到过类似情况,后来发现不是NCCL本身的问题,而是子进程里env.step()的耗时方差太大,一个agent掉队整组等它,timeout设再大也没用,你可以先给每个进程单独打日志看step时间。另外内存溢出八成是replay buffer或者trajectory存储没按进程隔离,每个worker都在攒全量数据,试试把GAE的计算放到collect之后统一做,别在采样循环里存太多中间变量。还有个小坑,MAMujoco的物理仿真默认开多线程,跟torch.distributed的进程数叠一起会抢CPU,导致通信线程被饿死,你可以在env创建时把mj_env的num_threads设成1。如果还卡死,建议把NCCL换成GLOO后端先验证逻辑,虽然慢但至少能定位是不是通信原语的问题。最后检查一下你是不是在spawn之后才初始化env,这会导致每个进程重新加载模型权重,显存翻倍,改成fork方式或者共享权重能缓解。
大概率是PettingZoo环境没做好多进程隔离,试试把env建在子进程里,别共享主进程的CUDA上下文。
看到这个我直接DNA动了,之前搞MAMujoco也差点被NCCL整崩溃。你这大概率不是环境同步的问题,而是torch.distributed默认的NCCL backend在PettingZoo这种带GIL锁和异步reset的环境里特别容易死锁,尤其是多agent共享一个进程时,通信组和env step的时序对不上就会超时。我后来是把所有agent的rollout收集和梯度同步都拆到独立的线程里,主进程只做env.step,然后给每个子进程单独设了torch.set_num_threads(1),内存溢出也跟这个有关,因为MAMujoco的observation空间会随agent数暴涨,你试过给每个进程设独立的shared memory buffer吗?另外建议把NCCL换成gloo先验证逻辑,虽然慢但能排除通信问题。还有个小坑,PettingZoo的parallel_env默认会做全局同步,4个agent时每个step的reset开销会翻倍,你可以试试手动把terminated和truncated分开处理,别让env自动重置。最后,如果还卡死,检查一下是不是MAPPO的value network输入拼接了全局state,那个在MAMujoco里维度是随agent数量线性增长的,显存爆炸很正常。
这个坑我太熟了,MAMujoco用PettingZoo的vector env时,多进程下每个子环境的状态空间如果不完全一致,NCCL很容易在某个step卡死,建议先确认一下agent的observation是否因为全局物理状态共享而出现隐式依赖,另外试试把torch.distributed换成gloo做通信后端,虽然慢点但能排查是不是NCCL本身的问题。内存溢出那个,大概率是PettingZoo的渲染缓冲没清干净,可以在每个episode结束手动调一下env.reset并gc.collect()。你用的MAPPO是用的共享critic吗?如果是的话,更新时同步全局state的shape对不对也很关键,我之前就是这里没对齐导致梯度同步超时。
试试把PettingZoo的vector env换成手动同步,MAPPO这玩意儿对数据新鲜度太敏感,NCCL超时多半是环境卡在reset上。
NCCL超时这个坑我太熟了,多半不是环境同步问题,而是MAMujoco的vector env在多进程下每个worker的reset或step没对齐,导致某些rank提前发完数据在那干等。你可以试试把PettingZoo的env包一层,在step后强制加个barrier同步,或者干脆把数据收集和训练拆成两个进程,用共享内存队列传experience,别让NCCL直接背锅。另外内存溢出很可能是每个进程都复制了一份完整的环境,4个agent的话试试用fork启动而不是spawn,能省不少内存。你用的是torch.distributed还是torch.multiprocessing?如果是前者,检查一下set_device是不是每个rank都设对了,之前我卡死就是GPU id没对应上。
遇到过一模一样的坑,MAMujoco这环境本身步进逻辑就有点问题,多agent同步的时候特别容易卡在某个环境的内部等待上,NCCL超时有时候根本不是你代码的锅。我后来是把PettingZoo的并行环境换成自己手写vectorized wrapper,每个子进程单独跑环境,然后用torch multiprocessing的queue传obs和action,绕开distributed那套集合通信,反而稳了很多。内存溢出的话,你检查下是不是每个进程都复制了一份完整的模型参数和优化器状态,尤其是PPO这种要存经验buffer的,建议把replay buffer挪到共享内存或者干脆分片存,别让每个进程都全量保留。另外你试试把NCCL的P2P level调低,设成NCCL_P2P_DISABLE=1,有些机器上能避免奇怪的同步死锁。还有个细节,MAMujoco的奖励计算是全局的,4个agent时梯度方差会大不少,你如果用的是MAPPO的centralized critic,记得把critic的输入做下normalization,不然loss容易炸导致训练停住。最后问下你用的是gymnasium还是老版gym的API?PettingZoo最近版本切换过接口,有些隐藏的时序问题会直接让分布式训练假死。
这问题我太熟了,MAMujoco的同步开销本来就大,4个agent用NCCL跑起来,通信量直接翻倍,timeout调大只是治标不治本。你大概率不是策略写错,而是环境步进和梯度更新之间的同步点没卡对,PettingZoo的并行环境默认用ray,跟torch.distributed的进程组混在一起容易死锁。我建议先试下把每个agent的env独立实例化,别共享一个ray remote,再用gloo跑CPU版debug,看能不能排除NCCL的干扰。另外内存溢出很可能是每个进程都复制了一份完整的环境状态,MAMujoco的obs维度又大,4个agent就是4倍显存,试试用shared memory或者把obs归一化放到CPU上算。还有个骚操作,把gradient accumulation打开,人为降低同步频率,虽然理论上会慢点,但能绕过很多诡异的卡死。你用的MAPPO是用的哪个开源实现?有些版本的老代码在multi-agent的data采集部分有hidden race condition,换一个最新的实现可能直接就好了。
NCCL超时这个坑我也踩过,八成不是环境同步的问题,而是MAMujoco的pettingzoo环境在子进程里没正确序列化,建议把环境创建放到每个worker的初始化函数里,别在全局搞。还有你试试把NCCL的GDRDMA关掉,有时候多机多卡反而没事,单机多卡会撞总线。内存溢出的话,大概率是replay buffer或者gae计算时把整个episode的tensor都堆显存了,改成增量式计算试试。实在不行就退回用subproc_vec_env那种同步采样,虽然慢点但稳定。
大概率是MAMujoco里子环境步进不同步把NCCL卡死了,试试把vector env的异步重置关掉,还有共享内存别开太大。
我最近也在搞类似的,不过用的是MAPPO+SMAC,感觉你这问题八成不是PyTorch本身,而是PettingZoo那套环境在多个进程里同步的时候出了岔子。MAMujoco的物理引擎本身就不太支持多进程并行,每个子进程各自跑环境,但全局状态没对齐,NCCL那边等数据等不到就超时了。
我建议你先别急着调timeout,先检查一下是不是每个进程里都重复初始化了环境,或者用了同一个随机种子导致步数不一致。另外,内存溢出很可能是PettingZoo的渲染buffer没释放,尤其是多agent的时候,每个环境都在攒obs,你不显式清的话,跑几百步就爆了。
我之前踩过一个坑,就是torch.distributed的初始化方式,如果你用的是spawn而不是launch,子进程里环境创建顺序和主进程不一致,也会卡死。你可以试试把所有环境初始化放在main函数里,然后只把模型参数传进去,别让环境跟分布式组网纠缠在一起。
还有个小技巧,如果你用gloo后端做CPU同步,NCCL只用来传梯度,能缓解不少通信压力,虽然慢点但至少不崩。你那个卡死是卡在loss backward还是collective call?如果是collective call,大概率是rank不同步,建议在每个step结束加个torch.distributed.barrier()强制对齐一下。
这问题我熟,之前搞MADDPG也差点被NCCL搞疯。你试试先把PettingZoo的env wrap成vectorized,然后每个进程只跑一个agent的rollout,别共享环境实例,感觉你八成是环境状态同步没做好。另外MAMujoco的action space维度差异大,分布式下gradient all-reduce容易堵,可以试试把gradient clip调小点,或者换GradientPipeline。内存溢出那个,看看是不是每个进程都复制了完整环境,用fork启动方式能省不少内存。
大概率是PettingZoo环境状态同步的锅,多agent下得把env实例也放进distributed的device_map里。
我上次是把所有agent的obs先集中到rank0再广播,虽然慢点但稳了,试试。
这问题我太熟了,MAPPO加PettingZoo的坑基本都在环境步进和分布式 rollout 不同步上。NCCL超时很多时候不是通信本身的问题,而是某个进程因为环境卡住没走到同步点,建议先把每个rank的env.step耗时打出来看看方差。内存溢出大概率是replay buffer或者gradient accumulation没按global batch size调,4个agent时每个进程的local batch要除以world_size。还有个野路子,把MAMujoco的vector env换成单进程串行采样,再配合gather做集中式训练,能绕开一半的通信问题。