NLP——LLM对齐微调-VeRL-rollout_corr_math解读

注:本文包含 AI 辅助创作


整体总结

  • 算法:从 REINFORCEPPO 再到 Decoupled PPO
  • Rollout 校正提供了一个统一框架,用于处理 RL 训练中的 一般性 off-policy 问题
  • 任何数据收集分布与训练分布不同的场景,适用场景包括:
    • Policy mismatch :不同精度(FP8 vs FP16 vs BF16 vs FP32),不同后端(vLLM vs SGLang vs FSDP vs Megatron)
    • Temporal lag :模型陈旧性,异步 Rollout 工作节点
    • Replay buffers :使用来自更早策略版本的轨迹进行训练
    • Off-policy 算法 :行为克隆,DAPO,专家演示
    • Data filtering :重加权,偏好学习(preference learning),课程学习
    • 补充理解:一些随机算子、前向操作误差等导致的不一致也可以在这里被修正

理论基础:从 REINFORCE 到 Decoupled PPO

REINFORCE:策略梯度基线

  • REINFORCE 算法(1992)是策略梯度方法的基础
  • 原始 REINFORCE(On-Policy)
    • 对于从当前策略 \(\pi_\theta\) 采样的轨迹 \(\tau = (s_0, a_0, s_1, a_1, \ldots, s_T, a_T)\),策略梯度为:
      $$
      \nabla_\theta J(\theta) = \mathbb{E}_{\tau \sim \pi_\theta} \left[ \sum_{t=0}^T \nabla_\theta \log \pi_\theta(a_t|s_t) \cdot A_t \right]
      $$
      • \(A_t\) 是时间步 \(t\) 的优势函数
  • Off-Policy REINFORCE
    • 当轨迹从不同的行为策略 \(\mu\) 采样时,作者对 联合轨迹分布 应用重要性采样:
      $$
      \nabla_\theta J(\theta) = \mathbb{E}_{\tau \sim \mu} \left[ \frac{P_{\pi_\theta}(\tau)}{P_\mu(\tau)} \sum_{t=0}^T \nabla_\theta \log \pi_\theta(a_t|s_t) \cdot A_t \right]
      $$
    • 其中轨迹级重要性权重为:
      $$
      \frac{P_{\pi_\theta}(\tau)}{P_\mu(\tau)} = \frac{p(s_0) \prod_{t=0}^T \pi_\theta(a_t|s_t) p(s_{t+1}|s_t, a_t)}{p(s_0) \prod_{t=0}^T \mu(a_t|s_t) p(s_{t+1}|s_t, a_t)} = \prod_{t=0}^T \frac{\pi_\theta(a_t|s_t)}{\mu(a_t|s_t)}
      $$
      • 转移动态 \(p(s_{t+1}|s_t, a_t)\) 和初始状态 \(p(s_0)\) 相互抵消,只留下每步动作概率比值的乘积
  • REINFORCE 的特点
    • 支持 Off-policy :可以通过重要性采样从任何行为策略学习
    • 无 Trust region :策略更新不受约束
  • PS:verl 中的实现: bypass_pg_is 预设实现了带截断重要性采样的 Off-policy REINFORCE

PPO:加入信任域控制

  • 近端策略优化(2017)增加了裁剪的替代目标:
    $$
    L_{\text{PPO} }(\theta) = -\mathbb{E}_{(s,a) \sim \mu} \left[ \min\left( r_t(\theta) A_t, \text{clip}(r_t(\theta), 1-\epsilon, 1+\epsilon) A_t \right) \right]
    $$
    • \(r_t(\theta) = \frac{\pi_\theta(a_t|s_t)}{\mu(a_t|s_t)}\),\(\epsilon\) 是裁剪范围(通常为 0.2)
  • PPO 的特点:
    • 两个策略 :\(\mu\)(用于裁剪的参考)和 \(\pi_\theta\)(被更新的策略)
    • 通过裁剪实现信任域 :通过比值 \(r_t(\theta) = \frac{\pi_\theta}{\mu}\) 限制策略更新幅度

Decoupled PPO :实现批次大小不变性(Batch Size Invariance)

  • 这里的 “批次大小不变性”,是指 Decoupled PPO 方法的一个关键特性
  • Batch Size Invariance 的定义:如果一个算法在 Batch Size 发生改变时,可以通过修改其他超参数(如 LR 等)来近似恢复原始 Batch Size 的训练行为,则称这个算法满足 Batch Size Invariance
  • 标准 PPO 通过比值 \(\frac{\pi_\theta}{\pi_{\text{old} } }\) 控制策略更新大小,其中 \(\pi_{\text{old} }\) 被假定同时是近端策略 和 行为策略
    • 这种耦合使得近端策略一直随着行为策略在变化,无法实现 Batch Size Invariance
      • 因为当 Rollout Batch Size 发生变化时,近端策略更新的频率也在发生变化,这个是无法恢复的
  • Decoupled PPO(Hilton,2021) 通过 Decoupled 两个角色 解决了 PPO 对批次大小的敏感性:
    • 1)Proximal policy \(\pi_{\text{prox} }\):PPO 裁剪的锚点策略(控制策略更新大小)
    • 2)Behavior policy \(\mu\):收集数据的策略(用于通过重要性采样进行 off-policy 校正)
    • 具体来说:Decoupled PPO 论文通过 COM(Center of Mass)来调节 Proximal policy 策略回看的步数(Batch Size 变小,则提升回看步数,几乎保证 Off-policy 幅度差不多)
  • Decoupled 这两个角色,得到 三策略公式
    $$
    L_{\text{DecoupledPPO} }(\theta) = -\mathbb{E}_{(s,a) \sim \mu} \left[ w_t \cdot \min\left( r_t(\theta) A_t, \text{clip}(r_t(\theta), 1-\epsilon, 1+\epsilon) A_t \right) \right]
    $$
    • 其中:
      • \(w_t = \frac{\pi_{\text{prox} }(a_t|s_t)}{\mu(a_t|s_t)}\):重要性采样权重(校正行为策略 \(\mu\))
        • 这里 \(\pi_{\text{prox} }\) 在训练期间是冻结的,因此 \(w_t\) 是常数(不需要 stopgrad 操作符)
      • \(r_t(\theta) = \frac{\pi_\theta(a_t|s_t)}{\pi_{\text{prox} }(a_t|s_t)}\):PPO 比值(控制相对于近端策略 \(\pi_{\text{prox} }\) 的策略更新大小)
  • 通过 Decoupled 实现了:
    • 批次大小不变性 :策略更新控制(通过 \(\pi_{\text{prox} }\))独立于数据聚合方式(理解:因为 Decoupled PPO 可以做到更新独立于数据 Batch Size)
    • 灵活的行为策略 :可以使用任意 \(\mu\)(不同工作节点、回放缓冲或陈旧 checkpoint)
    • 陈旧数据利用 :可通过重要性采样校正更早的轨迹
    • 保留裁剪 :针对 \(\pi_{\text{prox} }\) 的裁剪限制了更新幅度

