跳转至

rollout_corr_helper.py — 该文件实现了 Rollout 校正(Rollout Correction) 模块

文件概述

模块路径: verl.trainer.ppo.rollout_corr_helper

该文件实现了 Rollout 校正(Rollout Correction) 模块,用于解决 RL 训练中的 off-policy 问题。

Off-policy 问题主要来源于三个方面: 1. 策略不匹配:推理引擎(如 vLLM,BFloat16)和训练引擎(如 FSDP,FP32)的精度差异 2. 模型更新滞后:训练数据来自旧版本的模型 3. 一般分布偏移:数据收集和训练之间的分布差异

在训练流程中的位置

在 ray_trainer.py 的 fit() 中,计算优势之前调用,处理 rollout 策略和训练策略之间的差异。

关键代码讲解

1. 核心能力概述

该模块提供三类功能:

1. 重要性采样 (IS, Importance Sampling):
   - Token 级别: 每个 token 有独立的权重
   - 序列级别: 整条序列一个权重

2. 拒绝采样 (RS, Rejection Sampling):
   - 基于散度的过滤: 丢弃与当前策略差异过大的样本
   - 支持多种散度度量 (K1, K2, K3)

3. 指标跟踪:
   - KL 散度、PPL 比率等 off-policy 程度的度量

2. 重要性采样权重计算

重要性采样的核心思想:当训练数据来自策略 \(\pi_{\text{old}}\) 而非当前策略 \(\pi_\theta\) 时,需要用比率 \(\pi_\theta / \pi_{\text{old}}\) 来校正梯度。

\[ w = \frac{\pi_{\text{train}}(a|s)}{\pi_{\text{rollout}}(a|s)} \]

Token 级别:

\[ w_t = \exp\left(\log \pi_{\text{train},t} - \log \pi_{\text{rollout},t}\right) \]

序列级别:

\[ w = \exp\left(\sum_t \left(\log \pi_{\text{train},t} - \log \pi_{\text{rollout},t}\right)\right) \]

3. 拒绝采样

拒绝采样用于过滤掉"太离谱"的样本,支持多种散度度量:

K1 模式(对数比率):

\[ d_t = -\log r_t = \log \pi_{\text{rollout},t} - \log \pi_{\text{train},t} \]

基于比率的阈值: \([\text{lower},\; \text{upper}]\)

K2 模式(平方对数):

\[ d_t = \frac{1}{2}(\log r_t)^2 \]

上界阈值

K3 模式(KL 估计器):

\[ d_t = \exp(\log r_t) - 1 - \log r_t = r_t - 1 - \log r_t \]

更稳定的 KL 估计,\(K_3 \geq 0\) 总是非负

聚合方式包括: - token_*: Token 级别过滤 - seq_sum_*: 序列求和过滤 - seq_mean_*: 序列求平均过滤(长度归一化) - seq_max_*: 序列取最大值过滤

4. Bypass 模式和 Decoupled 模式

Bypass 模式 (2 策略):
  - π_rollout = π_old (直接复用 rollout 的 log_prob)
  - 训练时的比率 r = π_θ / π_rollout
  - 节省一次 log_prob 的重计算

Decoupled 模式 (3 策略):
  - π_rollout: 生成数据时的策略
  - π_old: 训练开始时重新计算的策略(proximal anchor)
  - π_θ: 训练中不断更新的策略
  - IS 权重校正 π_old 和 π_rollout 之间的差距

5. 设计特点

# 数值稳定性设计
# 1. 在 log 空间计算,避免上溢/下溢
# 2. 固定安全边界 exp(±20) 用于稳定指数运算
# 3. 指标计算不使用大型中间张量(防止 GPU OOM)

核心概念图

                    Rollout Policy (pi_rollout)
                    [vLLM BF16, 旧参数]
                           |
                      生成数据
                           |
                           v
                    +-------------+
                    | Trajectory  |  rollout_log_probs
                    | (s, a, r)   |
                    +-------------+
                           |
          +----------------+----------------+
          |                                 |
     Decoupled 模式                    Bypass 模式
          |                                 |
    pi_old = 重新计算             pi_old = rollout_log_probs
    (FP32 精度)                    (节省计算)
          |                                 |
    IS权重 = pi_old/pi_rollout     r = pi_theta/pi_rollout
          |                                 |
    + 拒绝采样过滤异常样本           + 拒绝采样过滤
          |                                 |
          +----------------+----------------+
                           |
                      PPO 策略更新

核心类/函数列表

名称 类型 作用
compute_rollout_correction_and_add_to_batch() function 计算 IS 权重和拒绝采样掩码,添加到 batch
apply_bypass_mode() function 应用 Bypass 模式(设置 old_log_probs = rollout_log_probs)
compute_is_weights() function 计算重要性采样权重
compute_rejection_mask() function 计算拒绝采样掩码
compute_offpolicy_metrics() function 计算 off-policy 相关指标

数据流和调用关系

ray_trainer.py: fit()
    |
    +-- (Bypass 模式)
    |   +-- apply_bypass_mode(batch, rollout_corr_config)
    |       +-- batch["old_log_probs"] = batch["rollout_log_probs"]
    |       +-- 设置 loss_type (ppo_clip / reinforce)
    |
    +-- (Decoupled 模式)
        +-- _compute_old_log_prob()  --> batch["old_log_probs"]
        +-- compute_rollout_correction_and_add_to_batch(batch, config)
            |
            +-- compute_is_weights()
            |   +-- Token 级别: exp(old_log_probs - rollout_log_probs)
            |   +-- 序列级别: exp(Σ(old_log_probs - rollout_log_probs))
            |   +-- 截断 + 可选的 batch 归一化
            |
            +-- compute_rejection_mask()
            |   +-- 基于散度的过滤(K1/K2/K3)
            |   +-- 聚合模式(token/seq_sum/seq_mean/seq_max)
            |
            +-- compute_offpolicy_metrics()
            |   +-- KL 散度、PPL 比率等
            |
            +-- batch["rollout_is_weights"] = is_weights
            +-- batch["response_mask"] *= rejection_mask

小结

rollout_corr_helper.py 解决了 RL 训练中一个关键的工程问题——推理和训练之间的策略不一致。它提供了:

  1. 重要性采样:用数学方法校正策略差异
  2. 拒绝采样:丢弃差异过大的异常样本
  3. 多种模式:Bypass(快速)和 Decoupled(精确)两种操作模式
  4. 数值稳定:log 空间计算 + 安全截断
  5. 丰富的预设:通过 RolloutCorrectionConfig 的工厂方法提供常用配置

这个模块来自论文 "When Speed Kills Stability: Demystifying RL Collapse from the Training-Inference Mismatch"。