技术摘要:PS-PPO:用于无评论员(Critic-Free)RLHF 的前缀采样 PPO
1. 问题陈述
基于人类反馈的强化学习(RLHF)是使大语言模型(LLM)与人类偏好对齐的标准方法。然而,使用 RLHF 训练 LLM 面临着显著的计算瓶ال瓶颈,特别是在无评论员(critic-free)方法(例如 GRPO、RLOO)中,这些方法避免了训练单独的价值网络。
在这些无评论员的方法中,单个标量奖励在生成完成后被分配给整个轨迹(completion)。随后,该奖励被均匀地广播到轨迹中的所有 token,以计算优势(advantages)。因此,策略更新需要针对每个 rollout 进行全轨迹反向传播。作者认为这是低效的,因为在许多推理任务(如逐步数学运算)中,中间的前缀往往已经包含了足以确定最终结果的充分信息。作者的实证分析(图 1)表明,轨迹的成功率往往在完成结束前就已趋于稳定,这意味着后缀部分的 token 携带了冗余的学习信号。因此,全轨迹更新会在这些冗余的后缀上浪费大量的计算资源和 GPU 显存。
2. 方法论:PS-PPO
作者提出了前缀采样近端策略优化(Prefix-Sampling PPO, PS-PPO),这是一种计算高效的无评论员方法,通过截断反向传播过程来利用时间上的冗余性,同时保持无偏的梯度估计量。
核心机制
PS-PPO 并非通过整个完成序列 o=[o1,…,oT] 进行反向传播,而是为每个轨迹采样一个截断时间步 H。策略更新仅在当前前缀 o≤H=[o1,…,oH] 上进行计算。
为了确保这种截断不会引入偏差,PS-PPO 采用了重加权方案:
- 截断分布: 定义了一个以 prompt 为条件的分布 ξ1:T,其中 ξt=Pr(H≥t∣x) 表示第 t 个 token 被保留的概率。该分布是非增的(ξ1≥ξ2≥⋯≥ξT)。
- 无偏估计量: 截断轨迹的梯度通过逆包含概率 1/ξt 进行重加权。所得估计量 G^(θ) 为:
G^(θ):=K1k=1∑Kt=1∑H(k)ξt1gt(k)(θ)
其中 gt(k) 是逐时间步的梯度。该估计量在 H 的采样上的期望值恢复了全轨迹更新。
优化截断分布
论文将 ξ1:T 的选择表述为一个凸优化问题。目标是在满足计算预算 B(期望的反向传播 token 数)的条件下,最小化截断梯度估计量的方差。
- 方差代理: 作者推导出了重要性采样引起的方差的一个易于处理的代理,将其近似为 ∑wt(x)(1/ξt−1),其中 wt(x) 代表第 t 个时间步的更新重要性。
- 重要性代理: 由于计算精确的得分范数 ∥∇θlogπθ∥2 需要进行完整反向传播(这会违背初衷),作者使用了一个基于输出头激活和奖励不确定性的仅前向代理。
- 奖励不确定性 (ut(x)): 通过给定前缀状态下的奖励方差来估计,该方差通过成功与失败的 rollout 之间下一 token 分布的距离来近似。
- 梯度范数代理 (γˉt): 从输出头梯度范数中导出。
- 优化: 该问题在满足单调性和预算约束的情况下,最小化 ∑ξtγˉt(x,t)ut(x)。解法使用**相邻违反者消除(Pool Adjacent Violators, PAV)**算法来强制执行单调性约束。
算法流程
- 为每个 prompt 生成 K 个 completion。
- 计算终止奖励并广播优势。
- 从 rollouts 中估计逐时间步的不确定性和梯度代理。
- 求解凸优化问题以获得最优的单调截断概率 ξ1:T∗。
- 根据 ξ∗ 定义的分布为每个轨迹采样一个截断点 H(k)。
6.仅对 t≤H(k) 的 token 进行反向传播,并使用 1/ξt∗ 对损失进行重加权。
3. 主要贡献
- PS-PPO 框架: 一种新型的无评论员 RLHF 方法,它引入了以 prompt 为条件的随机截断反向传播过程,显著降低了计算和内存成本。
- 无偏截断估计器: 推导出了一个通过包含概率重加权的估计器,该估计器在期望上恢复了全序列广播更新,且不需要额外的 rollout 或辅助价值模型。
- 优化的截断策略: 一种确定“在哪里”进行截断的原则性方法,通过求解一个平衡方差减少与计算预算的凸优化问题,并利用前向代理来获取梯度重要性和奖励不确定性。
- 实证验证: 证明了 PS-PPO 在数学推理基准测试(MATH500, AMC 2023, Minerva Math, AIME 2024/2025)上达到了与强力无评论员基线(GRPO, DAPO, RLOO)相当的准确率,同时大幅减少了训练时间和峰值 GPU 显存。
4. 实验结果
作者使用 Llama-3.1-8B-Instruct 和 Qwen2.5-Math-7B 在数学推理基准测试上评估了 PS-PPO。
- 性能: PS-PPO(优化版)在所有基准测试中实现的 Pass@1 准确率与 DAPO 和 S-GRPO 等强基线相当或略好。例如,在 MATH500 上使用 Llama-3.1-8B 时,PS-PPO 达到了 47.6%,而 DAPO 为 46.8%。
- 效率:
- 训练时间: 与基线相比,PS-PPO 将每步训练时间减少了 33%–45%。这主要是由于减少了前向/反向成本,抵消了计算截断分布带来的开销。
- 显存: 由于减少了截断序列的激活存储,峰值 GPU 显存使用量降低了 15%–17%。
- 扩展性: 随着最大完成长度(Tmax)的增加,效率提升更加显著。在 Tmax=4096 时,PS-PPO 在更新阶段比 S-GRPO 快 2.8 倍,比 DAPO 快 3.3 倍。
- 消融实验:
- 预算 (B): 预算 B=128 个 token 提供了准确率与速度之间的最佳平衡。较小的预算会增加方差;较大的预算虽然能提高准确率,但计算成本收益递减。
- 截断策略: “优化”策略(求解凸问题)明显优于均匀、时间先验和启发式截断策略,证实了自适应选择信息丰富的序列前缀至关重要。
- Rollouts (K): 增加 K 可以提高稳定性和性能,但会增加每步的计算量。研究确定 K=8 是一个实际的折中方案。
5. 重要性与声明
论文声称 PS-PPO 提供了一条在显著降低资源需求的情况下对齐大语言模型的可扩展路径。其核心意义在于观察到推理轨迹中的中间前缀通常包含足以确定最终结果的信息,使得全轨迹更新在计算上是冗余的。
通过通过随机前缀采样和重要性加权将学习信号与全序列长度解耦,PS-PPO 实现了:
- 降低硬件门槛: 更低的峰值显存和计算成本使得 RLHF 对学术界和资源受限的群体更加友好。
- 高效的长上下文训练: 该方法对于长推理轨迹(其全反向传播成本极高的情况)特别有效。
作者保持了谦逊的语气,指出 PS-PPO 是一种优化方法,旨在提高效率,但它本身并不决定模型的安全性或行为;这些仍取决于奖励规范、训练数据和部署保障措施。该方法在期望上保持了全序列目标的理论特性,确保了效率的提升不会以产生有偏学习为代价。