verl 中的实现:三策略框架

  • verl 库使用三种不同的策略实现 Decoupled PPO ,每种策略承担特定角色
  • 符号汇总
    • \(\pi_{\text{rollout} }\):行为策略(数据收集)
    • \(\pi_{\text{old} }\):近端策略(PPO 锚点)
    • \(\pi_{\theta}\):当前策略(被更新的策略)
    • \(\rho_t = \frac{\pi_{\text{old} }(a_t|s_t)}{\pi_{\text{rollout} }(a_t|s_t)}\):每 Token 的 IS 比值(校正 Drift 1)
    • \(r_t(\theta) = \frac{\pi_{\theta}(a_t|s_t)}{\pi_{\text{old} }(a_t|s_t)}\):PPO 比值(校正 Drift 2)
    • \(A_t\):Token \(t\) 处的优势
    • \(T\):序列中的有效 Token 集合
    • \(C_{\text{IS} }\):IS 权重的上阈值(例如 2.0)
    • \(C_{\text{RS-upper} }\):RS 掩码的上阈值(例如 2.0)
    • \(C_{\text{RS-lower} }\):RS 掩码的下阈值(通常为 \(1/C_{\text{RS-upper} }\))
    • \(\epsilon\):PPO 裁剪范围(通常为 0.2)

策略角色与符号

  • \(\pi_{\text{rollout} }\)(行为策略 \(\mu\))
    • 在 Rollout/数据收集阶段创建快照行为策略快照,生成用于训练的轨迹
      • 理解:一般在某个 Rollout 批次上的训练期间冻结(异步时则按照一定的时机来冻结)
    • 常见不匹配来源 :
      • 策略不匹配:相同权重,不同实现(精度、后端),补充:还有算子,forward 策略等
      • 时间滞后:来自异步工作节点的陈旧 checkpoint
      • 回放缓冲:来自更早迭代的历史数据
      • Off-policy 算法:专家演示、辅助策略(DAPO)
      • 数据过滤:重加权或过滤后的数据
  • \(\pi_{\text{old} }\)(近端策略 \(\pi_{\text{prox} }\))
    • PPO 裁剪的参考策略(在同一个批次上的所有 PPO 更新轮次中冻结)
      • PPO 裁剪的锚点(控制策略更新大小)
      • 当与 \(\pi_{\text{rollout} }\) 分离时:实现批次大小不变性和陈旧数据的有效利用
      • 理解:LLM 中这里通常是在同一个 Rollout 批次上的所有 PPO 更新轮次中冻结(异步时则是没多个连续的更新步之间冻结)
      • 理解:LLM 中常见的 Proximal Policy 跟原始论文 Decoupled PPO(Hilton,2021) 中的做法不同(原始论文使用了 COM(Center of Mass)的概念来得到 Proximal 策略(一个历史策略加权平均的策略)
    • 创建时机
      • Decoupled mode :在训练轮次开始时通过 actor.compute_log_prob() 计算
      • Bypass mode :设置为等于 \(\pi_{\text{rollout} }\)(跳过单独计算)
  • \(\pi_{\theta}\)(当前策略 Current Policy)
    • 在训练期间被优化的策略,每一步梯度更新都会被修改

Operating Modes,运行模式

  • Decoupled mode (三个策略)
    • 在每个训练轮次开始时单独计算 \(\pi_{\text{old} }\)
    • 具有三个策略的完整 Decoupled PPO (数学上正确)
    • 实现批次大小不变性
    • 分别校正 Drift 1(rollout→old)和 Drift 2(old→current)
  • Bypass mode (两个策略)
    • 设置 \(\pi_{\text{old} } = \pi_{\text{rollout} }\)(跳过单独计算)
    • 使用 \(\pi_{\text{rollout} }\) 同时作为行为策略和近端策略(数学上正确)
    • 近端策略等于行为策略,因此两者之间不需要 IS 校正
    • 更快(跳过 actor.compute_log_prob() 调用)
    • 不实现批次大小不变性

两种分布漂移

  • Drift 1: \(\pi_{\text{rollout} } \to \pi_{\text{old} }\)(Off-Policy 差距)
    • 数据收集策略与训练参考策略之间的分布偏移
    • 范围从可忽略(相同 checkpoint,微小差异)到严重(回放缓冲、专家数据)
    • 校正方式 :重要性采样权重 \(w_t = \frac{\pi_{\text{old} }(a_t|s_t)}{\pi_{\text{rollout} }(a_t|s_t)}\)
    • 当使用 Bypass mode 时可忽略
  • Drift 2: \(\pi_{\text{old} } \to \pi_{\theta}\)(策略更新漂移)
    • 训练期间策略参数更新带来的漂移
    • 随着 \(\pi_\theta\) 通过梯度下降更新而自然发生
    • 校正方式 :对比值 \(r_t(\theta) = \frac{\pi_\theta(a_t|s_t)}{\pi_{\text{old} }(a_t|s_t)}\) 进行 PPO 裁剪
    • 适用于 on-policy 和 off-policy 训练

算法组件与组合

  • verl 中的 Rollout 校正框架由 正交组件 构建而成,可以灵活组合:
    • 1)运行模式 :如何计算 \(\pi_{\text{old} }\)( Decoupled vs Bypass )
    • 2)损失函数 :PPO(带裁剪)vs 纯 IS(仅策略梯度)
    • 3)IS/RS 聚合级别 :Token、序列或几何 (Geometric)

运行模式: Decoupled vs Bypass

  • 运行模式决定近端策略 \(\pi_{\text{old} }\) 的计算方式
Decoupled mode (三个策略)
  • 参数:bypass_mode = false
  • 策略设置:
    • \(\pi_{\text{rollout} }\):行为策略(数据收集)
    • \(\pi_{\text{old} }\):近端策略(在训练轮次开始时通过 actor.compute_log_prob() 计算)
    • \(\pi_{\theta}\):当前策略(被更新)
  • IS 比值: \(\rho_t = \frac{\pi_{\text{old} }(a_t|s_t)}{\pi_{\text{rollout} }(a_t|s_t)}\)(校正 Drift 1:rollout→old)
  • PPO 比值: \(r_t(\theta) = \frac{\pi_{\theta}(a_t|s_t)}{\pi_{\text{old} }(a_t|s_t)}\)(校正 Drift 2:old→current)
  • 性质:
    • ✅ 实现批次大小不变性
    • ✅ 分别校正两种分布漂移
    • ✅ 高效利用陈旧数据
    • ❌ 需要额外的前向传播(actor.compute_log_prob()
Bypass mode (两个策略)
  • 参数: bypass_mode = true
  • 策略设置:
    • \(\pi_{\text{rollout} }\):行为策略(数据收集)
    • \(\pi_{\text{old} } = \pi_{\text{rollout} }\):近端策略等于行为策略
    • \(\pi_{\theta}\):当前策略(被更新)
  • 比值:
    • 使用 PPO-clip 损失loss_type = "ppo_clip",默认):
      • PPO 比值 \(r_t(\theta) = \frac{\pi_{\theta}(a_t|s_t)}{\pi_{\text{rollout} }(a_t|s_t)}\) 针对 Rollout 策略进行裁剪(IS 由比值处理)
    • 使用 REINFORCE 损失loss_type = "reinforce"):
      • IS 比值 \(\rho_t = \frac{\pi_{\theta}(a_t|s_t)}{\pi_{\text{rollout} }(a_t|s_t)}\) 在损失函数中即时计算
  • 性质:
    • ✅ 跳过 actor.compute_log_prob() 调用(更快)
    • ✅ 通过 IS/RS 处理 off-policy 校正(当使用带 IS/RS 的策略梯度时)
    • ✅ 使用两个策略而非三个( \(\pi_{\text{rollout} }\) = \(\pi_{\text{old} }\) )
    • ⚠️ 不像 Decoupled mode 那样将近端策略与行为策略分离

