Rollout Correction:生成策略和训练策略不一致怎么办
PPO/GRPO 默认假设:这批样本来自你以为的 old policy。但大规模 LLM RL 里,这个假设经常不干净。rollout 可能由 vLLM/SGLang 生成,训练由 FSDP/Megatron forward;rollout worker 的权重可能滞后;异步系统里样本产生时的 step 和训练时的 step 也可能不同。
Rollout Correction 处理的就是这种 mismatch:
pi_rollout 生成样本的行为策略
pi_old PPO clipping 的 anchor/proximal policy
pi_theta 当前正在更新的 actormismatch 从哪里来
- rollout server 还没同步到最新 actor 权重。
- vLLM/SGLang 推理路径与训练 forward 路径存在精度或实现差异。
- async rollout 中样本来自更旧的 checkpoint。
- replay/off-policy 数据来自历史策略或辅助策略。
- 为吞吐量批量更新 rollout 权重,不是每个 mini-batch 都同步。
如果训练时完全忽略 pi_rollout,就会把“数据来自谁”这件事说错。rollout correction 的核心是:要么显式修正 pi_rollout -> pi_old 的差距,要么干脆把 pi_old 设成 pi_rollout。
三策略框架
Decoupled mode
默认 bypass_mode=False 时,verl 会保留三种策略:
pi_rollout: rollout engine 生成样本,并可返回 rollout_log_probs
pi_old: actor.compute_log_prob(batch) 重算 old_log_probs
pi_theta: actor update 中当前 forward 得到 log_prob两个 ratio 分工不同:
用于 importance sampling,修正“样本来自 rollout,但训练 anchor 是 old”的差距。
用于 PPO clip,控制当前 actor 不要离 old 太远。
在 RayPPOTrainer.fit() 中,对应顺序是:
rollout returns rollout_log_probs # 需要 calculate_log_probs=True
old_log_probs = actor.compute_log_prob(batch)
batch = compute_rollout_correction_and_add_to_batch(batch, rollout_corr_config)
advantages = compute_advantage(...)
actor.update_actor(batch)Bypass mode
bypass_mode=True 时,verl 不再重算 old_log_probs,而是:
batch.batch["old_log_probs"] = batch.batch["rollout_log_probs"]
policy_loss.loss_mode = "bypass_mode"也就是:
这样 PPO ratio 直接变成:
优点是省掉一次 old logprob forward;代价是没有把 behavior policy 和 proximal policy 分开。
IS 和 RS:一个是加权,一个是过滤
rollout correction 有两类机制,源码刻意分开:
Importance Sampling
IS 权重是连续权重:
在 rollout_corr_helper.compute_rollout_correction_weights() 中,rollout_is 决定粒度:
token:每个 token 一个 ratio。sequence:把一条 response 的 log ratio 求和后 exponentiate,再广播回 token。
rollout_is_threshold 会截断极端权重;rollout_is_batch_normalize=True 会把 batch 内平均权重归一到 1,降低随机尺度波动。
Rejection Sampling
RS 是二值过滤,直接修改 response_mask:
modified_response_mask = response_mask * keep_maskrollout_rs 支持:
token_k1/k2/k3:token 级过滤。seq_sum_k1/k2/k3:按整条 response 的总 divergence 过滤。seq_mean_k1/k2/k3:按平均 divergence 过滤,对长度更友好。seq_max_k2/k3:只要某个 token 太离谱就过滤整条。
其中 k1/k2/k3 对应不同 KL/divergence 估计。源码里 log_ratio = old_log_prob - rollout_log_prob,然后构造:
token_k1 = -log_ratio
token_k2 = 0.5 * log_ratio ** 2
token_k3 = exp(log_ratio) - 1 - log_ratioIS 是“这条样本还用,但权重变小/变大”;RS 是“这部分 token 或整条 response 不参与训练”。
源码实现怎么读
1. ray_trainer.py 的分支
核心上下文在 RayPPOTrainer.fit():
rollout_corr_config = self.config.algorithm.get("rollout_correction", None)
bypass = rollout_corr_config and rollout_corr_config.get("bypass_mode", False)
if bypass:
apply_bypass_mode(batch, rollout_corr_config, actor.policy_loss)
else:
old_log_prob = self._compute_old_log_prob(batch)
batch = batch.union(old_log_prob)
if rollout_corr_config and "rollout_log_probs" in batch and not bypass:
batch, metrics = compute_rollout_correction_and_add_to_batch(batch, rollout_corr_config)也就是说:
- decoupled mode:driver 侧先算一次 IS/RS,并把
rollout_is_weights和修改后的response_mask放进 batch。 - bypass mode:actor loss 内部用当前
log_prob对rollout_log_probs计算 correction 和 metrics。
2. compute_rollout_correction_and_rejection_mask()
这是 helper 的统一入口:
log_ratio = old_log_prob - rollout_log_prob
rollout_is_weights = compute_rollout_correction_weights(...)
modified_response_mask = compute_rollout_rejection_mask(...)
metrics = compute_offpolicy_metrics(...)输入变量对应:
| 源码变量 | 符号 | 含义 |
|---|---|---|
old_log_prob | log pi_old 或 bypass 下 log pi_theta | 修正目标策略 |
rollout_log_prob | log pi_rollout | 行为策略 |
response_mask | mask | 有效 token |
rollout_is | IS 粒度 | token / sequence / None |
rollout_rs | RS 判据 | token_k1、seq_mean_k3 等 |
返回值有三类:
rollout_is_weights_proto:包含rollout_is_weights,如果没开 IS 则为None。modified_response_mask:应用 RS 后的新 mask。metrics:统一带rollout_corr/前缀。
3. compute_policy_loss_bypass_mode()
bypass mode 的 actor loss 注册名是 "bypass_mode"。它内部再看 loss_type:
loss_type == "ppo_clip":
old_log_prob = rollout_log_prob
compute_policy_loss_vanilla(...)
rollout_is_weights=None # PPO ratio 已经是 pi_theta / pi_rollout,不再重复乘 IS
loss_type == "reinforce":
compute_policy_loss_reinforce(...)
pg_losses = -advantages * log_prob * rollout_is_weights这里最容易误读的是:bypass + PPO clip 不显式乘 IS 权重,因为 ratio 本身已经直接比较 pi_theta / pi_rollout。再乘一次会 double count。
4. rollout correction 与 advantage 的顺序
decoupled mode 下,RS 会先修改 response_mask,再算 advantage:
batch.batch["response_mask"] = modified_response_mask
advantages = compute_advantage(batch, ...)所以过滤不仅影响 actor loss,也会影响 GRPO/RLOO 等 advantage 的 token 广播和 reward 聚合。读实验结果时要记住:RS 不是只在最后 loss 里遮一下。
配置最小例子
启用 rollout correction 必须让 rollout 计算 logprob:
actor_rollout_ref.rollout.calculate_log_probs=TrueDecoupled token IS:
algorithm.rollout_correction.rollout_is=token
algorithm.rollout_correction.rollout_is_threshold=2.0
algorithm.rollout_correction.bypass_mode=FalseBypass PPO clip:
algorithm.rollout_correction.bypass_mode=True
algorithm.rollout_correction.loss_type=ppo_clip
actor_rollout_ref.rollout.calculate_log_probs=TrueBypass REINFORCE with sequence IS:
algorithm.rollout_correction.bypass_mode=True
algorithm.rollout_correction.loss_type=reinforce
algorithm.rollout_correction.rollout_is=sequence
algorithm.rollout_correction.rollout_is_batch_normalize=True示例脚本在 examples/rollout_correction/run_qwen2_5_7b_fsdp.sh 和 run_qwen2_5_7b_fsdp_multi_rs.sh。
初学者应该盯哪些指标
compute_offpolicy_metrics() 会给出一组诊断:
rollout_corr/kl:KL(pi_rollout || pi_training)的直接估计。rollout_corr/k3_kl:更稳定的 K3 KL 估计。rollout_corr/training_ppl与rollout_corr/rollout_ppl:两边对已生成 token 的困惑度。rollout_corr/ppl_ratio:训练策略与 rollout 策略置信度差异。rollout_corr/chi2_token、rollout_corr/chi2_seq:IS 权重二阶矩,越大说明方差风险越大。rollout_corr/rollout_is_eff_sample_size:有效样本量,过低说明少数样本权重过大。rollout_corr/rollout_rs_masked_fraction:RS 过滤了多少 token。
这些指标比单看 reward 更能告诉你:训练是不是在你以为的分布上发生。
哪些地方不适合初学者硬啃
- 不要把
old_log_probs自动等同于rollout_log_probs。只有 bypass mode 明确这样设。 - 不要在 bypass + PPO clip 里再手动乘 IS 权重,源码已经避免 double counting。
- 不要只开
algorithm.rollout_correction却忘了actor_rollout_ref.rollout.calculate_log_probs=True。 - 不要把 RS 当成无副作用过滤。它会改变
response_mask,从而影响 advantage 和 loss。 - 不要把
rollout_corr/kl的正负号和 reference KL 混淆。这里比较的是 rollout policy 与 training/old/current policy,不是 actor 与 ref policy。
本节参考与延伸阅读
- verl 源码:
verl/trainer/ppo/rollout_corr_helper.py完整文件,尤其compute_rollout_correction_and_rejection_mask()、compute_rollout_correction_weights()、compute_rollout_rejection_mask()、compute_offpolicy_metrics()、apply_bypass_mode()。 - verl 源码:
verl/trainer/ppo/core_algos.py的compute_policy_loss_bypass_mode()、compute_policy_loss_reinforce()、compute_policy_loss_vanilla()。 - verl 源码:
verl/trainer/ppo/ray_trainer.py中 rollout correction、old logprob、KL、advantage 的主循环上下文。 - verl 官方文档:
docs/algo/rollout_corr.md、docs/algo/rollout_corr_math.md。 - 示例脚本:
examples/rollout_correction/run_qwen2_5_7b_fsdp.sh、examples/rollout_correction/run_qwen2_5_7b_fsdp_multi_rs.sh。 - 论文/网页:When Speed Kills Stability: Demystifying RL Collapse from the Training-Inference Mismatch。
- 论文:Trust Region Masking for Long-Horizon LLM Reinforcement Learning。
- 相关理论:Batch size-invariance for policy optimization。