PyTorch分布式环境下rank1与rank2点对点通信报错求助
PyTorch分布式环境下rank1与rank2点对点通信报错求助
我来帮你分析下这个问题,你遇到的NCCL连接重置错误其实是分布式通信组初始化的常见坑!
先看你代码里的关键问题:只有rank0创建了group2这个新通信组,其他rank根本没执行torch.distributed.new_group([1,2])这行代码。
在PyTorch分布式里,创建新通信组是全局同步操作——所有参与分布式训练的进程(哪怕不在这个新组里)都必须调用new_group。因为NCCL通信器的初始化需要通过全局KV store交换unique ID等信息,要是有进程没参与这个流程,就会导致KV store的请求无法被处理,最终抛出你看到的「Connection reset by peer」错误。
你之前用默认全局组能和rank0通信,是因为全局组的初始化是所有进程都参与的,而新组的初始化只在rank0执行,这就导致了同步异常。
给你几个具体的修复步骤,以及修正后的代码:
修复步骤
- 所有rank都要创建新通信组:不管是不是在
[1,2]这个目标组里,每个进程都得执行new_group调用,保证KV store同步正常。 - 通信时指定目标组:
send和recv要明确加上group=group2参数,告诉PyTorch用你新建的组来通信,而非默认的全局组。 - 确保分布式后端用NCCL:GPU之间的通信必须依赖NCCL后端,初始化时要明确指定。
- 张量绑定对应GPU设备:显式指定
cuda:{rank},避免设备不匹配问题。
修正后的代码示例
def runTpoly(rank, size, pp, cs, pkArithmetics_evals, pkSelectors_evals, domain): # 初始化分布式环境,明确指定NCCL后端 torch.distributed.init_process_group(backend='nccl', rank=rank, world_size=size) # 所有rank都执行新组创建,不管是否属于该组 group2 = torch.distributed.new_group([1,2]) if rank == 0: device = torch.device(f"cuda:{rank}") wo_eval_8n = torch.ones(SCALE * 8 * 1, 4, dtype=torch.int64, device=device) if rank == 1: device = torch.device(f"cuda:{rank}") wo_eval_8n = torch.ones(SCALE * 8 * 10, 4, dtype=torch.int64, device=device) wo_eval_8n = wo_eval_8n + wo_eval_8n # 指定用group2发送数据 torch.distributed.send(wo_eval_8n, dst=2, group=group2) if rank == 2: device = torch.device(f"cuda:{rank}") wo_eval_8n = torch.ones(SCALE * 8 * 10, 4, dtype=torch.int64, device=device) print(wo_eval_8n.size()) # 指定用group2接收数据 torch.distributed.recv(wo_eval_8n, src=1, group=group2) print(wo_eval_8n) if rank == 3: device = torch.device(f"cuda:{rank}") wo_eval_8n = torch.ones(SCALE * 10 * 10, 4, dtype=torch.int64, device=device) print(wo_eval_8n.size()) # 清理进程组 torch.distributed.destroy_process_group() if __name__ == "__main__": world_size = 4 # GPU数目 print(torch.__file__) pp, pk, cs = load("/home/whyin/data/9-data/") domain= Radix2EvaluationDomain.new(cs.circuit_bound()) torch.multiprocessing.spawn(runTpoly, args=(world_size,pp,cs,pk.arithmetics_evals,pk.selectors_evals,domain), nprocs=world_size, join=True)
额外注意点
- 如果你原来的
init_process函数已经做了init_process_group的工作,可以保留,但一定要确保后端是NCCL。 - 要保证rank1和rank2的通信张量形状、dtype、设备完全一致,这点你代码里已经做到了,没问题。
- 测试时可以先简化场景,比如只保留rank1和2的通信逻辑,排除其他干扰。
这样修改后,应该就能实现rank1和rank2绕过rank0的直接点对点通信了。
备注:内容来源于stack exchange,提问作者wynne yin
相关产品推荐
相关产品推荐

