跳转至

worker_group.py — verl/single_controller/base/worker_group.py

文件路径

verl/single_controller/base/worker_group.py

文件概述

这个文件定义了三个核心类:

  1. ResourcePool:描述集群资源分布(每个节点有多少进程)
  2. ClassWithInitArgs:延迟实例化包装器(存储类和参数,稍后再创建)
  3. WorkerGroup:Worker 组管理器,是 Controller 操控多个 Worker 的中枢

另外还有一个辅助函数 check_workers_alive 用于监控 Worker 存活状态。

关键代码讲解

1. ResourcePool:资源池

class ResourcePool:
    def __init__(self, process_on_nodes=None, max_colocate_count: int = 10, n_gpus_per_node=8) -> None:
        if process_on_nodes is None:
            process_on_nodes = []
        self._store = process_on_nodes
        self.max_colocate_count = max_colocate_count
        self.n_gpus_per_node = n_gpus_per_node

ResourcePool 描述集群的资源分布。核心属性 _store 是一个列表,每个元素表示一个节点上的进程数。

示例:

# 2个节点,第一个节点4个进程,第二个节点4个进程
pool = ResourcePool(process_on_nodes=[4, 4])
print(pool.world_size)  # 8(总共8个进程)

world_size 属性

@property
def world_size(self):
    return sum(self._store)

总进程数 = 所有节点进程数之和。

local_world_size_list 和 local_rank_list

def local_world_size_list(self) -> list[int]:
    nested_local_world_size_list = [
        [local_world_size for _ in range(local_world_size)] for local_world_size in self._store
    ]
    return [item for row in nested_local_world_size_list for item in row]

def local_rank_list(self) -> list[int]:
    nested_local_rank_list = [[i for i in range(local_world_size)] for local_world_size in self._store]
    return [item for row in nested_local_rank_list for item in row]

这两个方法生成扁平化的列表,方便为每个进程分配 local 信息。

示例: process_on_nodes = [2, 3] - local_world_size_list() = [2, 2, 3, 3, 3] - 前2个进程在节点0(local_world_size=2),后3个在节点1(local_world_size=3) - local_rank_list() = [0, 1, 0, 1, 2] - 节点0的两个进程 local_rank 分别为 0 和 1 - 节点1的三个进程 local_rank 分别为 0、1 和 2

max_colocate_count

self.max_colocate_count = max_colocate_count

一个 GPU 上最多可以共存多少个进程。在 RLHF 中,Actor、Critic、Reference 模型可能共享同一张 GPU(每个模型一个进程)。max_colocate_count=3 意味着一张 GPU 上最多放 3 个 Worker。

2. ClassWithInitArgs:延迟实例化

class ClassWithInitArgs:
    def __init__(self, cls, *args, **kwargs) -> None:
        self.cls = cls
        self.args = args
        self.kwargs = kwargs
        self.fused_worker_used = False

    def __call__(self) -> Any:
        return self.cls(*self.args, **self.kwargs)

这个类的设计模式是"延迟构造"(Lazy Instantiation)。它存储了: - cls:要实例化的类 - args、kwargs:构造参数

在分布式场景中,类的实例化需要在远程节点上进行,而不是在本地。所以先把"怎么创建"的信息存下来,到远程节点时再调用 __call__() 创建实例。

使用示例:

# 不是马上创建 MyWorker,而是存储创建信息
cia = ClassWithInitArgs(MyWorker, config=my_config)

# 后续在远程节点上调用 cia() 来实际创建
worker = cia()  # 等同于 MyWorker(config=my_config)

3. check_workers_alive:Worker 存活监控

def check_workers_alive(workers: list, is_alive: Callable, gap_time: float = 1) -> None:
    while True:
        for worker in workers:
            if not is_alive(worker):
                logging.warning(f"worker {worker} is not alive sending signal to main thread")
                signal.raise_signal(signal.SIGABRT)
        time.sleep(gap_time)

这个函数在后台线程中持续检查所有 Worker 是否存活。如果发现任何 Worker 死亡,就向主线程发送 SIGABRT 信号,导致程序终止。这是一种 fail-fast 策略:与其让训练在部分 Worker 死亡的情况下继续(可能导致难以调试的挂起),不如立即失败。

4. WorkerGroup:Worker 组管理器

class WorkerGroup:
    fused_worker_execute_fn_name = "_fuw_execute"

    def __init__(self, resource_pool: ResourcePool, **kwargs) -> None:
        self._is_init_with_detached_workers = resource_pool is None
        self.fused_worker_used = False

        if resource_pool is not None:
            self._procecss_dispatch_config = resource_pool()
        else:
            self._procecss_dispatch_config = None

        self._workers = []
        self._worker_names = []
        self._dispatch_info = {}
        self._collect_info = {}
        self._master_addr = None
        self._master_port = None
        self._checker_thread: threading.Thread = None

WorkerGroup 的初始化: - _is_init_with_detached_workers:如果 resource_pool 为 None,说明这个 WorkerGroup 是附加到已有的 Worker 上(而非自己创建新的) - _workers:所有 Worker 的列表 - _dispatch_info / _collect_info:缓存各 mesh 的分发/收集信息 - _checker_thread:后台存活检查线程

