跳转至

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 训练的通信基础,封装了进程组管理和集合通信操作。