跳转至

decorator.py — verl/single_controller/base/decorator.py

文件路径

verl/single_controller/base/decorator.py

文件概述

这是 single_controller 模块中最核心的文件之一。它定义了:

  1. Dispatch 枚举:数据分发策略(如广播、按 DP 切分等)
  2. Execute 枚举:执行模式(在所有 Worker 上执行,还是只在 rank 0 执行)
  3. @register 装饰器:用来标记 Worker 方法的分发和执行策略
  4. 一系列 dispatch/collect 函数:实现各种数据分发和结果收集的具体逻辑

简单来说,这个文件回答了一个问题:当 Controller 调用 Worker 的某个方法时,数据怎么分、结果怎么收?

关键代码讲解

1. MAGIC_ATTR 标记

MAGIC_ATTR = "attrs_3141562937"

这是一个"魔法属性"名称(用了一个不太可能冲突的数字后缀)。@register 装饰器会把分发配置信息存储在被装饰函数的这个属性上。后续 WorkerGroup 通过检查方法是否有这个属性,来判断该方法是否需要进行分布式分发。

2. Dispatch 枚举:数据分发策略

class Dispatch(DynamicEnum):
    _registry = {}
    _next_value = 0

def init_predefined_dispatch_mode():
    Dispatch.register("RANK_ZERO")
    Dispatch.register("ONE_TO_ALL")
    Dispatch.register("ALL_TO_ALL")
    Dispatch.register("DP_COMPUTE")
    Dispatch.register("DP_COMPUTE_PROTO")
    Dispatch.register("DP_COMPUTE_PROTO_WITH_FUNC")
    Dispatch.register("DP_COMPUTE_METRIC")
    Dispatch.register("DIRECT_ROLLOUT_METHOD")

Dispatch 继承自 DynamicEnum(可以动态注册新成员的枚举)。预定义的分发模式包括:

模式 含义
RANK_ZERO 仅发给 rank 0 的 Worker
ONE_TO_ALL 同一份数据广播给所有 Worker
ALL_TO_ALL 每个 Worker 收到对应的一份(参数本身就是列表)
DP_COMPUTE 按数据并行维度分发(参数已是列表形式)
DP_COMPUTE_PROTO 把 DataProto 按 DP 维度自动切分(支持 padding)
DP_COMPUTE_PROTO_WITH_FUNC 类似 DP_COMPUTE_PROTO,但第一个参数是函数
DP_COMPUTE_METRIC 数据按 DP 切分,但结果不合并
DIRECT_ROLLOUT_METHOD 特殊模式,用于 vLLM 外部执行器

3. Execute 枚举:执行模式

class Execute(DynamicEnum):
    _registry = {}
    _next_value = 0

def init_predefined_execute_mode():
    Execute.register("ALL")
    Execute.register("RANK_ZERO")

只有两种执行模式: - ALL:在所有 Worker 上执行 - RANK_ZERO:只在 rank 0 的 Worker 上执行

4. @register 装饰器

def register(dispatch_mode=Dispatch.ALL_TO_ALL, execute_mode=Execute.ALL, blocking=True, materialize_futures=True):
    def decorator(func):
        @wraps(func)
        def inner(*args, **kwargs):
            if materialize_futures:
                args, kwargs = _materialize_futures(*args, **kwargs)
            return func(*args, **kwargs)

        wrapper = async_inner if inspect.iscoroutinefunction(func) else inner
        attrs = {"dispatch_mode": dispatch_mode, "execute_mode": execute_mode, "blocking": blocking}
        setattr(wrapper, MAGIC_ATTR, attrs)
        return wrapper

    return decorator

@register 是一个参数化装饰器,工作流程:

  1. 接收 dispatch_mode、execute_mode、blocking、materialize_futures 四个参数
  2. 包装原函数,在调用前自动物化(materialize)所有 DataProtoFuture 类型的参数
  3. 将分发配置存储为函数的 MAGIC_ATTR 属性
  4. 支持同步和异步函数

使用示例:

class MyWorker(Worker):
    @register(dispatch_mode=Dispatch.DP_COMPUTE_PROTO)
    def train_step(self, data):
        return self.model(data)

这样标记后,WorkerGroup 会知道调用 train_step 时需要把 data 按 DP 维度切分给各 Worker。

5. 分发函数(dispatch functions)

dispatch_one_to_all:广播

def dispatch_one_to_all(worker_group, *args, **kwargs):
    args = tuple([arg] * worker_group.world_size for arg in args)
    kwargs = {k: [v] * worker_group.world_size for k, v in kwargs.items()}
    return args, kwargs

把每个参数复制 world_size 份。比如有 4 个 Worker,就把同一个参数复制 4 份。

dispatch_all_to_all:直接透传

def dispatch_all_to_all(worker_group, *args, **kwargs):
    return args, kwargs

不做任何变换,参数直接原样传递。

dispatch_dp_compute_data_proto:自动切分 DataProto

