worker_group.py — verl/single_controller/base/worker_group.py¶
文件路径¶
verl/single_controller/base/worker_group.py
文件概述¶
这个文件定义了三个核心类:
- ResourcePool:描述集群资源分布(每个节点有多少进程)
- ClassWithInitArgs:延迟实例化包装器(存储类和参数,稍后再创建)
- 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 属性¶
总进程数 = 所有节点进程数之和。
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¶
一个 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 模块中最重要的方法。 它的工作流程:
- 遍历用户定义的 Worker 类的所有方法
- 找到带有
MAGIC_ATTR标记的方法(即用@register装饰过的) - 读取方法的分发配置(dispatch_mode、execute_mode、blocking)
- 从注册表中获取对应的 dispatch_fn 和 collect_fn
- 调用
func_generator生成一个代理函数 - 把代理函数绑定到 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 是核心中的核心,它让用户能够透明地调用远程方法,无需关心数据分发和结果收集的细节。