Skip to content

Workers:单个 driver 怎样调度多卡模型计算

verl 的 trainer 是一个 single-controller driver:算法主循环写在一个 Python 进程里,但模型 forward、backward、rollout、ref logprob 都在 Ray worker 上跑。学习 worker 这一层,重点不是“又一个模型类”,而是看懂三件事:

  • trainer 怎么创建 RayWorkerGroup
  • @register 怎么把 worker 方法变成可远程调用的分布式 API。
  • worker 里面的 mesh 怎么决定一份 DataProto 要切给哪些 rank。

先补一点先验

把 PPO/GRPO 的一次训练想成两层:

text
driver / trainer:
  组织数据流:rollout -> reward -> old_log_prob/ref_log_prob -> advantage -> update_actor

worker / engine:
  执行重活:model forward -> loss -> backward -> optimizer step -> metrics

RayPPOTrainer.fit() 只是在正确时间调用 self.actor_rollout_wg.compute_log_prob(batch)self.actor_rollout_wg.update_actor(batch) 这类方法。真正“把 batch 切到 N 张卡、在哪些 rank 收集输出”的规则,来自 worker 方法上的 @register(...)

TrainingWorker:一个训练 engine 的统一外壳

verl/workers/engine_workers.py 里的 TrainingWorker 是最小训练单元。它初始化时会:

  1. initialize_global_process_group_ray() 建分布式进程组。
  2. EngineRegistry.new(...) 根据 engine_config.strategy 创建 FSDP、Megatron、VeOmni、TorchTitan、Automodel 等后端 engine。
  3. 通过 _register_dispatch_collect_info(mesh_name="train", dp_rank=..., is_collect=...) 告诉 controller:这个 worker 在 train mesh 里的 DP rank 是谁,哪些 rank 的输出需要收集。

可以把它理解成“把各种并行训练后端包成一套 RPC API”:

python
worker = TrainingWorker(config)
worker.reset()              # engine.initialize()
worker.set_loss_fn(ppo_loss)
worker.infer_batch(batch)   # logprob / value / ref 这类 forward-only
worker.train_mini_batch(batch)  # PPO epoch + mini-batch loop

train_mini_batch() 负责把一个 PPO batch 拆成 mini-batch、跑多个 epoch,再调用 train_batch()train_batch() 才是单次 optimizer step。infer_batch() 走 eval mode,用于 actor logprob、ref logprob、critic value、SFT eval loss 等。

ActorRolloutRefWorker:actor、rollout、ref 的混合 worker

ActorRolloutRefWorker 是 PPO/GRPO 最常见的 worker。它由 role 决定内部建哪些对象:

role内部对象
actorself.actor: TrainingWorker 和 checkpoint engine
rolloutself.rollout: BaseRollout
refself.ref: TrainingWorker
actor_rolloutactor + rollout + checkpoint engine
actor_rollout_refactor + rollout + ref

init_model() 的顺序很重要:

text
1. 如果有 ref:构建 ref TrainingWorker,注册 ref mesh
2. 如果有 actor:构建 actor TrainingWorker,设置 PPO 或 distillation loss,注册 actor mesh
3. 如果有 rollout:按 rollout.name/mode 创建 vLLM/SGLang/TRTLLM rollout adapter
4. 如果有 actor:创建 checkpoint engine,用于 actor -> rollout 权重同步

actor 和 ref 都复用 TrainingWorker,但它们注册的是不同 mesh。于是同一个外层 worker group 可以这样分发:

python
@register(dispatch_mode=make_nd_compute_dataproto_dispatch_fn(mesh_name="actor"))
def compute_log_prob(self, data):
    return self.actor.infer_batch(data)

@register(dispatch_mode=make_nd_compute_dataproto_dispatch_fn(mesh_name="ref"))
def compute_ref_log_prob(self, data):
    return self.ref.infer_batch(data)

@register(dispatch_mode=make_nd_compute_dataproto_dispatch_fn(mesh_name="actor"))
def update_actor(self, data):
    return self.actor.train_mini_batch(data)

这就是初学者容易漏掉的点:compute_log_probcompute_ref_log_prob 看起来都只是 forward,但它们可能使用不同模型、不同并行拓扑、不同 collect rank。

@register 到 WorkerGroup 的调用链

@register 本身不执行 Ray 调用。它只是给函数挂上元数据:

text
dispatch_mode: 数据怎么分发
execute_mode: 在所有 rank 执行还是只在 rank0 执行
blocking: driver 是否立刻 ray.get

真正绑定发生在 WorkerGroup._bind_worker_method()

text
worker class method with @register
  -> 读取 MAGIC_ATTR 里的 dispatch_mode / execute_mode / blocking
  -> 找到 dispatch_fn 和 collect_fn
  -> 用 RayWorkerGroup.func_generator 包一层本地方法
  -> setattr(worker_group, method_name, wrapped_rpc)

所以 trainer 里写:

python
output = self.actor_rollout_wg.update_actor(batch_td)

实际发生的是:

text
1. dispatch_fn(worker_group, batch_td) 把 batch 切成各 worker 的输入
2. execute_all("update_actor", shard_i) 调 Ray actor remote method
3. blocking=True 时 ray.get 等结果;blocking=False 时返回 future 风格结果
4. collect_fn(worker_group, outputs) 把各 rank 输出拼回 DataProto/TensorDict

Dispatch mode:为什么不能所有调用都广播

常见 dispatch 可以这样记:

dispatch直觉例子
ONE_TO_ALL同一条命令发给每个 workerinit_model()reset()save_checkpoint()
DP_COMPUTE调用方已经按 worker 数准备好 list/tuplecheckpoint engine 一类控制调用
DP_COMPUTE_PROTO按 worker world size 切 DataProto传统 DP 数据并行
make_nd_compute_dataproto_dispatch_fn(mesh_name=...)按某个 engine mesh 的 DP rank mapping 切数据actor/ref/train 计算主路径

make_nd_compute_dataproto_dispatch_fn("actor") 会懒查询每个 worker 的 actor mesh dispatch info。Megatron、FSDP、VeOmni 的并行形状不同,但只要 worker 注册了 dp_rankis_collect,controller 就能按统一协议切分和收集。

mesh:不是抽象名词,是 dispatch/collect 的路由表

在 worker 内部,mesh 至少承担两个问题:

text
dispatch: 第 i 个 Ray worker 应该拿哪份 DP shard?
collect: 哪些 rank 的输出代表这个 DP shard,可以回传给 driver?

TrainingWorker.__init__() 对 train mesh 注册:

python
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(),
)

ActorRolloutRefWorker.init_model() 再把内部 actor/ref 的信息挂到外层 worker:

python
self.set_dispatch_collect(mesh_name="actor", **self.actor.get_dispatch_collect())
self.set_dispatch_collect(mesh_name="ref", **self.ref.get_dispatch_collect())

这样 controller 不需要知道“Megatron PP 最后一段才有 logits”这类后端细节,只需要问 mesh:哪些 rank 收结果。

trainer 侧如何创建 worker group

RayPPOTrainer.init_workers() 大致做这几步:

text
ResourcePoolManager.create_resource_pool()
  -> 为 actor_rollout、critic、ref、rm 等角色分配 RayResourcePool

RayClassWithInitArgs(cls=ActorRolloutRefWorker, ...)
  -> 记录远端类和构造参数

RayWorkerGroup(resource_pool, ray_cls_with_init)
  -> 在 placement group 上启动 Ray actors
  -> _bind_worker_method() 绑定 @register 方法

worker_group.init_model()
  -> ONE_TO_ALL 初始化每个远端 worker

如果多个角色 colocate,verl 会创建 fused/colocated worker class。简单说,就是一个 Ray actor 进程里放多个子 worker,减少重复 CUDA/distributed context。然后 spawn(prefix_set) 把它们再暴露成 actor_rolloutcriticref 这些逻辑 worker group。

权重同步在 worker 层的位置

actor 更新后,rollout 服务必须拿到新权重。ActorRolloutRefWorker.update_weights() 有两条路径:

  • mode="naive":同步训练、trainer 和 rollout colocate 时,直接从 self.actor.engine.get_per_tensor_param() 导出 tensor,调用 self.rollout.update_weights(...)
  • naive:异步/解耦训练时,通过 checkpoint engine 的 send_weights() 把权重送到 rollout 侧。

这里也能看到 sleep/resume 的工程顺序:如果开启 free_cache_engine,先恢复 weights 内存,同步权重,再恢复 KV cache,避免权重同步时显存峰值太高。

这页原本不适合学习的地方

  • 只说了 TrainingWorkerActorRolloutRefWorker 的职责,但没有解释 @register -> WorkerGroup -> Ray remote -> collect 的调用链。
  • 把 dispatch mode 讲成概念,没有说明 make_nd_compute_dataproto_dispatch_fn(mesh_name=...) 如何依赖 worker 注册的 mesh 信息。
  • 缺少 colocate/fused worker 的动机:它不是语义层新角色,而是减少进程、CUDA context 和资源碎片的工程优化。
  • 缺少 update_weights(),导致读者看不懂 actor 训练和 rollout 推理为什么能交替使用同一份 policy。

本节参考与延伸阅读

  • 源码:verl/workers/engine_workers.py
  • 源码:verl/single_controller/base/decorator.py
  • 源码:verl/single_controller/base/worker_group.py
  • 源码:verl/single_controller/ray/base.py
  • 源码:verl/trainer/ppo/ray_trainer.py
  • 官方文档:docs/workers/engine_workers.rst
  • 官方文档:docs/workers/ray_trainer.rst
  • 官方文档:docs/single_controller.rst

面向源码阅读的 verl 学习文档。