def dispatch_dp_compute_data_proto(worker_group, *args, **kwargs):
    assert isinstance(worker_group, WorkerGroup)
    splitted_args, splitted_kwargs = _split_args_kwargs_data_proto_with_auto_padding(
        worker_group.world_size, *args, **kwargs,
    )
    return splitted_args, splitted_kwargs

自动把 DataProto 按 world_size 切成等份,并在必要时进行 padding(保证每份大小相等)。

6. 收集函数(collect functions)

collect_all_to_all:原样返回

def collect_all_to_all(worker_group, output):
    return output

collect_dp_compute_data_proto:合并结果

def collect_dp_compute_data_proto(worker_group, output):
    output = collect_dp_compute(worker_group, output)
    return _concat_data_proto_or_future(output)

把各 Worker 返回的 DataProto 结果合并(concat)成一个完整的 DataProto。

7. 分发模式注册表

DISPATCH_MODE_FN_REGISTRY = {
    Dispatch.ONE_TO_ALL: {
        "dispatch_fn": dispatch_one_to_all,
        "collect_fn": collect_all_to_all,
    },
    Dispatch.ALL_TO_ALL: {
        "dispatch_fn": dispatch_all_to_all,
        "collect_fn": collect_all_to_all,
    },
    Dispatch.DP_COMPUTE_PROTO: {
        "dispatch_fn": dispatch_dp_compute_data_proto,
        "collect_fn": collect_dp_compute_data_proto,
    },
    # ... 其他模式
}

这个全局字典将每种 Dispatch 枚举值映射到对应的分发函数和收集函数。WorkerGroup 在绑定方法时会查这个表。

8. N维分发(高级功能)

def dispatch_nd_compute(dp_rank_mapping: list[int], dp_size, worker_group, *args, **kwargs):

这是更灵活的分发方式,支持自定义 dp_rank 映射。当 Worker 的逻辑 dp_rank 和物理 rank 不一致时使用。比如在多模型共存的场景下,不同模型的 DP 划分方式可能不同。

9. 动态注册和更新分发模式

def register_dispatch_mode(dispatch_mode_name, dispatch_fn, collect_fn):
    dispatch_mode = Dispatch.register(dispatch_mode_name)
    DISPATCH_MODE_FN_REGISTRY[dispatch_mode] = {"dispatch_fn": dispatch_fn, "collect_fn": collect_fn}

def update_dispatch_mode(dispatch_mode, dispatch_fn, collect_fn):
    DISPATCH_MODE_FN_REGISTRY[dispatch_mode] = {"dispatch_fn": dispatch_fn, "collect_fn": collect_fn}

允许用户在运行时注册新的分发模式或修改已有模式。

核心类/函数列表

名称 类型 说明
Dispatch 动态枚举类 定义数据分发策略
Execute 动态枚举类 定义执行模式
register() 装饰器函数 标记 Worker 方法的分发/执行配置
dispatch_one_to_all() 函数 广播分发
dispatch_all_to_all() 函数 直接透传
dispatch_dp_compute() 函数 DP 维度分发(列表形式)
dispatch_dp_compute_data_proto() 函数 DP 维度分发(DataProto 自动切分)
collect_dp_compute_data_proto() 函数 合并 DataProto 结果
dispatch_nd_compute() 函数 N维分发,支持自定义 rank 映射
DISPATCH_MODE_FN_REGISTRY 全局字典 分发模式到函数的映射表
register_dispatch_mode() 函数 动态注册新分发模式
get_predefined_dispatch_fn() 函数 查询分发模式对应的函数
get_predefined_execute_fn() 函数 查询执行模式对应的函数名

与其他模块的关系

  • 被 worker.py 使用:Worker 类的方法用 @register 装饰
  • 被 worker_group.py 使用:WorkerGroup 在 _bind_worker_method 中查询 MAGIC_ATTR 和分发函数注册表
  • 被 ray/base.py 使用:Ray 后端在 func_generator 中调用 dispatch_fn 和 collect_fn
  • 依赖 verl.protocol:使用 DataProtoFuture 和 DataProto 进行数据切分和合并

数据分发流程图

用户调用 worker_group.method(data)
            |
            v
    dispatch_fn(worker_group, data)     <-- 根据 @register 的 dispatch_mode 选择
            |
            v
    [data_0, data_1, ..., data_n]       <-- 切分成 world_size 份
            |
            v
    execute_fn(method_name, data_i)     <-- 分发到各 Worker 执行
            |
            v
    [result_0, result_1, ..., result_n] <-- 各 Worker 返回结果
            |
            v
    collect_fn(worker_group, results)   <-- 根据 dispatch_mode 选择收集方式
            |
            v
    merged_result                       <-- 返回合并后的结果

小结

decorator.py 是 single_controller 的大脑,定义了分布式方法调用的策略模式。理解这个文件后,你就知道了: - Worker 的方法通过 @register 声明自己的分发策略 - 每种分发策略有对应的 dispatch_fn(分发数据)和 collect_fn(收集结果) - WorkerGroup 在运行时通过 MAGIC_ATTR 查找方法的配置,自动完成数据分发和结果收集