如何从PyTorch分布式(NCCL)通信组中获取组内全局rank列表?
关于PyTorch分布式组对象获取内部全局rank的解答
可以直接从创建好的ProcessGroup实例中获取组内包含的所有全局rank值,具体方法如下:
适用PyTorch 1.8及以上版本(推荐)
PyTorch对公开的ProcessGroup类提供了官方方法get_ranks(),直接调用即可返回组内所有全局rank构成的列表,顺序和创建组时传入的rank列表完全一致。
示例代码:
import torch.distributed as dist # 前提:已经完成分布式基础环境的初始化 # 构造包含全局rank 0、1、2、3的自定义分布式组 custom_group = dist.new_group([0,1,2,3]) # 获取组内所有全局rank all_global_ranks = custom_group.get_ranks() print(all_global_ranks) # 输出结果:[0, 1, 2, 3]
低版本PyTorch兼容方案
如果使用的PyTorch版本低于1.8,公开API未提供get_ranks(),可以临时访问私有属性_ranks获取对应值:
all_global_ranks = custom_group._ranks
注意:私有属性不保证跨版本兼容,非必要场景优先选择升级PyTorch版本使用公开API,避免后续版本迭代导致代码异常。
补充常用关联方法
如果需要获取当前进程在该自定义组内的局部rank(组内相对编号,区别于全局rank),可以调用rank()方法:
# 假设当前进程全局rank为1,在上述custom_group中的局部rank就是1 local_rank_in_group = custom_group.rank()
内容的提问来源于stack exchange,提问作者Qin Heyang
相关产品推荐
相关产品推荐