损失函数:PPO vs 策略梯度

PPO 损失(带裁剪)
  • 参数: loss_type = "ppo_clip"( Bypass mode 默认)
  • 损失函数:
    $$
    L_{\text{PPO} }(\theta) = -\mathbb{E}_t \left[ w_t \cdot \min\left( r_t(\theta) A_t, \text{clip}(r_t(\theta), 1-\epsilon, 1+\epsilon) A_t \right) \right]
    $$
    • \(w_t\):IS 权重(取决于聚合级别,见第 3.3 节)。在 Decoupled mode 下,\(w_t = \frac{\pi_{\text{old} } }{\pi_{\text{rollout} } }\),且 \(\pi_{\text{old} }\) 被冻结,因此 \(w_t\) 是常数(不需要 stopgrad)。在 Bypass mode 下使用 PPO 损失时,通常不单独计算 IS 权重
    • \(r_t(\theta) = \frac{\pi_{\theta}(a_t|s_t)}{\pi_{\text{old} }(a_t|s_t)}\):PPO 比值
    • \(\epsilon\):裁剪范围(通常为 0.2)
  • 性质:
    • 通过裁剪实现信任域控制
    • 限制策略更新幅度
    • RL 训练中的标准做法
策略梯度损失(带 IS/RS 校正)
  • 参数:loss_type = "reinforce"(需要 bypass_mode = true

  • 损失函数(以序列级 IS 为例):
    $$
    L_{\text{PG} }(\theta) = -\mathbb{E}_{(s,a) \sim \pi_{\text{rollout} } } \left[ \text{stopgrad}(w_{\text{seq} }(\theta)) \cdot \sum_{t \in T} \log \pi_{\theta}(a_t|s_t) \cdot A_t \right]
    $$

    • \(w_{\text{seq} }(\theta)\):样本权重(IS 或 RS,详见 §3.3-3.4)
    • 对于 IS:\(w_{\text{seq} }(\theta) = \min\left( \prod_{t \in T} \frac{\pi_{\theta}(a_t|s_t)}{\pi_{\text{rollout} }(a_t|s_t)}, C_{\text{IS} } \right)\)
    • 对于 RS:\(w_{\text{seq} }(\theta) \in \{0, 1\}\)(二元拒绝掩码)
    • stopgrad 操作符 :权重 \(w_{\text{seq} }(\theta)\) 使用 \(\pi_\theta\) 计算,但在计算 \(\nabla_\theta L\) 时被当作 常数系数
      • 这对于重要性采样的正确性是必需的(理论依据见下文)
  • 有效梯度:
    $$
    \nabla_\theta L_{\text{PG} } = -\mathbb{E}_{(s,a) \sim \pi_{\text{rollout} } } \left[ \text{stopgrad}(w_{\text{seq} }(\theta)) \cdot \sum_{t \in T} \nabla_\theta \log \pi_{\theta}(a_t|s_t) \cdot A_t \right]
    $$

  • stopgrad 的理论依据:

    • stopgrad 操作符是 数学上必需的 ,由重要性采样理论要求,而非实现细节,原因如下:

      • 基本原则 :重要性采样是一种 改变测度 的技术(将样本从一个分布重新加权以估计另一个分布下的期望),而不是优化重加权函数本身

      • 形式推导

        • 1)原始目标 :作者要优化 \(J(\theta) = \mathbb{E}_{\tau \sim \pi_\theta}[\sum_t A_t]\)

        • 2)Off-policy 设置 :作者只有来自 \(\pi_{\text{rollout} }\) 的样本,因此使用重要性采样:
          $$
          J(\theta) = \mathbb{E}_{\tau \sim \pi_{\text{rollout} } } \left[ \underbrace{\frac{P_{\pi_\theta}(\tau)}{P_{\pi_{\text{rollout} } }(\tau)} }_{w(\tau;\theta)} \sum_t A_t \right]
          $$

        • 3)计算策略梯度 :正确的梯度使用 重要性采样之前的策略梯度定理
          $$
          \begin{aligned}
          \nabla_\theta J(\theta) &= \nabla_\theta \mathbb{E}_{\tau \sim \pi_\theta}\left[\sum_t A_t\right] \\
          &= \mathbb{E}_{\tau \sim \pi_\theta} \left[\sum_t A_t \nabla_\theta \log \pi_\theta(a_t|s_t) \right] \quad \text{(策略梯度定理)} \\
          &= \mathbb{E}_{\tau \sim \pi_{\text{rollout} } } \left[ w(\tau;\theta) \sum_t A_t \nabla_\theta \log \pi_\theta(a_t|s_t) \right] \quad \text{(测度变换)}
          \end{aligned}
          $$

          • 在最后一行中,\(w(\tau;\theta)\) 作为 乘法系数 来自测度变换,而不是被微分的东西
        • 4)没有 stopgrad 会出什么问题 :如果作者在损失中朴素地计算 \(\nabla_\theta \left[w(\theta) \log \pi_\theta \right]\),会得到:
          $$
          \nabla_\theta \left[w(\theta) \log \pi_\theta \right] = \underbrace{\log \pi_\theta \cdot \nabla_\theta w(\theta)}_{\text{错误:偏置项} } + \underbrace{w(\theta) \cdot \nabla_\theta \log \pi_\theta}_{\text{正确:IS 加权梯度} }
          $$

          • 第一项 \(\log \pi_\theta \cdot \nabla_\theta w(\theta)\) 是计算技巧(使用损失乘以 log 概率)的产物,不是真正策略梯度的一部分
          • 它会使梯度估计产生偏差,并优化与 \(J(\theta)\) 不同的目标
        • 5)实现要求 :在 PyTorch 中,为了只计算第二项,作者必须使用:

          1
          loss = -advantages * log_prob * rollout_is_weights.detach()  # 对权重 stopgrad
          • 如果不使用 .detach(),autograd 会同时计算两项,给出错误的梯度
  • 直观理解 :IS 权重 \(w(\theta)\) 告诉我们 “在估计 \(\pi_\theta\) 下的梯度时,这个样本应该被多大程度地信任”

    • 更新 \(\theta\) 以最大化重新加权后的目标,但并不更新 \(\theta\) 以最大化权重本身,因为那将是循环论证(优化校正因子而不是实际目标)
  • 性质:

    • 算法 :带 IS/RS 校正的 Off-policy 策略梯度
    • 损失类型( Bypass mode 下的 loss_type 配置选项):
      • "ppo_clip"(默认):PPO 裁剪目标
        • \(L = -\mathbb{E}[\min(r \cdot A, \text{clip}(r) \cdot A)]\),其中 \(r = \pi_\theta / \pi_{\text{rollout} }\)
        • 注意:应用 IS 权重(PPO 比值已经处理了它,否则会重复计数)
      • "reinforce":带显式 IS 权重的纯策略梯度,无 PPO 裁剪
        • \(L = -\mathbb{E}[w \cdot \log \pi_\theta(a|s) \cdot A]\),其中 \(w = \pi_\theta / \pi_{\text{rollout} }\)
    • 始终使用 Bypass mode :直接比较 \(\pi_\theta\) 和 \(\pi_{\text{rollout} }\)
    • 快速 :单次前向传播
  • 实现: core_algos.py 中的 compute_policy_loss_bypass_mode()compute_policy_loss_reinforce()

