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}}\) 来校正梯度。
Token 级别:
序列级别:
3. 拒绝采样¶
拒绝采样用于过滤掉"太离谱"的样本,支持多种散度度量:
K1 模式(对数比率):
基于比率的阈值: \([\text{lower},\; \text{upper}]\)
K2 模式(平方对数):
上界阈值
K3 模式(KL 估计器):
更稳定的 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. 设计特点¶
核心概念图¶
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 训练中一个关键的工程问题——推理和训练之间的策略不一致。它提供了:
- 重要性采样:用数学方法校正策略差异
- 拒绝采样:丢弃差异过大的异常样本
- 多种模式:Bypass(快速)和 Decoupled(精确)两种操作模式
- 数值稳定:log 空间计算 + 安全截断
- 丰富的预设:通过
RolloutCorrectionConfig的工厂方法提供常用配置
这个模块来自论文 "When Speed Kills Stability: Demystifying RL Collapse from the Training-Inference Mismatch"。