Skip to content

Trainer Loop:按一次 batch 的生命线读 RayPPOTrainer.fit()

verl/trainer/ppo/ray_trainer.pyRayPPOTrainer.fit() 是学习 verl trainer/data 工程最重要的一条主线。它不是“PPO 公式文件”,而是一个单控制器 driver 把 rollout、reward、logprob、value、advantage、actor/critic update 和 rollout 权重同步串起来的地方。

这一页建议你带着一个问题读:一个 prompt batch 从 dataloader 出来,到 actor 参数被更新,中间到底被哪些对象接住、切开、合并、送去远端 worker?

先验知识

读这段源码前,至少要知道 5 件事:

  1. RL post-training 的一个训练 step 不是普通的 forward -> loss -> backward。它通常是 prompt -> rollout 生成 response -> reward 打分 -> 计算 old/ref logprob 和 value -> 算 advantage -> 更新 actor/critic -> 同步 rollout 权重
  2. PPO 需要 old_log_probs 作为近端优化锚点,通常还需要 ref_log_prob 做 KL,GAE 路径还需要 critic 的 values
  3. GRPO/RLOO 等 group-based 算法不一定要 critic,但需要知道哪些 responses 来自同一个 prompt。verl 用 uid 这个 non-tensor 字段做分组。
  4. DataProto 是 trainer 和 worker 之间的 batch 信封。tensor 字段在 batch,Python/object metadata 在 non_tensor_batch,控制信息在 meta_info
  5. actor_rollout_wg.compute_log_prob(...) 这种调用看起来像本地函数,其实会通过 single_controller dispatch 到 Ray workers,再 collect 回 driver。

本页原先不适合小白的地方

原说明已经讲到了主流程,但有几个学习断点:

  • batch 的生命线不够细。小白容易知道“有 rollout、有 reward”,但不知道 batch 什么时候 repeat、什么时候 union、什么时候变成 TensorDict
  • fit() 里很多分支没有给阅读优先级,比如 REMAX、rollout correction、profiling、spec decode。初学者应该先抓主路径,再回来看这些增强功能。
  • 缺少关键字段表。读 trainer 时最容易迷路的是字段名,例如 token_level_scorestoken_level_rewardsold_log_probsadvantages 分别在哪一步出现。
  • 缺少 worker 调用边界。_compute_old_log_prob() 不是在 driver 上算模型 forward,而是把 DataProto -> TensorDict -> worker group -> TensorDict -> DataProto 绕了一圈。

源码实现怎么读

建议按这个顺序读源码:

  1. 先读 apply_kl_penalty()compute_advantage()。这两个函数短,但能告诉你 trainer 期待 batch 里有哪些字段。
  2. 再读 RayPPOTrainer.init_workers()。先确认 actor/rollout/ref/critic/reward manager/LLM server/checkpoint manager 分别从哪里建起来。
  3. 然后读 _compute_old_log_prob()_compute_ref_log_prob()_compute_values()_update_actor()_update_critic()。它们展示了 DataProtoTensorDict 的边界。
  4. 最后读 fit()。不要一开始就陷入 profiling、REMAX、rollout correction、dump generation。先用下面的生命线标注主路径。

一次 batch 的生命线

0. 训练开始前:worker 和 rollout server 已经就位

init_workers() 负责创建资源池、worker group 和 rollout/reward/teacher/checkpoint 管理器:

text
ResourcePoolManager.create_resource_pool()
  -> create_colocated_worker_cls(...)
  -> RayWorkerGroup(...)
  -> wg_dict.spawn(...)
  -> actor_rollout_wg.init_model()
  -> RewardLoopManager(...)
  -> LLMServerManager.create(...)
  -> AgentLoopManager.create(...)
  -> CheckpointEngineManager(...)
  -> checkpoint_manager.sleep_replicas()

这一步说明:fit() 开始时并不是临时创建模型。模型 worker、rollout server、reward loop manager 都已经挂在 trainer 上了。

1. checkpoint 和初始 rollout 权重

fit() 一开始做 logger、checkpoint 和权重同步:

python
self.global_steps = 0
self._load_checkpoint()
self.checkpoint_manager.update_weights(self.global_steps)

这里的 update_weights(0) 很关键。actor 权重在训练 worker 上,rollout server 用来生成 response。第一批 rollout 之前必须让 rollout server 拿到初始 actor 权重。

2. dataloader batch 进入 DataProto

每次循环从 self.train_dataloader 取出 batch_dict

python
batch = DataProto.from_single_dict(batch_dict)
batch.meta_info["temperature"] = config.actor_rollout_ref.rollout.temperature
batch.non_tensor_batch["uid"] = np.array([...], dtype=object)