IS/RS 聚合级别

  • 聚合级别决定每 Token 概率比值如何组合成 IS 权重和/或拒绝掩码
  • 这个选择 独立于运行模式 ,可以在 Decoupled mode 或 Bypass mode 下使用任何聚合级别
Token 级聚合
  • IS 权重:
    $$w_t = \min(\rho_t, C_{\text{IS} })$$

    • 其中
      • Decoupled:
        $$\rho_t = \frac{\pi_{\text{old} }(a_t|s_t)}{\pi_{\text{rollout} }(a_t|s_t)}$$
      • 或 Bypass /纯 IS
        $$\rho_t = \frac{\pi_{\theta}(a_t|s_t)}{\pi_{\text{rollout} }(a_t|s_t)}$$
  • 参数配置:

    1
    2
    rollout_is = "token"  # IS 权重
    rollout_rs = "token_k1" # 可选:拒绝采样(比值边界)
  • 性质:

    • 每 Token 独立截断
    • 方差低于序列级(每 Token 的比值乘积被分别限定)
    • 偏差-方差权衡
      • Token 级校正偏差大,方差小:
        • Token 级矫正的 偏差 为 \(O(T^2 \Delta_{\max})\)
        • 其中 \(T\) 是序列长度,\(\Delta_{\max}\) 是每 Token 最大策略散度
        • 当 Rollout 策略与训练策略显著偏离时,这种偏差会变得显著
      • 序列级校正偏差为0,方差大
        • 数学上来看是无偏的,但方差更高
    • 典型阈值:1.5 - 5.0
    • 可选的批归一化 批归一化-batch-normalization:在所有 Token 权重上归一化,确保 \(\mathbb{E}[\tilde{w}_t] = 1\)(降低方差)
    • 使用时机 :当 Rollout 策略保持在训练策略的信任域内时,Token 级效果良好
      • 当不匹配严重时,偏差变得不可接受,此时应使用序列级校正
  • 损失函数(REINFORCE + Token IS):
    $$
    L_{\text{REINFORCE+TIS} }(\theta) = -\mathbb{E}_t \left[ \text{stopgrad}(w_t) \cdot \log \pi_\theta(a_t|s_t) \cdot A_t \right]
    $$

    • \(w_t = \min(\rho_t, C_{\text{IS} })\) 是截断后的 Token 级 IS 权重
    • stopgrad 操作符确保在计算 \(\nabla_\theta L\) 时,权重被视为常数
    • 此公式也可以通过将 REINFORCE 梯度替换为裁剪的替代目标来与 PPO 裁剪结合
  • 实现:

    • IS 权重:rollout_corr_helper.py 中的 compute_rollout_correction_weights()
    • 损失:core_algos.py 中的 compute_policy_loss()
序列级聚合
  • IS 权重:
    $$ w_{\text{seq} } = \min\left( \prod_{t \in T} \rho_t, C_{\text{IS} } \right) = \min\left( \exp\left(\sum_{t \in T} \log \rho_t\right), C_{\text{IS} } \right)$$

    • 广播到所有 Token
  • 参数配置:

    1
    2
    rollout_is = "sequence"  # IS 权重
    rollout_rs = "seq_sum_k1" # 可选:拒绝采样
  • 性质:

    • 乘法聚合
    • 对异常值比 Token 级更敏感
    • 典型阈值:2.0 - 10.0
    • 可选的批归一化 批归一化-batch-normalization:在序列均值上归一化(每个序列一个权重)
  • 术语说明:

    • Seq-TIS(序列级截断 IS) :将序列比值 \(\rho(\tau)\) 裁剪为 \(\min(\rho(\tau), C)\)
      • 通过从所有样本中提取信号来最大化信息效率
      • 适用于数据干净、不匹配程度适中的情况
    • Seq-MIS(序列级掩码 IS) :拒绝(掩码)\(\rho(\tau) > C\) 的序列,而不是裁剪
      • 相当于一个硬信任域过滤器
      • 适用于严重不匹配,或分布尾部 “有毒(toxic)”(包含垃圾/对抗样本而非信号)的情况
  • 损失函数(REINFORCE + 序列 IS):
    $$
    L_{\text{REINFORCE+SeqIS} }(\theta) = -\mathbb{E}_t \left[ \text{stopgrad}(w_{\text{seq} }) \cdot \log \pi_\theta(a_t|s_t) \cdot A_t \right]
    $$

    • \(w_{\text{seq} }\) 广播到序列中的所有 Token
    • stopgrad 操作符确保正确的 IS 梯度计算
    • 此公式也可以通过将 REINFORCE 梯度替换为裁剪的替代目标来与 PPO 裁剪结合
