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

如何在torch.distributed中收集非Tensor对象(如字典)?

在torch.distributed中收集非Tensor对象的实现方法

要在torch.distributed中收集各rank的非Tensor对象(比如你的字典),可以通过序列化+Tensor通信的方式实现,以下是两种可行方案:

方案一:基于all_gather的通用序列化收集法

这种方法适用于任意可序列化的Python对象,支持单节点/多节点分布式场景,核心是将对象序列化后转成Tensor完成分布式收集,再反序列化合并。

实现代码

import torch
import torch.distributed as dist
import pickle
from typing import Dict

def init_distributed():
    dist.init_process_group("nccl")  # 根据你的后端调整,比如gloo
    rank = dist.get_rank()
    world_size = dist.get_world_size()
    return rank, world_size

def gather_non_tensor_objects(obj) -> list:
    """收集所有rank的非Tensor对象,返回所有对象的列表"""
    # 1. 序列化当前rank的对象
    serialized_obj = pickle.dumps(obj)
    # 2. 收集所有rank的序列化字节长度
    obj_len = torch.tensor(len(serialized_obj), dtype=torch.int64, device=f"cuda:{dist.get_rank()}")
    all_obj_lens = [torch.zeros_like(obj_len) for _ in range(dist.get_world_size())]
    dist.all_gather(all_obj_lens, obj_len)
    max_len = max([x.item() for x in all_obj_lens])
    # 3. 将序列化字节填充到最大长度,转成Tensor
    padded_bytes = serialized_obj.ljust(max_len, b'\x00')
    obj_tensor = torch.tensor(list(padded_bytes), dtype=torch.uint8, device=f"cuda:{dist.get_rank()}")
    # 4. 收集所有rank的Tensor
    all_obj_tensors = [torch.zeros_like(obj_tensor) for _ in range(dist.get_world_size())]
    dist.all_gather(all_obj_tensors, obj_tensor)
    # 5. 反序列化每个对象
    gathered_objs = []
    for tensor, length in zip(all_obj_tensors, all_obj_lens):
        bytes_data = tensor.cpu().numpy().tobytes()[:length.item()]
        gathered_objs.append(pickle.loads(bytes_data))
    return gathered_objs

if __name__ == "__main__":
    rank, world_size = init_distributed()
    # 每个rank生成自己的字典
    local_dict = {rank: 1} if rank % 2 == 0 else {}  # 模拟你给出的P0/P2/P4等持有字典的情况
    # 收集所有rank的对象
    all_objs = gather_non_tensor_objects(local_dict)
    # 合并所有非空字典
    merged_dict = {}
    for obj in all_objs:
        merged_dict.update(obj)
    # 打印结果
    print(f"P {rank}: {local_dict}")
    if rank == 0:  # 仅在rank0打印合并后的结果,也可改为所有rank都打印
        print(f"All: {merged_dict}")
    dist.destroy_process_group()

使用说明

  • 用torchrun启动时,示例命令:torchrun --nproc_per_node=10 your_script.py(对应你的10个rank场景)
  • 序列化用pickle,若担心安全性或需要支持更多对象类型,可替换为cloudpickle

方案二:基于点对点通信的定向收集

如果只需要把所有rank的字典收集到某一个rank(比如rank0),可以用点对点的dist.recv和dist.send,逻辑更简单:

实现代码

import torch
import torch.distributed as dist
import pickle

def init_distributed():
    dist.init_process_group("nccl")
    rank = dist.get_rank()
    world_size = dist.get_world_size()
    return rank, world_size

if __name__ == "__main__":
    rank, world_size = init_distributed()
    local_dict = {rank: 1} if rank % 2 == 0 else {}
    merged_dict = {}
    if rank == 0:
        # rank0先加入自己的字典
        merged_dict.update(local_dict)
        # 接收其他所有rank的字典
        for src_rank in range(1, world_size):
            serialized_obj = dist.recv(tensor=None, src=src_rank)
            remote_dict = pickle.loads(serialized_obj.cpu().numpy().tobytes())
            merged_dict.update(remote_dict)
    else:
        # 其他rank把序列化后的字典发送给rank0
        serialized_obj = pickle.dumps(local_dict)
        tensor = torch.tensor(list(serialized_obj), dtype=torch.uint8, device=f"cuda:{rank}")
        dist.send(tensor, dst=0)
    # 打印结果
    print(f"P {rank}: {local_dict}")
    if rank == 0:
        print(f"All: {merged_dict}")
    dist.destroy_process_group()

关于Manager失效的原因

你之前尝试的multiprocessing.Manager是基于共享内存的进程间通信工具,仅适用于单节点内的普通多进程场景。而torchrun启动的是分布式进程(可能跨节点),进程间不共享内存,因此Manager无法跨rank同步数据,自然失效。

内容的提问来源于stack exchange,提问作者juliozhao97

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.21 15:27:45