如何在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
相关产品推荐
相关产品推荐