几何均值聚合 (Geo-RS)
  • 几何均值比值:
    $$ \rho_{\text{geo} } = \exp\left( \frac{1}{|T|} \sum_{t \in T} \log \rho_t \right) = \left(\prod_{t \in T} \rho_t\right)^{1/|T|} $$

    • 广播到所有 Token
  • 参数配置:

    1
    2
    rollout_is = null  # 无 IS 权重,纯拒绝
    rollout_rs = "seq_mean_k1" # 几何均值拒绝采样(比值边界)
  • 性质:

    • 长度不变性(按序列长度归一化)
    • 理想比值 = 1.0(策略匹配)
    • 典型边界:"0.999_1.001"(约 ±0.1%)
    • 仅用于拒绝采样,不用于 IS 加权
      • 理解:因为经过几何平均处理过的 IS 在数学上是不准确的(PS:GSPO 本身在数学上是不准确的)
  • 长度陷阱问题:

    • 标准 IS 估计器存在系统性的 长度偏差 ,会惩罚长序列,重要性比值 \(\rho(y)\) 是乘性的:
      $$
      \rho(y) = \prod_{t=1}^T \frac{\pi(y_t|y_{ < t})}{\mu(y_t|y_{ < t})}
      $$
      • 因为样本是从行为策略 \(\mu\) 中采样的,于是与 KL 散度类似,这个比值往往是小于等于 1 的,对概率比值的连乘(或者实现时更多是对概率比值对数的连加,然后取指数)会导致这个差异逐步累计,最终造成长度越长,得到的结果越极端(一般是接近于 0)
    • 假设新策略 \(\pi\) 与 \(\mu\) 略有不同,平均每 Token 比值约为 1.1(注:这个例子举的不好,每 Token 比值一般是小于 1 的):
      • 短序列(10 Token): \(\rho \approx 1.1^{10} \approx 2.6\) → 在阈值内,保留
      • 长序列(100 Token): \(\rho \approx 1.1^{100} \approx 13,780\) → 超过阈值,拒绝
    • 这会造成 上下文崩塌 (Context Collapse) :模型偏向于学习短而浅的答案,拒绝长推理链(即使每步质量相同)
      • 对于推理模型(CoT)和 Agent,这实际上惩罚了“思考太久”
  • Geo-RS 解决方案:

    • 几何级拒绝按序列长度归一化,将广延量(总概率乘积)转换为强度量(平均每 Token 漂移):
      $$
      \rho_{\text{geo} }(y) = \rho(y)^{1/T}
      $$
    • 现在两个序列具有相同的“信任分数”:
      • 短(10 Token): \((1.1^{10})^{1/10} = 1.1\)
      • 长(100 Token): \((1.1^{100})^{1/100} = 1.1\)
  • 为什么要用紧阈值?

    • 对于 100 个 Token,每个 Token 的对数比值为 0.01:
      • 算术乘积比值:\(e^{100 \times 0.01} \approx 2.7\)
      • 几何比值:\(e^{0.01} \approx 1.010\)
    • 比值边界 "0.999_1.001" 会拒绝平均每 Token 对数偏差超过约 0.1% 的序列
  • 损失函数(REINFORCE + 几何 RS):
    $$
    L_{\text{GeoRS} }(\theta) = -\mathbb{E}_{(s,a) \mid \text{seq} \in \mathcal{A}_{\text{geo} } } \left[ \sum_{t \in T} \log \pi_\theta(a_t|s_t) \cdot A_t \right]
    $$

    • 其中 \(\mathcal{A}_{\text{geo} } = \{ \text{seq} : C_{\text{RS-lower} } \leq \rho_{\text{geo} } \leq C_{\text{RS-upper} } \}\) 是接受集合(拒绝掩码)
    • 不使用 IS 权重,因此不需要 stopgrad
    • 此公式也可以通过将 REINFORCE 梯度替换为裁剪的替代目标来与 PPO 裁剪结合
  • 组合估计器(Geo-RS-Token-TIS):

    • 为获得最佳效果,将 几何过滤器(长度不变的有效性检查)与 Token 级 IS 权重(更低方差)结合:
      $$
      \hat{g}_{\text{geo-rs-token-tis} }(y) = \underbrace{\mathbb{I}\left( C_{\text{low} } \le \rho(y)^{1/T} \le C_{\text{high} } \right)}_{\text{几何过滤器} } \cdot \prod_t \min(\rho_t, C) \cdot f(y)
      $$
    • 参数配置:
      • rollout_rs="seq_mean_k1" and rollout_is="token"
K2 散度聚合
  • 每 Token 统计量:
    $$
    K2_t = \frac{1}{2} \left(\log \rho_t\right)^2
    $$

    • \(\rho_t = \frac{\pi_{\text{old} }(a_t|s_t)}{\pi_{\text{rollout} }(a_t|s_t)}\),实现中将 \(\log \rho_t\) 裁剪到 \([-20, 20]\) 以保证数值安全
  • 序列聚合(共享相同的每 Token \(K2_t\)):

    • seq_sum_k2:\(K2_{\text{sum} } = \sum_{t \in T} K2_t\)
    • seq_mean_k2:\(K2_{\text{mean} } = \frac{1}{|T|} \sum_{t \in T} K2_t\)
    • seq_max_k2:\(K2_{\text{max} } = \max_{t \in T} K2_t\)
  • 参数配置:

    1
    2
    3
    rollout_is = null            # 可选:与 Token IS 权重搭配以降低方差
    rollout_rs = "token_k2" # 或 "seq_sum_k2", "seq_mean_k2", "seq_max_k2"
    rollout_rs_threshold = 2.0 # 仅正上界
  • 性质:

    • 在 \(\log \rho_t\) 上的对称二次惩罚;当策略匹配时为零
    • 在小策略漂移下近似 \(\tfrac{1}{2}\operatorname{Var}[\log \rho]\),因此是匹配度的平滑检测器
    • 仅有上阈值:典型范围对于 token_k2 为 1.5-3.0,对于 seq_mean_k2 为 2.0-2.5,对于 seq_sum_k2 为 2.5-4.0
    • seq_max_k2 即使在序列其余部分干净时也能隔离单 Token 尖峰
    • 可与 Token 级 IS 权重(rollout_is="token")共存,在保留有用样本的同时裁剪方差
  • 组合估计器(K2-RS-Token-TIS):

    • 对于组合过滤和加权,令 \(K2_{\text{agg} }\) 表示选定的聚合(token、sum、mean 或 max):
      $$
      \hat{g}_{\text{k2-rs-token-tis} }(y) = \underbrace{\mathbb{I}\left( K2_{\text{agg} }(y) \le C_{\text{k2} } \right)}_{\text{K2 过滤器} } \cdot \prod_t \min(\rho_t, C) \cdot f(y)
      $$
    • 参数配置:rollout_rs="seq_mean_k2"(或其他 k2 模式)配合 rollout_is="token" 实现
K3 散度聚合
  • 序列级 K3 散度:
    $$
    K3_{\text{seq} } = \frac{1}{|T|} \sum_{t \in T} \left( \rho_t - \log \rho_t - 1 \right)
    $$

    • \(\rho_t = \frac{\pi_{\text{old} }(a_t|s_t)}{\pi_{\text{rollout} }(a_t|s_t)}\) 是每 Token 比值
  • K3 等于反向 KL: 在期望上,\(K3 = \text{KL}(\pi_{\text{rollout} } | \pi_{\text{old} })\),推导如下:

    • \(\mathbb{E}_{\pi_\text{rollout} }[\rho] = 1\)
    • \(\mathbb{E}_{\pi_\text{rollout} }[\log \rho] = -\text{KL}(\pi_{\text{rollout} } | \pi_{\text{old} })\)
    • 因此:\(K3 = 1 - (-\text{KL}) - 1 = \text{KL}(\pi_{\text{rollout} } | \pi_{\text{old} })\)
  • 参数配置:

    1
    2
    rollout_is = null          # 无 IS 权重,纯拒绝
    rollout_rs = "seq_mean_k3" # K3 拒绝采样
  • 性质:

    • K3 散度每 Token 总是 >= 0(当 \(\rho = 1\) 时等于 0)
    • 比几何比值检查更稳定,因为每个 Token 项都是非负的
    • 仅有上阈值(因为 K3 >= 0,无需下阈值)
    • 典型阈值:0.001 - 0.01
  • 为什么用 K3 而不是几何比值?

    • 几何比值使用平均对数比值;微小的数值偏差可能使符号翻转
    • \(K3 = \mathbb{E}[\rho - log \rho - 1]\) 每 Token 非负,提供更平滑的检测器
    • 两者估计相同的量:KL( \(\pi_{\text{rollout} }\) || \(\pi_{\text{old} }\) )
    • 对于小散度,K3 ≈ 0.5 × Var(log_ratio)
  • 组合估计器(K3-RS-Token-TIS):

    • 为获得最佳效果,将 K3 过滤器与 Token 级 IS 权重结合:
      $$
      \hat{g}_{\text{k3-rs-token-tis} }(y) = \underbrace{\mathbb{I}\left( K3_{\text{seq} } \le C_{\text{k3} } \right)}_{\text{K3 过滤器} } \cdot \prod_t \min(\rho_t, C) \cdot f(y)
      $$
    • 通过组合 rollout_rs="seq_mean_k3"rollout_is="token" 实现

