Skip to content

Rollout Correction:生成策略和训练策略不一致怎么办

PPO/GRPO 默认假设:这批样本来自你以为的 old policy。但大规模 LLM RL 里,这个假设经常不干净。rollout 可能由 vLLM/SGLang 生成,训练由 FSDP/Megatron forward;rollout worker 的权重可能滞后;异步系统里样本产生时的 step 和训练时的 step 也可能不同。

Rollout Correction 处理的就是这种 mismatch:

text
pi_rollout  生成样本的行为策略
pi_old      PPO clipping 的 anchor/proximal policy
pi_theta    当前正在更新的 actor

mismatch 从哪里来

  • 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 会保留三种策略:

text
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 分工不同:

ρt=πold(at|st)πrollout(at|st)

用于 importance sampling,修正“样本来自 rollout,但训练 anchor 是 old”的差距。

rt(θ)=πθ(at|st)πold(at|st)

用于 PPO clip,控制当前 actor 不要离 old 太远。

RayPPOTrainer.fit() 中,对应顺序是:

text
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,而是:

text
batch.batch["old_log_probs"] = batch.batch["rollout_log_probs"]
policy_loss.loss_mode = "bypass_mode"

也就是:

πold=πrollout

这样 PPO ratio 直接变成:

rt(θ)=πθ(at|st)πrollout(at|st)

优点是省掉一次 old logprob forward;代价是没有把 behavior policy 和 proximal policy 分开。

IS 和 RS:一个是加权,一个是过滤

rollout correction 有两类机制,源码刻意分开:

Importance Sampling

IS 权重是连续权重:

wt=clip(exp(logπoldlogπrollout),0,C)

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

text
modified_response_mask = response_mask * keep_mask

rollout_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,然后构造:

text
token_k1 = -log_ratio
token_k2 = 0.5 * log_ratio ** 2
token_k3 = exp(log_ratio) - 1 - log_ratio

IS 是“这条样本还用,但权重变小/变大”;RS 是“这部分 token 或整条 response 不参与训练”。

源码实现怎么读

1. ray_trainer.py 的分支

核心上下文在 RayPPOTrainer.fit()

text
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_probrollout_log_probs 计算 correction 和 metrics。

2. compute_rollout_correction_and_rejection_mask()

这是 helper 的统一入口:

text
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_problog pi_old 或 bypass 下 log pi_theta修正目标策略
rollout_log_problog pi_rollout行为策略
response_maskmask有效 token
rollout_isIS 粒度token / sequence / None
rollout_rsRS 判据token_k1seq_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

text
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:

text
batch.batch["response_mask"] = modified_response_mask
advantages = compute_advantage(batch, ...)

所以过滤不仅影响 actor loss,也会影响 GRPO/RLOO 等 advantage 的 token 广播和 reward 聚合。读实验结果时要记住:RS 不是只在最后 loss 里遮一下。

配置最小例子

启用 rollout correction 必须让 rollout 计算 logprob:

text
actor_rollout_ref.rollout.calculate_log_probs=True

Decoupled token IS:

text
algorithm.rollout_correction.rollout_is=token
algorithm.rollout_correction.rollout_is_threshold=2.0
algorithm.rollout_correction.bypass_mode=False

Bypass PPO clip:

text
algorithm.rollout_correction.bypass_mode=True
algorithm.rollout_correction.loss_type=ppo_clip
actor_rollout_ref.rollout.calculate_log_probs=True

Bypass REINFORCE with sequence IS:

text
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.shrun_qwen2_5_7b_fsdp_multi_rs.sh

初学者应该盯哪些指标

compute_offpolicy_metrics() 会给出一组诊断:

  • rollout_corr/klKL(pi_rollout || pi_training) 的直接估计。
  • rollout_corr/k3_kl:更稳定的 K3 KL 估计。
  • rollout_corr/training_pplrollout_corr/rollout_ppl:两边对已生成 token 的困惑度。
  • rollout_corr/ppl_ratio:训练策略与 rollout 策略置信度差异。
  • rollout_corr/chi2_tokenrollout_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 学习文档。