数据截至 (上游 commit cfacd76a0bdd)
01 · 单控制器编程模型
这一章讲什么: verl 最核心、也最值得抄走的设计——
@register装饰器、Dispatch 模式、WorkerGroup 方法绑定。读完你会明白 driver 上那行self.actor_rollout_wg.update_actor(batch)到底发生了什么。
1. 它要解决的小问题
分布式训练代码通常长这样(多控制器 / SPMD):每张卡跑同一份脚本,靠 if rank == 0 区分行为,靠 all_reduce 通信。这在纯预训练里很好用,因为只有一个模型、一条数据流。
但 RL 后训练有 4 个模型(actor / critic / ref / reward)和一条分支很多的数据流:先生成、再打分、再算优势、再更新两个模型。用 SPMD 写出来会变成一坨 if rank == ... 的意大利面,而且换个算法就要重写通信。
verl 的选择:混合。 控制流用单控制器(一个 driver 顺序发号施令),计算用多控制器(每个 WorkerGroup 内部还是 SPMD)。难点就一句话:
怎么让 driver 上一次普通的方法调用,自动变成「把 batch 切成 N 份 → 分别发给 N 个 rank → 等结果 → 拼回一个 batch」?
2. 思路:把「怎么切、怎么收」标注在方法上
直觉很简单——数据怎么分发是方法的属性,不是调用点的属性。
update_actor要按数据并行切分 → 每个 DP rank 拿 1/N 数据,结果拼回来。init_model是全体一起做同一件事 → 每个 rank 拿一模一样的参数,结果取一份。save_checkpoint也是全体动作,但语义不同。
于是 verl 用装饰器把这个属性写在方法定义上,driver 侧调用时什么都不用管:
# 示意,非源码
class MyWorker(Worker):
@register(dispatch_mode=Dispatch.DP_COMPUTE_PROTO) # 按 DP 切分,结果拼接
def update_actor(self, data):
return train_one_step(data)
@register(dispatch_mode=Dispatch.ONE_TO_ALL) # 全体同参数
def init_model(self):
...
# driver 侧:看起来就是单机调用
wg = RayWorkerGroup(resource_pool, MyWorker)
out = wg.update_actor(batch) # batch 自动切 N 份,out 自动拼回来
重点看:调用点没有任何分布式痕迹。这就是「单控制器」体验的全部来源。
3. 图示:一次调用的完整展开
driver 侧 worker 侧 (N 个 Ray actor)
───────── ────────────────────────
wg.update_actor(batch)
│
▼
① dispatch_fn(wg, batch)
把 batch 切成 [b0, b1, ... bN-1]
│
▼
② execute_fn("update_actor", b0..bN-1)
ray actor 逐个 .remote(bi) ─────────► rank0.update_actor(b0)
rank1.update_actor(b1)
...
│ rankN.update_actor(bN-1)
▼ │
③ ray.get(...) 等待 │ 各自返回 oi
│ ◄────────────────────────────────────┘
▼
④ collect_fn(wg, [o0..oN-1])
拼成一个输出
│
▼
返回给调用者
这四步的真实实现只有十几行,在 func_generator(verl/single_controller/ray/base.py:49)里:
def __call__(this, *args, **kwargs):
args, kwargs = dispatch_fn(self, *args, **kwargs)
padding_count = kwargs.pop(_padding_size_key, 0)
output = execute_fn(method_name, *args, **kwargs)
if blocking:
output = ray.get(output)
output = collect_fn(self, output)
这段就是「单控制器」这四个字的全部机械原理。注意它还顺手处理了 padding 回收——因为 batch 未必能被 world_size 整除,见 §5。
4. 注册表:有哪些 Dispatch 模式
模式定义在 DISPATCH_MODE_FN_REGISTRY(verl/single_controller/base/decorator.py:308)。
| 模式 | 分发行为 | 收集行为 | 典型用途 |
|---|---|---|---|
ONE_TO_ALL | 同一份参数 复制给所有 rank | 返回所有 rank 的结果列表 | init_model、save_checkpoint、to(device) |
ALL_TO_ALL | 原样透传(调用者自己已按 rank 备好) | 原样返回 | 底层控制类调用 |
DP_COMPUTE | 要求参数已是长度 = world_size 的列表 | 返回列表 | 手工分发场景 |
DP_COMPUTE_PROTO | 把 DataProto 按 world_size 均分(带自动 padding) | 各 rank 结果 concat 成一个 DataProto | 早期的 actor/critic 计算 |
DP_COMPUTE_PROTO_WITH_FUNC | 第一个参数是函数、广播;其余按 DP 切 | concat | 传自定义算子进 worker |
DP_COMPUTE_METRIC | 按 DP 切 | 只收集不拼接(指标是 dict) | 指标回传 |
DIRECT_ROLLOUT_METHOD | 直接报错 | 直接报错 | 占位,禁止误用(dummy_direct_rollout_call) |
Dispatch 本身是 DynamicEnum,可以运行时注册新模式(register_dispatch_mode,verl/single_controller/base/decorator.py:338),所以 recipe 作者能加自己的切分策略而不改框架。
4.1 更聪明的一档:按 device mesh 懒查询
上面 DP_COMPUTE_PROTO 有个隐含假设:world_size == DP size。一旦开了张量并行(TP)或流水并行(PP),这就不成立了——8 个进程可能只有 2 个 DP 组。
verl 的解法是 make_nd_compute_dataproto_dispatch_fn(mesh_name)(verl/single_controller/base/decorator.py:300):不再假设,而是问 worker 自己。
第一次调用 wg.update_actor(...)
│
├─► dispatch_lazy_compute_data_proto("actor", wg, batch)
│ │
│ ├─ wg._dispatch_info 里没有 "actor"?
│ │ └─► 向所有 rank 广播 _query_dispatch_info("actor")
│ │ 拿回 [0,0,1,1,2,2,3,3] ← 每个 rank 的 dp_rank
│ │ 缓存起来,以后不再问
│ │
│ └─ dp_size = max(mapping)+1 = 4,按 4 份切,再按 mapping 铺到 8 个 rank
│
└─► collect 时同理,用 collect_mask 只收「该出数据的那个 rank」
谁来登记这个 mesh?worker 自己在初始化时登记(_register_dispatch_collect_info,verl/single_controller/base/worker.py:86)。TrainingWorker.__init__ 里就有:
self._register_dispatch_collect_info(
mesh_name="train",
dp_rank=self.engine.get_data_parallel_rank(),
is_collect=self.engine.is_mp_src_rank_with_outputs(),
)
(verl/workers/engine_workers.py:145-149)
妙在哪: driver 完全不需要知道底层是 FSDP 还是 Megatron、TP 开了几路。引擎自己知道自己的 DP rank,driver 只管问。这条抽象让同一份 trainer 代码同时跑得动 FSDP 和 Megatron。
ActorRolloutRefWorker 登记了 actor 和 ref 两个 mesh(verl/workers/engine_workers.py:583 附近的 set_dispatch_collect),对应 compute_log_prob 与 compute_ref_log_prob 两条不同的分发路径(verl/workers/engine_workers.py:687、:645)。
5. 关键细节:自动 padding
batch_size 常常不能被 world_size 整除。verl 的处理在 _split_args_kwargs_data_proto_with_auto_padding(verl/single_controller/base/decorator.py:91):
- 算出补几条:
padding_size = chunks - (len % chunks)。 - 复制已有样本补齐,切分。
- 把
padding_size塞进 kwargs 的特殊 key(_padding_size_key)。 func_generator收集完成后按这个数把尾巴切掉。
注意这个开关是按 DataProto 实例控制的(DataProto.is_padding_enabled(),verl/protocol.py:840),不是全局开着——所以只有明确启用了 auto_padding 的数据才会被补。
6. 方法是怎么「长」到 WorkerGroup 上的
WorkerGroup._bind_worker_method(verl/single_controller/base/worker_group.py:185)做的事就一句话:扫描 worker 类的所有方法,凡是带 MAGIC_ATTR 的,就在 WorkerGroup 实例上 setattr 一个同名代理函数。
MAGIC_ATTR = "attrs_3141562937"(verl/single_controller/base/decorator.py:23)——故意取个圆周率味的怪名字,避免和用户自定义属性撞车。这是个很实用的小技巧。
MyWorker 类 RayWorkerGroup 实例
────────── ───────────────────
@register(...) dir() 扫描
def update_actor ── 有 MAGIC ──► setattr(wg, "update_actor", Functor())
def _helper ── 没有 ──► 跳过
@register(...)
def init_model ── 有 MAGIC ──► setattr(wg, "init_model", Functor())
代理函数是用 type(method_name, (Functor,), {})() 造出来的(verl/single_controller/ray/base.py:67),注释说明理由是「用类型名传递方法名以获得更好的可观测性」——异常栈里能直接看到是哪个方法炸的。
7. 共置:一个进程装多个角色
RL 的 4 个模型如果各占一批 GPU,利用率会很惨。verl 默认把 actor / ref(以及可选的 critic)塞进同一个进程。
实现是 create_colocated_worker_cls(verl/single_controller/ray/base.py:984),做三件事:
class_dict = {"actor_rollout_ref": ActorRolloutRefWorker, "critic": TrainingWorker}
│
▼
① 动态造一个 WorkerDict 类,__init__ 里实例化字典里每个 worker
(用 DISABLE_WORKER_INIT=1 环境变量避免重复初始化分布式环境)
│
▼
② 把每个内层 worker 的方法以「前缀_方法名」绑到 WorkerDict 上
actor_rollout_ref_update_actor / critic_train_mini_batch ...
│
▼
③ ray.remote(WorkerDict) 起进程
然后 RayWorkerGroup.spawn(prefix_set)(verl/single_controller/ray/base.py:714)把这一组进程再切成多个逻辑 WorkerGroup:每个 group 只暴露自己前缀的方法,并把前缀去掉。
物理:8 个 Ray actor,每个进程里既有 actor 又有 critic
│
spawn({"actor_rollout_ref", "critic"})
│
┌───────────┴───────────┐
▼ ▼
actor_rollout_wg critic_wg
.update_actor() .train_mini_batch()
(同 8 个进程) (同 8 个进程)
driver 侧因此可以写 self.actor_rollout_wg.update_actor(...) 和 self.critic_wg.train_mini_batch(...),读起来像两个独立集群,实际上共用同一批进程和同一批 GPU。这是 verl 显存效率的第一个来源。
源码里
create_colocated_worker_cls和_bind_workers_method_to_parent都标了# deprecated, switching to FusedWorker(verl/single_controller/ray/base.py:915、:987),但 V1 trainer 当前仍在用它(verl/trainer/ppo/v1/trainer_base.py:293)。FusedWorker 路径(create_colocated_worker_raw_cls,:1035)已存在但未成为主线。
8. 资源池:GPU 怎么分
ResourcePoolManager(verl/single_controller/ray/base.py:185)把 Ray placement group 包了一层,语义是「一个池 = 一组 [每节点 GPU 数] * 节点数」。
V1 的默认分法在 _init_resource_pool_mgr(verl/trainer/ppo/v1/trainer_base.py:733):
| 池名 | 装什么 | 何时单独开 |
|---|---|---|
global_pool | actor + rollout + ref + critic | 总是 |
reward_pool | 奖励模型 | reward.reward_model.enable_resource_pool=true |
teacher_pool | 蒸馏教师模型 | 开启 on-policy distillation |
还有一个容易忽略但很实用的细节:sort_placement_group_by_node_ip(verl/single_controller/ray/base.py:70)在建 worker 前按节点 IP 给 placement group 排序。原因写在 docstring 里——FSDP checkpoint 按 rank 分片存本地盘,如果重启后 rank 和节点的对应关系变了,恢复就会读错分片。排序让 rank↔节点映射在集群不变时保持稳定。 这是那种「不踩过坑写不出来」的代码。
9. 代码地图
| 主题 | 文件路径 | 符号名 |
|---|---|---|
| 注册装饰器与魔法属性 | verl/single_controller/base/decorator.py | register、MAGIC_ATTR、Dispatch、Execute |
| 分发/收集函数注册表 | verl/single_controller/base/decorator.py | DISPATCH_MODE_FN_REGISTRY、register_dispatch_mode |
| 按 mesh 懒查询的分发 | verl/single_controller/base/decorator.py | make_nd_compute_dataproto_dispatch_fn、dispatch_lazy_compute_data_proto、collect_lazy_compute_data_proto |
| 自动 padding | verl/single_controller/base/decorator.py | _split_args_kwargs_data_proto_with_auto_padding |
| 方法绑定 | verl/single_controller/base/worker_group.py | WorkerGroup._bind_worker_method |
| 调用展开的四步 | verl/single_controller/ray/base.py | func_generator |
| Ray worker 组 | verl/single_controller/ray/base.py | RayWorkerGroup、RayClassWithInitArgs |
| 逻辑切组 | verl/single_controller/ray/base.py | RayWorkerGroup.spawn、spawn_fused |
| 共置类合成 | verl/single_controller/ray/base.py | create_colocated_worker_cls、_bind_workers_method_to_parent |
| 资源池 | verl/single_controller/ray/base.py | RayResourcePool、ResourcePoolManager、sort_placement_group_by_node_ip |
| worker 侧 mesh 登记 | verl/single_controller/base/worker.py | Worker._register_dispatch_collect_info、_query_dispatch_info、query_collect_info |