Batch Normalization, 批归一化

  • Batch Normalization 是一种可选的方差缩减技术,将 IS 权重归一化到每个批次内均值为 1.0

  • 参数配置

    1
    rollout_is_batch_normalize = True  # 默认:False
  • 归一化公式(聚合感知):

    • 对于 Token 级 IS
      $$
      \tilde{w}_t = \frac{w_t}{\frac{1}{\sum_{i,t} m_{i,t} } \sum_{i,t} w_{i,t} \cdot m_{i,t} }
      $$
      • \(w_{i,t}\) 是截断后的 Token IS 权重
      • \(m_{i,t}\) 是 Response 掩码,归一化在 所有 Token 上进行
    • 对于 序列级 IS
      $$
      \tilde{w}_i = \frac{w_i}{\frac{1}{B}\sum_{j=1}^B \bar{w}_j}
      $$
      • \(\bar{w}_j = \frac{1}{T_j}\sum_{t=1}^{T_j} w_{j,t} \cdot m_{j,t}\) 是每个序列的均值(一个序列中的所有 Token 权重相同),归一化在 序列 上进行
  • 性质:

    • 在截断 之后 应用,以保留截断的语义
    • 确保每个批次内 \(\mathbb{E}[\tilde{w}] = 1\)
    • 聚合感知
      • Token 级在 Token 上归一化
      • 序列级在序列上归一化
    • 使用 masked_mean 以尊重填充 Token
    • 通过消除随机的批次级尺度波动来降低梯度量级方差
  • 指标:

    • rollout_is_batch_norm_factor:应用的归一化因子(归一化前的批次均值)

拒绝采样 (RS)

  • 拒绝采样可以添加到 任何 运行模式和聚合级别的组合中,它修改 response_mask 以排除离群的 Token/序列

  • 参数配置示例:

    1
    2
    3
    4
    5
    6
    7
    8
    rollout_rs = "token_k1"    # Token 级比值边界
    rollout_rs_threshold = "0.6_1.6"

    rollout_rs = "seq_sum_k1" # 序列对数比值之和
    rollout_rs_threshold = "0.5_2.0"

    rollout_rs = "seq_mean_k3" # 序列 K3 散度均值
    rollout_rs_threshold = 0.01
  • 接受集合:

    • Token 级 :\(\mathcal{A}_{\text{token} } = \{ t : C_{\text{RS-lower} } \leq \rho_t \leq C_{\text{RS-upper} } \}\)
    • 序列级 :\(\mathcal{A}_{\text{seq} } = \{ \text{seq} : C_{\text{RS-lower} } \leq \prod_{t \in T} \rho_t \leq C_{\text{RS-upper} } \}\)
    • 几何级 :\(\mathcal{A}_{\text{geo} } = \{ \text{seq} : C_{\text{RS-lower} } \leq \rho_{\text{geo} } \leq C_{\text{RS-upper} } \}\)
  • 性质:

    • 与 IS 加权分离(可以只使用 RS 而不使用 IS)
    • 降低有效样本量
    • 过滤极端离群值
  • 代码实现: rollout_corr_helper.py 中的 compute_rollout_rejection_mask()

组合矩阵

  • 关键洞察: 估计器(如何计算 IS/RS)和运行模式( Decoupled PPO vs Bypass PG)是 正交的
  • 任何估计器都可以与任何运行模式组合
估计器 × 运行模式
  • 具体配置实例:
    估计器 配置 兼容模式
    Token-TIS rollout_is="token" Decoupled PPO , Bypass PG
    Seq-TIS rollout_is="sequence" Decoupled PPO , Bypass PG
    Seq-MIS rollout_is="sequence" + rollout_rs="seq_sum_k1" Decoupled PPO , Bypass PG
    Geo-RS rollout_rs="seq_mean_k1"(几何均值) Decoupled PPO , Bypass PG
    Geo-RS-Token-TIS rollout_is="token" + rollout_rs="seq_mean_k1" Decoupled PPO , Bypass PG
    K3-RS rollout_rs="seq_mean_k3" Decoupled PPO , Bypass PG
    K3-RS-Token-TIS rollout_is="token" + rollout_rs="seq_mean_k3" Decoupled PPO , Bypass PG
  • 注意: 在 Bypass mode 下,loss_type 控制损失函数
    • 使用 “ppo_clip”(默认)或 “reinforce”
可用的预设方法
  • 预设方法展示:
    预设方法 估计器 模式 性质
    Decoupled PPO 模式(3 个策略: \(\pi_{\text{rollout} }\) , \(\pi_{\text{old} }\) , \(\pi_{\theta }\))
    decoupled_token_is() Token-TIS Decoupled PPO 每 Token IS 权重
    decoupled_seq_is() Seq-TIS Decoupled PPO 序列级 IS 权重
    decoupled_seq_is_rs() Seq-MIS Decoupled PPO 序列 IS + 序列 RS
    decoupled_geo_rs() Geo-RS Decoupled PPO 几何 RS
    decoupled_geo_rs_token_tis() Geo-RS-Token-TIS Decoupled PPO 几何过滤器 + Token IS
    K3 KL 估计器(对小 KL 值更稳定)
    decoupled_k3_rs() K3-RS Decoupled PPO K3 拒绝,无 IS 权重
    decoupled_k3_rs_token_tis() K3-RS-Token-TIS Decoupled PPO K3 过滤器 + Token 裁剪权重
    Bypass mode (PPO-clip)(比值处理 IS,RS 掩码离群值)
    bypass_ppo_clip() - Bypass (PPO-clip) 仅 PPO-clip
    bypass_ppo_clip_geo_rs() Geo-RS Bypass (PPO-clip) PPO-clip + Geo-RS(比值)
    bypass_ppo_clip_k3_rs() K3-RS Bypass (PPO-clip) PPO-clip + K3-RS
    Bypass mode (REINFORCE)(显式 IS 权重,无 PPO 裁剪)
    bypass_pg_is() Seq-TIS Bypass (REINFORCE) REINFORCE + Seq IS
    bypass_pg_geo_rs() Geo-RS Bypass (REINFORCE) REINFORCE + Geo-RS(比值)
    bypass_pg_geo_rs_token_tis() Geo-RS-Token-TIS Bypass (REINFORCE) REINFORCE + Geo 过滤器 + Token IS
    其他
    disabled() - - 仅指标
  • 注意:Bypass mode 设置 \(\pi_{\text{old} }\) = \(\pi_{\text{rollout} }\) ,并使用 loss_type 选择损失函数