这一步的含义:

  • batch_dict 里 torch tensor 进入 batch.batch,例如 input_idsattention_maskposition_idsprompts
  • numpy/object 数据进入 batch.non_tensor_batch,例如 multi_modal_inputsdata_sourcereward_modelextra_info
  • uid 是 group 算法的分组锚点。后面 compute_advantage() 的 GRPO 路径会用 data.non_tensor_batch["uid"]

3. 取出生成需要的 gen_batch

_get_gen_batch() 不是简单复制 batch。它会把不适合送去生成端的 non-tensor 字段 pop 掉,但保留 reward 相关字段:

python
gen_batch = self._get_gen_batch(batch)
gen_batch.meta_info["global_steps"] = self.global_steps
gen_batch_output = gen_batch.repeat(repeat_times=rollout_n, interleave=True)

如果 rollout.n=8,一个 prompt 会复制成 8 条生成请求。interleave=True 后顺序更像:

text
prompt0, prompt0, ..., prompt1, prompt1, ...

这对 GRPO 很自然,因为同一个 uid 的多条 response 相邻出现。

4. rollout 生成 response

主路径调用:

python
combined_gen_output = self.async_rollout_manager.generate_sequences(combined_gen_batch)
self.checkpoint_manager.sleep_replicas()

现在源码里默认是 async rollout manager。rollout.mode: async 是配置默认值,init_workers() 里也直接设置 self.async_rollout_mode = True。生成完成后 sleep replicas 是显存动作:rollout server 生成时占用 KV cache 和推理显存,训练 actor/critic 前要把空间让出来。

生成输出通常带回:

  • responses
  • 新的 attention_mask
  • 可能的 rollout_log_probs
  • 多轮或工具调用相关 non-tensor 信息
  • meta_info["timing"]

5. 原 batch 对齐 response 数,再 union()

生成输出长度是 原 prompt 数 * rollout.n。所以原 batch 也要 repeat:

python
batch = batch.repeat(repeat_times=rollout_n, interleave=True)
batch = batch.union(gen_batch_output)
batch.batch["response_mask"] = compute_response_mask(batch)

这里是理解 DataProto 的核心例子:原 batch 有 prompt、uid、reward metadata;rollout 输出有 response、attention mask 等生成结果。union() 把两个同 batch size 的信封合成一个。

6. 可选:按 token 数重排 batch

如果 trainer.balance_batch=True

python
self._balance_batch(batch, metrics=metrics)

它根据 attention_mask 估计每条样本 token 工作量,按 data parallel rank 做负载均衡,然后调用 batch.reorder(global_idx)。这会改变 batch 顺序,但不会破坏 GRPO 分组,因为优势计算用的是 uid

7. reward:先得到分数,再变成训练奖励

reward 阶段有两层:

python
if self.use_rm and "rm_scores" not in batch.batch:
    batch_reward = self._compute_reward_colocate(batch)
    batch = batch.union(batch_reward)

reward_tensor, reward_extra_infos_dict = extract_reward(batch)
batch.batch["token_level_scores"] = reward_tensor

token_level_scores 是原始 reward 分数。它还不是最终用于 advantage 的 token_level_rewards

如果开启 reward-side KL:

python
batch, kl_metrics = apply_kl_penalty(batch, kl_ctrl, kl_penalty)

apply_kl_penalty() 做的是:

python
kld = kl_penalty(old_log_probs, ref_log_prob) * response_mask
token_level_rewards = token_level_scores - beta * kld

如果不开启 algorithm.use_kl_in_reward,则直接:

python
batch.batch["token_level_rewards"] = batch.batch["token_level_scores"]

8. old logprob:把 rollout 样本放回 actor 当前权重下重算

PPO 更新前需要 old policy 的 logprob。默认主路径会重算:

python
old_log_prob, old_log_prob_mfu = self._compute_old_log_prob(batch)
batch = batch.union(old_log_prob)

_compute_old_log_prob() 的内部链路更像这样:

python
batch_td = batch.to_tensordict()
batch_td = left_right_2_no_padding(batch_td)
tu.assign_non_tensor(batch_td, calculate_entropy=True, compute_loss=False)
output = self.actor_rollout_wg.compute_log_prob(batch_td)
log_probs = no_padding_2_padding(tu.get(output, "log_probs"), batch_td)
return DataProto.from_tensordict({"old_log_probs": log_probs})

小白读这里要注意:模型 forward 在 worker 上,driver 只是准备 TensorDict、发 RPC、接回结果。

如果开启 rollout correction 的 bypass mode,则可能直接用 rollout 端返回的 rollout_log_probs 作为 old_log_probs

9. ref logprob 和 values

如果需要 reference policy:

python
ref_log_prob = self._compute_ref_log_prob(batch)
batch = batch.union(ref_log_prob)

如果需要 critic:

python
values = self._compute_values(batch)
batch = batch.union(values)

need_reference_policy(config) 通常由 actor.use_kl_lossalgorithm.use_kl_in_reward 决定。need_critic(config) 通常和 algorithm.adv_estimator 相关,GAE 需要 critic,GRPO 通常不需要。