Worker 存活检查

def start_worker_aliveness_check(self, every_n_seconds=1) -> None:
    self._block_until_all_workers_alive()
    self._checker_thread = threading.Thread(
        target=check_workers_alive, args=(self._workers, self._is_worker_alive, every_n_seconds)
    )
    self._checker_thread.start()

启动存活检查前,先阻塞等待所有 Worker 都就绪,然后启动后台守护线程。

_bind_worker_method:核心绑定逻辑

def _bind_worker_method(self, user_defined_cls, func_generator):
    method_names = []
    for method_name in dir(user_defined_cls):
        try:
            method = getattr(user_defined_cls, method_name)
            assert callable(method)
        except Exception:
            continue

        if hasattr(method, MAGIC_ATTR):
            attribute = getattr(method, MAGIC_ATTR)
            dispatch_mode = attribute["dispatch_mode"]
            execute_mode = attribute["execute_mode"]
            blocking = attribute["blocking"]

            # 获取分发函数
            if isinstance(dispatch_mode, Dispatch):
                fn = get_predefined_dispatch_fn(dispatch_mode=dispatch_mode)
                dispatch_fn = fn["dispatch_fn"]
                collect_fn = fn["collect_fn"]
            else:
                dispatch_fn = dispatch_mode["dispatch_fn"]
                collect_fn = dispatch_mode["collect_fn"]

            # 获取执行函数
            execute_mode = get_predefined_execute_fn(execute_mode=execute_mode)
            wg_execute_fn_name = execute_mode["execute_fn_name"]
            execute_fn = getattr(self, wg_execute_fn_name)

            # 生成并绑定代理方法
            func = func_generator(self, method_name,
                dispatch_fn=dispatch_fn, collect_fn=collect_fn,
                execute_fn=execute_fn, blocking=blocking)
            setattr(self, method_name, func)
            method_names.append(method_name)

    return method_names

这是整个 single_controller 模块中最重要的方法。 它的工作流程:

  1. 遍历用户定义的 Worker 类的所有方法
  2. 找到带有 MAGIC_ATTR 标记的方法(即用 @register 装饰过的)
  3. 读取方法的分发配置(dispatch_mode、execute_mode、blocking)
  4. 从注册表中获取对应的 dispatch_fn 和 collect_fn
  5. 调用 func_generator 生成一个代理函数
  6. 把代理函数绑定到 WorkerGroup 上,使其成为 WorkerGroup 的方法

效果: 用户可以直接调用 worker_group.train_step(data),WorkerGroup 会自动执行:数据分发 -> 远程执行 -> 结果收集。

绑定流程图:

Worker 类定义:
    @register(dispatch_mode=Dispatch.DP_COMPUTE_PROTO)
    def train_step(self, data): ...

          |  _bind_worker_method 遍历发现这个方法
          v

从 MAGIC_ATTR 读取:
    dispatch_mode = Dispatch.DP_COMPUTE_PROTO
    execute_mode = Execute.ALL

          |
          v

查注册表获取:
    dispatch_fn = dispatch_dp_compute_data_proto
    collect_fn = collect_dp_compute_data_proto
    execute_fn = self.execute_all

          |
          v

func_generator 生成代理函数:
    def proxy(*args, **kwargs):
        args, kwargs = dispatch_fn(self, *args, **kwargs)
        output = execute_fn("train_step", *args, **kwargs)
        output = ray.get(output)
        output = collect_fn(self, output)
        return output

          |
          v

绑定到 WorkerGroup:
    setattr(self, "train_step", proxy)

核心类/函数列表

名称 类型 说明
ResourcePool 类 描述集群资源分布
ResourcePool.world_size 属性 总进程数
ResourcePool.local_world_size_list() 方法 各进程的 local world size
ResourcePool.local_rank_list() 方法 各进程的 local rank
ClassWithInitArgs 类 延迟实例化包装器
check_workers_alive() 函数 后台监控 Worker 存活
WorkerGroup 类 Worker 组管理器
WorkerGroup._bind_worker_method() 方法 将 Worker 方法绑定为 WorkerGroup 的代理方法
WorkerGroup.start_worker_aliveness_check() 方法 启动后台存活监控
WorkerGroup.world_size 属性 Worker 数量

与其他模块的关系

  • 依赖 decorator.py:使用 MAGIC_ATTR、Dispatch、get_predefined_dispatch_fn、get_predefined_execute_fn
  • 管理 Worker 实例:_workers 列表存储所有受管理的 Worker
  • 被 ray/base.py 继承:RayWorkerGroup 继承 WorkerGroup,RayResourcePool 继承 ResourcePool
  • 被用户代码使用:Trainer 通过 WorkerGroup 调用 Worker 上的方法

小结

worker_group.py 定义了分布式管理层的三个核心抽象:

  • ResourcePool 回答"有多少资源可用"
  • ClassWithInitArgs 回答"怎么在远程创建 Worker"
  • WorkerGroup 回答"怎么管理和调用这些 Worker"

其中 _bind_worker_method 是核心中的核心,它让用户能够透明地调用远程方法,无需关心数据分发和结果收集的细节。