通过最优系数校准在强化学习中联合训练多令牌预测
可验证奖励强化学习 (RLVR) 已成为提升大语言模型推理能力的标准范式,而多令牌预测 (MTP) 是预训练中广泛采用的模块。将两者结合是自然思路,但当前 RL 实践会分离 MTP 梯度,因为联合训练会降低性能。本文从优化视角重新审视这一失败。 我们证明,MTP 对 RL 目标的逐步骤影响可分解为两项:一阶相关性和二阶扰动惩罚。该分解统一了三种 MTP 训练模式:Detach、Cross-Entropy loss 和 Policy loss,并解释了每种模式成功或失败的原因。进一步分析 Policy loss 发现,尽管其符合直觉,性能仍会下降:相关性项衰减而二次惩罚项持续存在。 基于此分析,我们提出最优系数校准 (OCC),一种自适应方案,通过对数概率代理在线跟踪最优系数,且代价可忽略。在六个竞赛级数学推理基准上,OCC 一致达到或超过 Detach 基线,实现了改进的联合 MTP-RL 训练性能。
论文精读
TL;DR 通过分解 MTP 对 RL 目标的效应,揭示策略损失退化原因,并提出在线自适应校准最佳系数的 **OCC** 方法,实现 MTP 梯度与 RL 的稳定联合训练,在数学推理基准上超越传统分离方案。
问题
问题背景
基于可验证奖励的强化学习(RLVR)已成为提升大语言模型推理能力的主流范式,而多令牌预测(MTP)在预训练中被广泛采用。将 MTP 直接集成到 RL 训练中自然合理,但直接联合训练会损害模型性能,迫使实践中切断 MTP 梯度。
现有方法局限
- Detach 模式:完全冻结 MTP 梯度,放弃了预训练阶段积累的有利归纳偏置,导致信息浪费。
- 固定辅助损失:尝试用交叉熵(CE)或策略损失(Policy Loss)作为辅助目标,但缺乏理论指导。CE 损失的一阶相关项消失(见论文附录 B),无法提供有效梯度;策略损失虽直觉合理,其有益的一阶项随训练迅速衰减,而二阶扰动惩罚持续存在,导致后期性能退化。
- 无动态调整:固定系数无法适应训练过程的动态变化,难以平衡主任务与辅助任务。
问题难点与重要性
RL 目标与 MTP 辅助损失间的交互可分解为一阶相关性与二阶惩罚。难点在于:
- 动态博弈:早期相关项主导,提供有益信号,后期衰减后惩罚项主导,优化方向恶化。
- 缺乏理论指导:无法在线感知这种动态变化,难以自适应调节系数。
该问题阻碍了预训练与 RL 后训练的有效协同,对于追求极致数学推理能力的工程实践至关重要。
行业类比:类似多任务学习中的梯度冲突——辅助任务在不当时刻介入会破坏主任务收敛,需要自适应权衡机制。
核心洞察
- 将MTP引入RL训练造成的性能退化归因为**一阶相关项衰减与二阶惩罚项持续**的解耦视角。不同于以往直接分离梯度或替换损失的工程技巧,该分析从优化动力学层面揭示了策略损失的内在矛盾:早期有用的相关系数会随策略收敛而消失,但二阶惩罚却不会减少,最终拖垮性能。这种量化归因为其他辅助模块与RL的联合训练提供了可迁移的诊断范式。
- **最优系数校准(OCC)**通过仅需log概率梯度代理的在线估计,以极低成本自适应维持MTP辅助任务的理想强度。与固定系数或需要额外超参扫描的方案相比,OCC不需调参即可在多个benchmark上稳定超越detach基线,且跨RL算法和基座模型泛化,对实际部署MTP+RL管线极具工程价值。
方法
问题设定与输入
在 RLVR (Reinforcement Learning from Verifiable Rewards) 框架下训练语言模型时,通常联合使用 Multi-Token Prediction (MTP) 模块——该模块在预训练中已被证明能提升表示质量,但在 RL 阶段,直接联合优化 MTP 与策略梯度会导致性能下降,因此当前实践普遍 detach (切断) MTP 分支的梯度,使其不参与 RL 更新。
本方法的核心输入是:一个带 MTP 头的策略模型、RL 优化器以及可验证奖励信号。目标是在不牺牲 RL 性能的前提下,让 MTP 梯度也参与联合训练。
关键模块:从理论分解到自适应系数校准
作者首先将 MTP 对 RL 目标的单步影响分解为两项:
- 一阶相关项:衡量 MTP 预测方向与策略梯度方向的对齐程度,本质是 MTP 能否为 RL 提供有益的驱动信号;
- 二阶扰动项:类似于一个二次惩罚,限制因 MTP 引入的策略分布偏移。
这一分解统一解释了三种常见训练策略的成败:
- Detach 同时消除相关项与惩罚项,故稳定但未利用 MTP 信息;
- Cross-Entropy (CE) loss 使得相关项消失,仅剩惩罚项,导致性能始终低于 detach 基线;
- Policy loss 保留了相关项,直觉上更合理,但实验发现相关项在训练后期迅速衰减,而惩罚项持续存在,最终仍损害 RL 优化。
基于此,本文提出 Optimal Coefficient Calibration (OCC):在每一步训练中,自适应地调节 MTP 损失项的权重系数 λ,使得一阶相关收益与二阶惩罚达到最优平衡。闭式最优 λ 的表达式涉及未来状态的期望值,难以直接计算。OCC 的关键技巧在于引入 log-probability 梯度代理:利用 MTP 输出头的 log 概率梯度的统计量,作为无偏、低成本的代理信号,在线估计最优系数,几乎不增加训练开销。
输出与效果
OCC 输出一个动态标量 λ,用于缩放 MTP 辅助损失。与固定系数或手动衰减策略不同,OCC 能自动在相关项较强时增大权重、在惩罚占优时减小权重,使 MTP 梯度始终正向贡献。在多项竞赛级数学推理基准上,OCC 联合训练的表现一致匹配甚至超越 detach 基线,且无需任何额外超参搜索。
差异点:OCC 首次从优化分解角度阐明 MTP-RL 联合训练的失效机理,并提供了一种理论上可解释、实现上零额外成本的自适应权重校准方案,完全避免固定系数调度带来的调试负担。
实验
实验设计
论文在六个竞赛级数学推理基准上评估联合训练效果,覆盖多种 RL 算法(如 PPO、GRPO 等)和不同的基座模型(Mistral、DeepSeek 等)。对比三种 MTP 训练模式:
- Detach:梯度完全分离,作为当前实践中的标准基线
- Cross-Entropy (CE) Loss:使用固定的 CE 损失函数进行联合训练
- Policy Loss:直接用策略梯度将 MTP 纳入 RL 目标
提出的 OCC (Optimal Coefficient Calibration) 在 Policy loss 框架下引入自适应系数,通过一个对数概率梯度代理(log-probability proxy)在线估计最优权重,以极低的计算开销动态平衡一阶相关项与二阶扰动惩罚项。
关键发现
- CE loss 在所有基准上一致劣于 Detach,证实单纯引入 MTP 信号会破坏 RL 训练。
- OCC 能够匹配或超越 Detach 基线,在多个设置下实现更优的最终性能,且这一优势在不同 RL 算法和基座模型上具可复现性。
- Policy loss 即使精细调节固定系数,也无法达到 OCC 的水平,验证了自适应方案的必要性:早期相关项主导时需要较大系数,后期惩罚项主导时权重必须衰减,固定值无法兼顾。
- OCC 的训练时间开销可忽略,且对超参数不敏感。
与基线的对比解读
Detach 通过阻断 MTP 梯度避免了性能退化,但完全丢弃了 MTP 可能带来的加速或正则化收益。OCC 的核心贡献在于识别并量化了“破坏”的来源——二阶扰动惩罚项随训练进行而持续存在,一阶有益相关项却在后期衰减。OCC 的动态系数同步追踪这一变化,使得训练能保留有益的探索信号,同时抑制有害的二次扰动。相较于 Policy loss 的静态加权,OCC 实质上是在每个优化步上求解一个闭式最优参数,从而在不牺牲训练效率的前提下,实现了对简单梯度分离的稳健改进。这一分析框架也为其他联合训练场景提供了借鉴:当两个目标存在冲突梯度时,动态权重校准可能比硬性分离更有效。
行业影响
落地场景
Multi-Token Prediction (MTP) 与 Reinforcement Learning from Verifiable Rewards (RLVR) 联合训练,主要面向需要多步推理且存在可验证奖励信号的场景:
- 数学 / 代码推理产品:教育辅助工具、编程助手等,需模型在复杂问题求解中保持长链一致性,MTP 能促进未来 token 规划,而 RLVR 提供结局奖励。
- 智能客服与对话系统:当多轮交互存在最终成功标准(如工单解决率)时,联合训练可让模型在生成当前回答时已优化后续轮次的结果。
- 内容生成与审核平台:例如对长文本摘要或报告,若最终有事实准确性评分作为奖励,MTP 有助于提升整体连贯性,避免中期发散。
商业价值
- 降低训练成本,提升样本效率:论文提出的 Optimal Coefficient Calibration (OCC) 以可忽略代价自适应调整联合损失系数,避免了工程上繁琐的手动调参和反复实验,直接缩短训练周期。
- 提升模型最终性能:在六个竞赛级数学推理基准上,OCC 一致达到或超越梯度解耦基线(Detach),这意味着在不牺牲 RL 收益的前提下可额外获得 MTP 的正则化优势,最终模型准确率更高,直接关联用户满意度和产品竞争力。
- 增强模型鲁棒性:通过动态平衡一阶相关与二阶惩罚,训练更稳定,减少因梯度冲突导致的模型退化,降低线上事故风险。
与现有产品/工作流的集成
- 训练框架轻量插件:OCC 仅需维护一个对数概率代理(log-probability proxy)来实时估计最佳系数,不改变模型结构,可直接嵌入如 Hugging Face TRL、DeepSpeed 等 RLHF/RLVR 流程,作为 loss 计算的一个自适应权重。
- 无需额外数据:系数校准完全基于当前 batch 的在线梯度统计,不需要存储额外状态或修改数据管道,现有 CI/CD 训练流水线几乎无改动。
- 兼容多种 RL 算法:论文验证在 PPO 和 GRPO 上均有效,表明该方法可作为通用模块,适配不同产品团队已有的 RL 优化器选择。
具体落地 Use Case
- 编程教育平台(如 LeetCode 类产品的 AI 导师):利用 RLVR 以编译/测试结果作为奖励,同时通过 MTP 使模型在生成一行代码时已“预判”后续几行,减少语法错误或逻辑断裂。OCC 确保训练平稳,直接提升解题正确率和教学体验。
- 金融报告自动生成:对给定数据生成分析报告,最终以合规性、关键指标覆盖度作为可验证奖励。MTP 帮助模型保持段落间因果链条,OCC 则让训练避免因多目标梯度冲突导致事实幻觉,提升报告质量,降低人工审核成本。
局限
- OCC 的有效性依赖于 log-probability 代理对一阶相关性的准确估计,当动作空间巨大或奖励稀疏导致梯度噪声较高时,该代理可能引入偏差,使自适应系数偏离理论最优。此外,理论分解假定 MTP 头与主模型共享相同底层表示,若未来模型采用异构 MTP 架构,推导需重新验证,限制了方法的即插即用性。
- 实验设计局限于竞赛级数学推理(如 MATH、GSM8K 等六个基准),且仅在 7B–13B 参数规模上测试。对于代码生成、指令遵循等更广泛的 RLVR 场景,以及更大规模模型(如 70B+)上的表现缺乏验证。同时,未与动态任务权重、梯度投影等自适应训练策略进行对比,难以判断 OCC 是否具有通用优势。
- 作者在 Limitations 部分指出,OCC 目前仅适用于 RL 后训练阶段,未探索在预训练中结合 MTP 的联合优化。另外,log-prob 代理计算虽被描述为开销可忽略,但在长序列生成时,额外的前向传递和存储需求可能在内存受限的部署环境中成为瓶颈,实际工程成本需更精细评估。