跳到主要内容

数据截至 (上游 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_generatorverl/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_REGISTRYverl/single_controller/base/decorator.py:308)。

模式分发行为收集行为典型用途
ONE_TO_ALL同一份参数复制给所有 rank返回所有 rank 的结果列表init_modelsave_checkpointto(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_modeverl/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_infoverl/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 登记了 actorref 两个 mesh(verl/workers/engine_workers.py:583 附近的 set_dispatch_collect),对应 compute_log_probcompute_ref_log_prob 两条不同的分发路径(verl/workers/engine_workers.py:687:645)。


5. 关键细节:自动 padding

batch_size 常常不能被 world_size 整除。verl 的处理在 _split_args_kwargs_data_proto_with_auto_paddingverl/single_controller/base/decorator.py:91):

  1. 算出补几条:padding_size = chunks - (len % chunks)
  2. 复制已有样本补齐,切分。
  3. padding_size 塞进 kwargs 的特殊 key(_padding_size_key)。
  4. func_generator 收集完成后按这个数把尾巴切掉。

注意这个开关是按 DataProto 实例控制的(DataProto.is_padding_enabled()verl/protocol.py:840),不是全局开着——所以只有明确启用了 auto_padding 的数据才会被补。


6. 方法是怎么「长」到 WorkerGroup 上的

WorkerGroup._bind_worker_methodverl/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_clsverl/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 FusedWorkerverl/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 怎么分

ResourcePoolManagerverl/single_controller/ray/base.py:185)把 Ray placement group 包了一层,语义是「一个池 = 一组 [每节点 GPU 数] * 节点数」。

V1 的默认分法在 _init_resource_pool_mgrverl/trainer/ppo/v1/trainer_base.py:733):

池名装什么何时单独开
global_poolactor + rollout + ref + critic总是
reward_pool奖励模型reward.reward_model.enable_resource_pool=true
teacher_pool蒸馏教师模型开启 on-policy distillation

还有一个容易忽略但很实用的细节:sort_placement_group_by_node_ipverl/single_controller/ray/base.py:70)在建 worker 前按节点 IP 给 placement group 排序。原因写在 docstring 里——FSDP checkpoint 按 rank 分片存本地盘,如果重启后 rank 和节点的对应关系变了,恢复就会读错分片。排序让 rank↔节点映射在集群不变时保持稳定。 这是那种「不踩过坑写不出来」的代码。


9. 代码地图

主题文件路径符号名
注册装饰器与魔法属性verl/single_controller/base/decorator.pyregisterMAGIC_ATTRDispatchExecute
分发/收集函数注册表verl/single_controller/base/decorator.pyDISPATCH_MODE_FN_REGISTRYregister_dispatch_mode
按 mesh 懒查询的分发verl/single_controller/base/decorator.pymake_nd_compute_dataproto_dispatch_fndispatch_lazy_compute_data_protocollect_lazy_compute_data_proto
自动 paddingverl/single_controller/base/decorator.py_split_args_kwargs_data_proto_with_auto_padding
方法绑定verl/single_controller/base/worker_group.pyWorkerGroup._bind_worker_method
调用展开的四步verl/single_controller/ray/base.pyfunc_generator
Ray worker 组verl/single_controller/ray/base.pyRayWorkerGroupRayClassWithInitArgs
逻辑切组verl/single_controller/ray/base.pyRayWorkerGroup.spawnspawn_fused
共置类合成verl/single_controller/ray/base.pycreate_colocated_worker_cls_bind_workers_method_to_parent
资源池verl/single_controller/ray/base.pyRayResourcePoolResourcePoolManagersort_placement_group_by_node_ip
worker 侧 mesh 登记verl/single_controller/base/worker.pyWorker._register_dispatch_collect_info_query_dispatch_infoquery_collect_info