distributed.py — 分布式训练工具¶
文件路径: verl/utils/distributed.py
文件概述¶
distributed.py 封装了 PyTorch 分布式训练的基础设施,包括进程组初始化、rank 信息获取、分布式通信操作等。是多 GPU 训练的基础。
背景知识¶
在分布式训练中: - rank: 每个进程的唯一编号(0, 1, 2, ...) - world_size: 参与训练的进程总数 - ProcessGroup: 一组进程的通信组,可以在组内进行集合通信(如 all_reduce)
核心函数详解¶
1. 初始化分布式环境¶
def initialize_global_process_group(timeout_second=36000):
"""初始化全局进程组"""
if not dist.is_initialized():
backend = get_nccl_backend() # 自动选择 nccl 或 hccl
dist.init_process_group(backend=backend, timeout=timedelta(seconds=timeout_second))
必须在任何分布式操作之前调用。它读取环境变量(RANK、WORLD_SIZE、MASTER_ADDR、MASTER_PORT)来建立通信。
2. 获取 rank 信息¶
def get_rank():
return dist.get_rank() if dist.is_initialized() else 0
def get_world_size():
return dist.get_world_size() if dist.is_initialized() else 1
安全获取当前进程的 rank 和总进程数,在非分布式环境下返回默认值。
3. 分布式通信操作¶
def allgather_dict_tensors(tensor_dict, world_size, group=None):
"""对字典中所有张量执行 all_gather"""
def broadcast_dict_tensor(tensor_dict, src=0, group=None):
"""从 src rank 广播张量字典到所有 rank"""
这些函数将 PyTorch 的基础通信原语包装为更方便的字典操作,因为 verl 中数据通常以 {key: tensor} 形式传递。
核心函数列表¶
| 函数 | 说明 |
|---|---|
initialize_global_process_group() |
初始化分布式环境 |
get_rank() / get_world_size() |
获取 rank 信息 |
allgather_dict_tensors() |
对张量字典执行 all_gather |
broadcast_dict_tensor() |
广播张量字典 |
与其他模块的关系¶
device.py提供后端选择(get_nccl_backend())- 被
fsdp_utils.py、ulysses.py等并行模块依赖 - 被 trainer 在启动时调用初始化
小结¶
distributed.py 是多 GPU 训练的通信基础,封装了进程组管理和集合通信操作。