额外支持的组合(手动配置)
  • 这些组合 完全支持 ,但需要手动配置:

    • 1. Token IS + Token RS

      1
      2
      3
      4
      5
      6
      config = RolloutCorrectionConfig(
      rollout_is="token",
      rollout_is_threshold=2.0,
      rollout_rs="token_k1",
      rollout_rs_threshold="0.5_2.0",
      )
      • 性质: Token 级 IS 权重 + Token 级 RS 掩码
    • 2. 纯 Token RS

      1
      2
      3
      4
      5
      config = RolloutCorrectionConfig(
      rollout_is=None,
      rollout_rs="token_k1",
      rollout_rs_threshold="0.5_2.0",
      )
      • 性质: 仅 Token 级 RS 掩码,无 IS 权重
    • 3. 纯序列 RS

      1
      2
      3
      4
      5
      config = RolloutCorrectionConfig(
      rollout_is=None,
      rollout_rs="seq_sum_k1",
      rollout_rs_threshold="0.5_2.0",
      )
      • 性质: 仅序列级 RS 掩码,无 IS 权重
  • 关键性质:

    • 任何 IS 聚合级别(Token/序列)都可以在 Decoupled 或 Bypass mode 下使用
    • 拒绝采样可以添加到任何组合中
    • 几何聚合通常仅用于 RS(不用于 IS 加权)
    • 纯 RS(bypass_pg_rs)使用 Bypass + 几何 RS,搭配 loss_type="reinforce" 用于 REINFORCE(无 IS 权重)
    • 上表中的所有组合都是有效的,并且实现支持

常见实现错误

不正确的 LLM-RL 实现(无 Rollout 校正的 PPO)
  • 理论: 朴素的 LLM-RL 实现错误地应用 PPO,忽略实际 Rollout 策略 ,并假设 \(\pi_{\text{old} } = \pi_{\text{rollout} }\)
  • 注意: 这种错误的实现模式在 When Speed Kills Stability: Demystifying RL Collapse from the Training-Inference Mismatch 中被识别为 LLM-RL 系统训练不稳定的关键原因,从而推动了本 Rollout 校正框架的发展
  • 损失函数:
    $$
    L_{\text{PPO} }(\theta) = -\mathbb{E}_t \left[ \min\left( r_t(\theta) A_t, \text{clip}(r_t(\theta), 1-\epsilon, 1+\epsilon) A_t \right) \right]
    $$
    • \(r_t(\theta) = \frac{\pi_{\theta}(a_t|s_t)}{\pi_{\text{old} }(a_t|s_t)}\)(忽略 \(\pi_{\text{rollout} }\))
  • 为什么这是错误的:
    • 忽略 \(\pi_{\text{rollout} }\) :使用 \(\pi_{\text{old} }\) 作为行为策略,而不是实际的 \(\pi_{\text{rollout} }\)
    • 策略不匹配 :在 LLM-RL 中,Rollout 通常使用与训练不同的精度/后端/checkpoint,即使权重相同,也会导致 \(\pi_{\text{rollout} } \neq \pi_{\text{old} }\)
    • 不是 PPO 的错 :PPO 本身是正确的;问题在于错误的假设
  • 正确的替代方案:
    • 1)Decoupled mode :三个策略,带从 \(\pi_{\text{rollout} }\) 到 \(\pi_{\text{old} }\) 的 IS 校正
    • 2)Bypass mode :两个策略,使用 \(\pi_{\text{rollout} }\) 同时作为行为策略和近端策略
    • 3)Bypass + 策略梯度模式 :两个策略,带 IS/RS 校正,无 PPO 裁剪
  • 实现: core_algos.py 中的 compute_policy_loss()

Off-Policy 诊断指标

  • 这些指标量化 off-policy 漂移的严重程度
  • 符号说明: 指标使用 \(\rho_t = \frac{\pi_{\text{old} }(a_t|s_t)}{\pi_{\text{rollout} }(a_t|s_t)}\)
    • 在 Bypass mode 下,\(\pi_{\text{old} } = \pi_{\text{rollout} }\),因此指标使用 \(\rho_t = \frac{\pi_{\theta} }{\pi_{\text{rollout} } }\) 来测量 rollout→current 漂移

KL 散度

  • 直接 KL 估计器:
    $$
    \text{KL}(\pi_{\text{rollout} } | \pi_{\text{old} }) = \mathbb{E}_{t \sim \pi_{\text{rollout} } } \left[ \log \pi_{\text{rollout} }(a_t|s_t) - \log \pi_{\text{old} }(a_t|s_t) \right]
    $$
  • K3 KL 估计器(替代公式):
    $$
    \text{KL}_{\text{K3} } = \mathbb{E}_{t \sim \pi_{\text{rollout} } } \left[ \rho_t - \log \rho_t - 1 \right]
    $$
    • 其中 \(\rho_t = \frac{\pi_{\text{old} }(a_t|s_t)}{\pi_{\text{rollout} }(a_t|s_t)}\)

困惑度 (Perplexity)

  • 旧策略困惑度:
    $$
    \text{PPL}_{\text{old} } = \exp\left( -\frac{1}{|T|} \sum_{t \in T} \log \pi_{\text{old} }(a_t|s_t) \right)
    $$
  • Rollout 策略困惑度:
    $$
    \text{PPL}_{\text{rollout} } = \exp\left( -\frac{1}{|T|} \sum_{t \in T} \log \pi_{\text{rollout} }(a_t|s_t) \right)
    $$
  • PPL 比值(几何均值 IS 权重的倒数):
    $$
    \text{PPL}_{\text{ratio} } = \frac{\text{PPL}_{\text{old} } }{\text{PPL}_{\text{rollout} } } = \exp\left( -\frac{1}{|T|} \sum_{t \in T} \log \rho_t \right) = \left(\prod_{t \in T} \rho_t\right)^{-1/|T|}
    $$
  • 解释: 值 > 1 表示 \(\pi_{\text{old} }\) 对观察到的动作赋予的概率低于 \(\pi_{\text{rollout} }\)(存在分布偏移)

卡方散度

  • 衡量 IS 权重分布的二阶矩
  • Token 级:
    $$
    \chi^2_{\text{token} } = \mathbb{E}_{t \sim \pi_{\text{rollout} } } \left[ \rho_t^2 \right] - 1
    $$
  • 序列级:
    $$
    \chi^2_{\text{seq} } = \mathbb{E}_{\text{seq} \sim \pi_{\text{rollout} } } \left[ \left(\prod_{t \in T} \rho_t\right)^2 \right] - 1
    $$
  • 解释:
    • \(\chi^2 = 0\):策略相同
    • \(\chi^2 > 0\):值越高表示 off-policy 分布偏移越严重
  • 实现: rollout_corr_helper.py 中的 compute_offpolicy_metrics()

