Jet-RL:通过统一训练与推演精度流实现同策略 FP8 强化学习

Jet-RL: Enabling On-Policy FP8 Reinforcement Learning with Unified Training and Rollout Precision Flow

arXiv: 2601.14243v1

论文信息

标题: Jet-RL: Enabling On-Policy FP8 Reinforcement Learning with Unified Training and Rollout Precision Flow

作者: Haocheng Xi, Charlie Ruan, Peiyuan Liao, et al.

发布日期: 2026-01-20

arXiv ID: 2601.14243v1

PDF 链接: 下载 PDF

3 分钟速览

  • 研究问题:本文解决强化学习训练大语言模型时,BF16 训练搭配 FP8 推理(rollout)的策略在长序列生成和困难任务上出现的训练不稳定甚至灾难性崩溃。
  • 核心方法:提出 Jet‑RL 框架,让训练前向和 rollout 推理使用完全相同的 FP8 量化精度流,消除训练–推理精度不匹配,使整个过程真正保持 on‑policy。
  • 关键结果:在 8B 模型上 Jet‑RL 实现端到端训练 1.16 倍加速,同时精度退化通常小于 1%,远优于 BF16‑Train‑FP8‑Rollout 的 5% 以上退化(表 2、表 3)。
  • 主要局限:论文基于 8B–32B 模型,受资源限制未完成 14B–32B 的完整扩展实验;当张量并行度较高时 FP8 推理加速效果下降(表 4);极端任务下仍有约 2–3% 的性能退化。
  • 适合读者:从事大模型强化学习训练、推理优化、量化加速的工程师和研究人员;对分布式训练和 on‑policy 算法有基础理解的读者。

论文背景和研究动机

强化学习(RL)已成为推动大语言模型复杂推理能力的关键范式,特别是在链式思维(CoT)长序列生成中,模型需要通过多步推理来探索解空间。然而,RL 训练中 actor 模型的 rollout 阶段(即自回归生成回复)占据了端到端训练时间的 70% 以上(图 2),成为主要瓶颈。当生成长度超过 8K tokens 时,rollout 延迟占比可高达 75%,严重拖慢整体效率。

为了加速 rollout,工业界普遍采用 BF16 训练 + FP8 推理 的策略:在更新权重时保留 BF16 精度以保证训练稳定性,在生成 rollout 时转换为 FP8 以利用低精度算力的速度优势。这一方法被 VeRL、SLIME、NeMo‑RL 等主流 RL 框架广泛使用。但本文通过系统性分析发现,这种简单混合精度策略存在两个致命缺陷。首先,长序列生成时性能急剧退化:当 rollout 长度从 4K 增加到 8K 乃至 16K 时,BF16‑Train‑FP8‑Rollout 的模型准确率迅速下降,甚至在 Qwen2.5‑7B 上出现完全不收敛的情况(图 3)。其次,困难任务下劣化更严重:当用弱基座模型或高难度数学数据集训练时,BF16 与 FP8 训练的曲线很快分叉(图 4),导致最终性能远低于纯 BF16 基线。

作者分析认为,上述问题的根源在于 训练与 rollout 的精度不一致。RL 训练严格依赖 on‑policy 假设,即 rollout 产生的轨迹应与当前策略一致。而 BF16‑Train‑FP8‑Rollout 使得训练前向(BF16)和推理前向(FP8)形成两套不同的精度图,随序列增长累积小误差,在困难任务中被进一步放大,形成严重的 off‑policy 训练,最终导致优化发散。为此,论文提出必须 统一训练与 rollout 的量化精度流,才能实现稳健的低精度 RL 训练。

核心方法和技术细节

Jet‑RL 通过以下三个层面实现 on‑policy 的 FP8 RL 训练。

统一的 FP8 精度流

Jet‑RL 将模型的前向计算视为一个有向图 G\mathcal{G},节点为算子或权重,边表示张量的传播及精度属性(图 5)。在 BF16‑Train‑FP8‑Rollout 中,训练图 Gtrainfwd\mathcal{G}_{\text{train}}^{\text{fwd}} 精度为 BF16,而推理图 Ginfer\mathcal{G}_{\text{infer}} 中的线性层输入被量化为 FP8,两者明显不一致。Jet‑RL 强制 Ginfer\mathcal{G}_{\text{infer}} 成为 Gtrainfwd\mathcal{G}_{\text{train}}^{\text{fwd}} 的子图,即训练前向中同样使用 FP8 量化,使生成 rollout 的精度流与训练完全一致,从根本上消除策略不匹配。唯一区别是训练中保留 BF16 的 master 权重以积累更新,但前向计算时按同样的量化方案读取 FP8 权重。

细粒度 GEMM 量化方案

Jet‑RL 对所有线性层的三个 GEMM(前向 FProp、反向梯度计算 WGrad 和 DGrad)均采用 FP8 运算以加速。为防止 per‑tensor 量化在训练中引发数值不稳定,论文设计了混合粒度的量化策略(图 6):

  • 权重 采用 128×128128 \times 128 的 per‑block 量化;
  • 激活与梯度 采用 1×1281 \times 128 的 per‑group 量化(沿 token 维度每 128 个元素一组)。

