You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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执行,这就导致了同步异常。

给你几个具体的修复步骤,以及修正后的代码:

修复步骤

  1. 所有rank都要创建新通信组:不管是不是在[1,2]这个目标组里,每个进程都得执行new_group调用,保证KV store同步正常。
  2. 通信时指定目标组:send和recv要明确加上group=group2参数,告诉PyTorch用你新建的组来通信,而非默认的全局组。
  3. 确保分布式后端用NCCL:GPU之间的通信必须依赖NCCL后端,初始化时要明确指定。
  4. 张量绑定对应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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.04.14 12:53:01