Trainer Loop:按一次 batch 的生命线读 RayPPOTrainer.fit()
verl/trainer/ppo/ray_trainer.py 的 RayPPOTrainer.fit() 是学习 verl trainer/data 工程最重要的一条主线。它不是“PPO 公式文件”,而是一个单控制器 driver 把 rollout、reward、logprob、value、advantage、actor/critic update 和 rollout 权重同步串起来的地方。
这一页建议你带着一个问题读:一个 prompt batch 从 dataloader 出来,到 actor 参数被更新,中间到底被哪些对象接住、切开、合并、送去远端 worker?
先验知识
读这段源码前,至少要知道 5 件事:
- RL post-training 的一个训练 step 不是普通的
forward -> loss -> backward。它通常是prompt -> rollout 生成 response -> reward 打分 -> 计算 old/ref logprob 和 value -> 算 advantage -> 更新 actor/critic -> 同步 rollout 权重。 - PPO 需要
old_log_probs作为近端优化锚点,通常还需要ref_log_prob做 KL,GAE 路径还需要 critic 的values。 - GRPO/RLOO 等 group-based 算法不一定要 critic,但需要知道哪些 responses 来自同一个 prompt。verl 用
uid这个 non-tensor 字段做分组。 DataProto是 trainer 和 worker 之间的 batch 信封。tensor 字段在batch,Python/object metadata 在non_tensor_batch,控制信息在meta_info。actor_rollout_wg.compute_log_prob(...)这种调用看起来像本地函数,其实会通过single_controllerdispatch 到 Ray workers,再 collect 回 driver。
本页原先不适合小白的地方
原说明已经讲到了主流程,但有几个学习断点:
- batch 的生命线不够细。小白容易知道“有 rollout、有 reward”,但不知道
batch什么时候 repeat、什么时候 union、什么时候变成TensorDict。 fit()里很多分支没有给阅读优先级,比如 REMAX、rollout correction、profiling、spec decode。初学者应该先抓主路径,再回来看这些增强功能。- 缺少关键字段表。读 trainer 时最容易迷路的是字段名,例如
token_level_scores、token_level_rewards、old_log_probs、advantages分别在哪一步出现。 - 缺少 worker 调用边界。
_compute_old_log_prob()不是在 driver 上算模型 forward,而是把DataProto -> TensorDict -> worker group -> TensorDict -> DataProto绕了一圈。
源码实现怎么读
建议按这个顺序读源码:
- 先读
apply_kl_penalty()和compute_advantage()。这两个函数短,但能告诉你 trainer 期待 batch 里有哪些字段。 - 再读
RayPPOTrainer.init_workers()。先确认 actor/rollout/ref/critic/reward manager/LLM server/checkpoint manager 分别从哪里建起来。 - 然后读
_compute_old_log_prob()、_compute_ref_log_prob()、_compute_values()、_update_actor()、_update_critic()。它们展示了DataProto和TensorDict的边界。 - 最后读
fit()。不要一开始就陷入 profiling、REMAX、rollout correction、dump generation。先用下面的生命线标注主路径。
一次 batch 的生命线
0. 训练开始前:worker 和 rollout server 已经就位
init_workers() 负责创建资源池、worker group 和 rollout/reward/teacher/checkpoint 管理器:
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 和权重同步:
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:
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_ids、attention_mask、position_ids、prompts。- numpy/object 数据进入
batch.non_tensor_batch,例如multi_modal_inputs、data_source、reward_model、extra_info。 uid是 group 算法的分组锚点。后面compute_advantage()的 GRPO 路径会用data.non_tensor_batch["uid"]。
3. 取出生成需要的 gen_batch
_get_gen_batch() 不是简单复制 batch。它会把不适合送去生成端的 non-tensor 字段 pop 掉,但保留 reward 相关字段:
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 后顺序更像:
prompt0, prompt0, ..., prompt1, prompt1, ...这对 GRPO 很自然,因为同一个 uid 的多条 response 相邻出现。
4. rollout 生成 response
主路径调用:
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:
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:
self._balance_batch(batch, metrics=metrics)它根据 attention_mask 估计每条样本 token 工作量,按 data parallel rank 做负载均衡,然后调用 batch.reorder(global_idx)。这会改变 batch 顺序,但不会破坏 GRPO 分组,因为优势计算用的是 uid。
7. reward:先得到分数,再变成训练奖励
reward 阶段有两层:
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_tensortoken_level_scores 是原始 reward 分数。它还不是最终用于 advantage 的 token_level_rewards。
如果开启 reward-side KL:
batch, kl_metrics = apply_kl_penalty(batch, kl_ctrl, kl_penalty)apply_kl_penalty() 做的是:
kld = kl_penalty(old_log_probs, ref_log_prob) * response_mask
token_level_rewards = token_level_scores - beta * kld如果不开启 algorithm.use_kl_in_reward,则直接:
batch.batch["token_level_rewards"] = batch.batch["token_level_scores"]8. old logprob:把 rollout 样本放回 actor 当前权重下重算
PPO 更新前需要 old policy 的 logprob。默认主路径会重算:
old_log_prob, old_log_prob_mfu = self._compute_old_log_prob(batch)
batch = batch.union(old_log_prob)_compute_old_log_prob() 的内部链路更像这样:
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:
ref_log_prob = self._compute_ref_log_prob(batch)
batch = batch.union(ref_log_prob)如果需要 critic:
values = self._compute_values(batch)
batch = batch.union(values)need_reference_policy(config) 通常由 actor.use_kl_loss 或 algorithm.use_kl_in_reward 决定。need_critic(config) 通常和 algorithm.adv_estimator 相关,GAE 需要 critic,GRPO 通常不需要。
10. advantage 在 driver 上算
advantage 是轻量 tensor 计算,直接在 driver process 上做:
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_rewards、values、response_mask算advantages和returns。grpo:用uid按 prompt group 归一化 outcome reward。- 其他 estimator:走
core_algos.get_adv_estimator_fn()注册表,并按需传uid、reward_baselines、sum_pi_squared等字段。
11. update critic,再 update actor
critic 更新:
critic_output = self._update_critic(batch)actor 更新:
actor_output = self._update_actor(batch)两个函数都会先把 DataProto 转成 TensorDict,做 remove-padding 变换,再通过 worker group 调用远端训练方法。actor 里还会把 mini-batch size 乘上 rollout.n:
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 更新后:
if should_save:
self._save_checkpoint()
self.checkpoint_manager.update_weights(self.global_steps)这是一次 batch 生命线的闭环。训练 worker 上 actor 参数更新了,下一次 rollout 必须使用更新后的权重。
主路径伪代码
去掉 profiling、dump、REMAX、rollout correction 后,fit() 主路径可以压缩成:
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)字段出现时间表
| 字段 | 出现位置 | 学习意义 |
|---|---|---|
uid | fit() 创建 batch 后 | 标记同 prompt 多 response 的 group |
responses | rollout 输出 | 模型生成的答案 token |
response_mask | compute_response_mask() | 只在 response token 上算 KL、reward、loss |
rm_scores | reward model 或 reward loop | 模型 reward 的 token 级结果 |
token_level_scores | extract_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_rewards | KL penalty 后 | 真正进入 advantage 的 reward |
advantages / returns | compute_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.rst、docs/examples/ppo_code_architecture.rst、docs/hybrid_flow.rst。 - 论文/网页:HybridFlow: A Flexible and Efficient RLHF Framework, arXiv:2409.19256。