10. advantage 在 driver 上算

advantage 是轻量 tensor 计算,直接在 driver process 上做:

python
batch = compute_advantage(
    batch,
    adv_estimator=config.algorithm.adv_estimator,
    gamma=config.algorithm.gamma,
    lam=config.algorithm.lam,
    num_repeat=config.actor_rollout_ref.rollout.n,
    config=config.algorithm,
)

compute_advantage() 分流:

  • gae:用 token_level_rewardsvaluesresponse_maskadvantagesreturns
  • grpo:用 uid 按 prompt group 归一化 outcome reward。
  • 其他 estimator:走 core_algos.get_adv_estimator_fn() 注册表,并按需传 uidreward_baselinessum_pi_squared 等字段。

11. update critic,再 update actor

critic 更新:

python
critic_output = self._update_critic(batch)

actor 更新:

python
actor_output = self._update_actor(batch)

两个函数都会先把 DataProto 转成 TensorDict,做 remove-padding 变换,再通过 worker group 调用远端训练方法。actor 里还会把 mini-batch size 乘上 rollout.n

python
ppo_mini_batch_size = actor.ppo_mini_batch_size * rollout.n

所以配置里的 actor.ppo_mini_batch_size 是 prompt 视角还是 response 视角,读源码时要非常小心。当前 trainer 在 update 前乘了 rollout.n,让 worker 看到放大后的 response batch。

12. save checkpoint,然后同步 rollout 权重

actor 更新后:

python
if should_save:
    self._save_checkpoint()

self.checkpoint_manager.update_weights(self.global_steps)

这是一次 batch 生命线的闭环。训练 worker 上 actor 参数更新了,下一次 rollout 必须使用更新后的权重。

主路径伪代码

去掉 profiling、dump、REMAX、rollout correction 后,fit() 主路径可以压缩成:

python
for batch_dict in train_dataloader:
    batch = DataProto.from_single_dict(batch_dict)
    batch.non_tensor_batch["uid"] = make_uid(len(batch))

    gen_batch = _get_gen_batch(batch)
    gen_batch = gen_batch.repeat(rollout.n, interleave=True)
    gen_output = async_rollout_manager.generate_sequences(gen_batch)
    checkpoint_manager.sleep_replicas()

    batch = batch.repeat(rollout.n, interleave=True)
    batch = batch.union(gen_output)
    batch.batch["response_mask"] = compute_response_mask(batch)

    reward_tensor, reward_info = extract_reward(batch)
    old_log_prob = _compute_old_log_prob(batch)
    batch = batch.union(old_log_prob)

    if use_reference_policy:
        batch = batch.union(_compute_ref_log_prob(batch))
    if use_critic:
        batch = batch.union(_compute_values(batch))

    batch.batch["token_level_scores"] = reward_tensor
    if algorithm.use_kl_in_reward:
        batch, kl_metrics = apply_kl_penalty(batch, kl_ctrl)
    else:
        batch.batch["token_level_rewards"] = batch.batch["token_level_scores"]

    batch = compute_advantage(batch, algorithm.adv_estimator, ...)

    if use_critic:
        _update_critic(batch)
    if global_steps >= trainer.critic_warmup:
        _update_actor(batch)
        checkpoint_manager.update_weights(global_steps)

字段出现时间表

字段出现位置学习意义
uidfit() 创建 batch 后标记同 prompt 多 response 的 group
responsesrollout 输出模型生成的答案 token
response_maskcompute_response_mask()只在 response token 上算 KL、reward、loss
rm_scoresreward model 或 reward loop模型 reward 的 token 级结果
token_level_scoresextract_reward()未扣 KL 的 reward 分数
old_log_probs_compute_old_log_prob()PPO ratio 的分母或近端锚点
ref_log_prob_compute_ref_log_prob()KL reward 或 KL loss 的 reference
values_compute_values()GAE/critic 路径使用
token_level_rewardsKL penalty 后真正进入 advantage 的 reward
advantages / returnscompute_advantage()actor loss 和 critic loss 的训练目标

本节参考与延伸阅读

  • 源码:verl/trainer/ppo/ray_trainer.py,重点读 RayPPOTrainer.fit()init_workers()_validate()apply_kl_penalty()compute_advantage()
  • 源码:verl/protocol.py,理解 DataProto.from_single_dict()union()repeat()reorder()to_tensordict()
  • 源码:verl/workers/engine_workers.py,理解 ActorRolloutRefWorker.compute_log_prob()compute_ref_log_prob()update_actor()TrainingWorker.train_mini_batch()
  • 官方 docs:docs/workers/ray_trainer.rstdocs/examples/ppo_code_architecture.rstdocs/hybrid_flow.rst
  • 论文/网页:HybridFlow: A Flexible and Efficient RLHF Framework, arXiv:2409.19256。

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