总结与决策指南

方法汇总表

  • 整体方法汇总:
    方法 理论 策略数 PPO 裁剪 IS 校正 正确性 速度
    Bypass mode( \(\pi_{\text{old} }\) = \(\pi_{\text{rollout} }\) ,loss_type 选择算法)
    loss_type="ppo_clip"(默认) PPO(比值 = \(\pi_{\theta }\)/ \(\pi_{\text{rollout} }\) ) 2(rollout, \(\theta\)) 仅 RS 掩码(比值处理 IS) ✅ 正确
    loss_type="reinforce" Off-policy REINFORCE 2(rollout, \(\theta\)) ✅(显式 IS 权重) ✅ 正确
    Bypass mode 预设(PPO-clip)
    bypass_ppo_clip 仅 PPO 2(rollout, \(\theta\)) - ✅ 正确
    bypass_ppo_clip_geo_rs PPO + Geo-RS 2(rollout, \(\theta\)) Geo-RS 掩码(比值) ✅ 正确
    Bypass mode 预设(REINFORCE)
    bypass_pg_is REINFORCE + Seq-TIS 2(rollout, \(\theta\)) ✅ Seq-TIS ✅ 正确
    bypass_pg_geo_rs REINFORCE + Geo-RS 2(rollout, \(\theta\)) 仅 Geo-RS(比值) ✅ 正确
    bypass_pg_geo_rs_token_tis REINFORCE + Geo RS + Token IS 2(rollout, \(\theta\)) ✅ Geo-RS-Token-TIS ✅ 正确
    Decoupled PPO 模式(IS 权重 = \(\pi_{\text{old} }\) / \(\pi_{\text{rollout} }\) )
    decoupled_token_is Decoupled PPO 3(rollout, old, \(\theta\)) ✅ Token-TIS ✅ 正确 标准
    decoupled_seq_is Decoupled PPO 3(rollout, old, \(\theta\)) ✅ Seq-TIS ✅ 正确 标准
    decoupled_seq_is_rs Decoupled PPO + RS 3(rollout, old, \(\theta\)) ✅ Seq-MIS ✅ 正确 标准
    decoupled_geo_rs Decoupled PPO + Geo-RS 3(rollout, old, \(\theta\)) Geo-RS 仅(比值) ✅ 正确 标准
    decoupled_geo_rs_token_tis Decoupled PPO + Geo RS + Token IS 3(rollout, old, \(\theta\)) ✅ Geo-RS-Token-TIS ✅ 正确 标准
    错误(供参考)
    朴素 LLM-RL 错误的 PPO 使用 2(old, \(\theta\)) ⚠️ 错误 标准
  • 注意:
    • Bypass mode 设置 \(\pi_{\text{old} }\) = \(\pi_{\text{rollout} }\) ,并使用 loss_type 选择损失函数:
      • "ppo_clip"(默认):PPO 裁剪比值(IS 由比值 = \(\pi_{\theta }\)/ \(\pi_{\text{rollout} }\) 处理,不显式使用 IS 权重以避免重复计数)
      • "reinforce":显式 IS 权重以 \(w \cdot \log \pi \cdot A\) 的形式应用
    • 两种损失类型都受益于拒绝采样(RS),它可以掩码掉分布外的样本

估计器层级

  • 这些估计器定义 IS 权重和拒绝掩码的计算方式
  • 它们与运行模式( Decoupled PPO vs Bypass 策略梯度)正交,并且可以与任一种组合
    估计器 配置 机制 最适合
    Token-TIS rollout_is="token" 裁剪每 Token 比值 偏差可接受时方差较低的 IS
    Seq-TIS rollout_is="sequence" 裁剪序列比值 \(\rho(\tau) \to \min(\rho(\tau), C)\) 数据干净、不匹配程度适中的情况;无偏
    Seq-MIS rollout_is="sequence" + rollout_rs="seq_sum_k1" 拒绝 \(\rho(\tau) > C\) 的序列 严重不匹配;过滤“有毒尾部”(垃圾数据)
    Geo-RS rollout_rs="seq_mean_k1" 基于几何均值比值 exp(E[log(r)]) 拒绝 长度不变的信任域
    Geo-RS-Token-TIS rollout_is="token" + rollout_rs="seq_mean_k1" 几何过滤器 + Token IS 权重 基于比值的长度归一化 + 低方差 IS
    K3-RS rollout_rs="seq_mean_k3" 基于 K3 KL 散度拒绝 小 KL 值;平滑检测器
    K3-RS-Token-TIS rollout_is="token" + rollout_rs="seq_mean_k3" K3 过滤器 + Token IS 权重 小 KL + 低方差 IS
  • 注意: 每个估计器都可以与以下任一种运行模式结合使用:
    • Decoupled PPObypass_mode=false):三个策略,带 PPO 裁剪
    • Bypass modebypass_mode=true):两个策略,可配置损失类型
      • loss_type="ppo_clip"(默认):PPO 裁剪目标(IS 通过比值,应用 RS 掩码)
      • loss_type="reinforce":带显式 IS 权重的 REINFORCE

基于场景的方法特点

  • 按 off-policy 严重程度选择估计器:
    • 可忽略(相同 checkpoint,微小差异):无需 IS 校正;使用 Bypass mode 以提高效率
    • 中等(异步工作节点,轻微陈旧):Token-TIS 提供每 Token IS 校正,方差较低
    • 严重(回放缓冲,陈旧数据):Seq-TIS 或 Seq-MIS 提供序列级 IS 校正;当高权重样本可能包含垃圾数据时使用 Seq-MIS
  • 按序列长度选择估计器:
    • 短序列(标准对话):Seq-TIS 最优
    • 长序列(CoT,Agent):K1-RS 或 K1-RS-Token-TIS 以避免长度陷阱
  • 选择运行模式:
    • 需要批次大小不变性 :使用 Decoupled mode (bypass_mode=false
    • 需要计算效率 :使用 Bypass mode (bypass_mode=true)以跳过 old_log_prob 计算
    • 无需 PPO 裁剪 :使用 Bypass mode 并设置 loss_type="reinforce"

Decoupled mode vs Bypass mode

  • Decoupled mode(单独计算 old_log_prob):
    • 实现完整的 Decoupled PPO ,使用三个策略(数学上正确)
    • 分别测量和校正 Drift 1(rollout→old)和 Drift 2(old→current)
    • 实现批次大小不变性和陈旧数据的高效利用
    • 支持准确的 off-policy 指标监控
  • Bypass mode(设置 \(\pi_{\text{old} } = \pi_{\text{rollout} }\)):
    • 使用 \(\pi_{\text{rollout} }\) 同时作为行为策略和近端策略(数学上正确)
    • 计算效率:跳过单独的 old_log_prob 计算
    • 不实现批次大小不变性(近端策略取决于数据收集方式)