前向 FProp 算子输入激活以 1×1281 \times 128 量化为行主序,权重以 128×128128 \times 128 量化为列主序,直接服用 DeepGEMM 等高效 FP8 矩阵乘内核。反向 DGrad 与 FProp 等价,可复用相同配置;WGrad 则需第一矩阵 1×1281 \times 128、第二矩阵 128×1128 \times 1 量化,并被融合在一起减少重复量化开销。为了维持训练精度,反向传播中 梯度的传输仍保留 BF16,仅计算内核内部量化,这借鉴了此前 COAT 等工作的经验(见第 4.1 节)。

系统实现

Jet‑RL 以 vLLM 作为推理引擎、VeRL 作为 RL 训练框架,利用 DeepGEMM 的高效 FP8 内核以及 Triton 实现的量化、转置和融合算子完成整个流程。权重在参数更新阶段即被量化为 FP8,推理引擎无需每次同步时重新校准,避免了 PTQ 中高昂的数据依赖型校准开销,也进一步强化了 on‑policy 的一致。

创新点和贡献

  1. 首次系统性揭露 BF16‑Train‑FP8‑Rollout 的脆弱性:论文通过多模型、多长度、多数据集的实验表明,这种看似无害的混合精度会导致 off‑policy 训练,在长 rollout 和高难度任务中引发灾难性退化,推翻了其 “不影响精度” 的惯常认知。

  2. 提出 on‑policy FP8 RL 框架 Jet‑RL:核心创新在于将精度流统一为训练与推理共享的 FP8 图,真正满足 on‑policy 假设,使 RL 训练的稳定性与纯 BF16 相当,而加速收益不丢失。

  3. 设计适应 RL 训练的细粒度量化策略:针对 RL 中频繁的正反向计算,采用分块和分组量化的组合,并融合 WGrad 的量化操作,兼顾了计算加速与训练稳定性。

  4. 系统的加速‑精度权衡评估:提供了全面的吞吐量加速数据(rollout 1.07×–1.33×,训练 1.41×,端到端 1.16×)及精度对比(退化通常 <1%),为实际部署提供明确指引。

实验结果分析

精度稳定性:在 8K rollout 长度下,BF16‑Train‑FP8‑Rollout 导致 Llama3.1‑8B 平均得分从 23.2 降至 13.0(绝对下降 10.2%),Qwen2.5‑7B 甚至完全不收敛。Jet‑RL 则将退化控制在 1.0% 以内,部分模型(如 Qwen3‑8B‑Base)平均得分仅比 BF16 基线低 1.1%(表 2)。在 16K 长度和 DeepMATH 高难度数据上,Jet‑RL 同样稳健,相对退化通常 ≤2.7%,而 BF16‑Train‑FP8‑Rollout 再次出现不收敛或 10.3% 的巨幅下跌(表 3)。

加速效果:rollout 阶段 FP8 对 BF16 的加速随着模型规模增大而提升,32B 模型 TP=2 时达 1.33×,但 TP 增至 4 时因通信开销下降到 1.07×(表 4)。整体上,Jet‑RL 对 Qwen3‑8B 实现了 1.54× 的 actor 更新加速、1.80× 的参考模型推理加速,训练阶段综合提升 1.41×,端到端步时加速 1.16×。这些数据说明 Jet‑RL 在保持精度的同时确实能有效压缩训练时间。

值得注意的是,Jet‑RL 在某些强基座、简单任务上甚至略微超过 BF16 基线(如 Llama3.1‑8B 上平均得分 +2.0%),这可能是因为少量量化噪声起到了隐式正则化的作用,但论文并未对此展开分析。

实践建议

对于希望将 FP8 量化落地到 LLM 强化学习训练中的团队,Jet‑RL 提供了明确的设计路线。

  • 不要采用分离精度方案:避免 BF16 训练与 FP8 rollout 的混合,其在长序列和高难度任务上的风险已被充分证实(图 3、图 4)。若现有推理引擎已支持 FP8,务必同时改造训练流程以保证精度流统一。

  • 选择合适的量化粒度:对于 7B–8B 模型,权重 128×128128 \times 128 块量化、激活 1×1281 \times 128 组量化的方案在精度和速度间取得良好平衡。训练时注意将前向量化的激活储存在 FP8 中以节省显存,反向梯度则保持 BF16 精度,这对防止梯度下溢至关重要。

  • 控制张量并行度:FP8 推理加速在低 TP 度时收益最大,TP=4 时通信开销会显著侵蚀加速(表 4)。因此,在 GPU 数量充裕的情况下,可优先采用更低的 TP 度来发挥 FP8 的计算优势。

  • 系统集成的实现要点:Jet‑RL 的权重量化与参数更新耦合在一起,可避免每次权重同步后重新校准的开销。实际部署时,可将权重转换为 FP8 的步骤嵌入到 VeRL 等训练框架的 optimizer step 之后,并通过共享内存或 RDMA 直接将量化权重推送给 vLLM 引擎,保持训练与推理的同步一致。

  • 关注更大模型的收益:虽然论文未给出 14B 以上模型的端到端加速,但 rollout 加速随模型规模增大的趋势(表 4)暗示,Jet‑RL 在 30B+ 模型上可能带来更显著的整体提速。资源允许时应优先在超大模型上验证 Jet‑RL。