接力棒:轨迹接力的在线策略蒸馏 Pass the Baton: Trajectory-Relayed On-Policy Distillation
在学生推理出错处让老师短暂接管再交还,修复在线策略蒸馏的前缀失败
前置知识
在线策略蒸馏(On-Policy Distillation, OPD)
区别于用老师生成的数据做监督微调(SFT)或离线知识蒸馏,OPD 让「学生模型」用自己的策略 rollout 出轨迹 $y=(y_1,\dots,y_N)\sim\pi_{\bar\theta}(\cdot|x)$,再由老师在学生实际访问到的前缀 $h_t=(x,y_{<t})$ 上对每个 token 提供密集的 token 级指导。由于监督始终落在学生自己的状态分布上,能有效缓解训练-推理分布偏移,是当前强到弱蒸馏的标准做法。
Relay-OPD 正是在标准 OPD 的框架上做改造,不理解 OPD 的轨迹采样和 token 级优势就无法理解它要修复的「前缀失败」和它最终的优化目标。
前缀失败(Prefix Failure)
在长链推理中,学生一旦在早期某个 token 走错方向(比如对一个错误结论继续推导),由于自回归特性,后续所有生成都建立在这个偏差之上,形成长段「被误导的延续」。这些延续既产生不可靠甚至有害的监督信号,又浪费大量训练算力。
前缀失败是本文要解决的核心病灶,所有方法设计(handoff 触发、teacher leg、relay budget)都是针对它而做的,是理解论文动机的主线。
推测解码(Speculative Decoding)
一种加速推理的技术:用一个小的 draft 模型批量生成若干候选 token,再用大的 target 模型一次性验证;对每个 draft token $a^S_t$ 以 $\alpha_t=\min(1,\pi_{tgt}(a^S_t|h)/\pi_{\bar\theta}(a^S_t|h))$ 的概率接受,拒绝时从残差分布重采样。其正确性保证被接受的 token 严格服从 target 分布。本文把学生当 draft、老师当 target,统一到一个引擎里实现 relay rollout。
Relay-OPD 的高效实现完全建立在推测解码之上,是其能在训练中无缝切换 student/teacher leg、且零额外开销计算触发信号的关键,跳过它就难以理解 §3.3 的精确性论证。
反向 KL 散度与单样本优势(Reverse-KL Single-Sample Advantage)
标准 OPD 在 token $y_t$ 上用单样本估计反向 KL:$\hat D^{(t)}_{RKL}=\log\pi_{\bar\theta}(y_t|h_t)-\log\pi_T(y_t|h_t)$,取其负号作为优势 $A^{OPD}_t=\log\pi_T(y_t|h_t)-\log\pi_{\bar\theta}(y_t|h_t)$。正值鼓励学生提高该 token 概率,负值则抑制。反向 KL 具有 mode-seeking(选模)特性,使学生能选择性吸收老师的纠正信号。
Relay-OPD 在 teacher leg 上坚持用反向 KL 单样本目标而不用前向 KL,这一点是消融实验中拉开差距(46.96 vs 44.08)的关键,必须先理解 KL 优势的语义。
研究动机
在线策略蒸馏虽然把监督建立在学生自己的轨迹上、缓解了训练-推理分布偏移,但也把学生的失败一起搬了进来。在长链数学推理里这会导致严重的前缀失败:学生一旦在早期某个位置选错方向(例如对一个不成立的方案继续往下算),自回归生成会让整条后半段轨迹都建立在偏差之上,产生又长又跑偏的延续。这些延续不仅得到的 token 级监督越来越不可靠、甚至有害,还会白白消耗大量训练算力。作者用 Qwen3-4B-Instruct-2507 作老师、Qwen3-1.7B-Non-Thinking 作学生在 DAPO-Math-17K 上做诊断:学生单跑平均准确率仅 27.73%,而老师可达 60.55%。已有补救办法都有结构性短板:固定长度截断(ESR、FastOPD)在死板位置一刀切,无法识别真正失败的位置;离线重写(TRD)只能在 rollout 结束后修补,且常留下明显的改写痕迹(18.96% 的轨迹含『original solution』『rewrite』等改写残留);token 级混合(SKD)依据的是泛化的分布差异,而非「推理方向是否已经走错」的明确信号,导致它一旦陷入重复生成模式就难以打破。
本文的目标是作者希望设计一种「在线、即时」的纠正机制:当推理失败正在发生时就把它截住并修好,而不是事后修补;更重要的是,要在哪里干预、要不要干预这件事,必须由推理状态本身(也就是当前前缀)来决定,而不依赖任何外部验证器、过程标签或奖励模型。这个机制还要满足两个看似矛盾的要求——既要把纠正集中在前缀失败真正起源的早期关键位置(实验显示把干预往后挪会显著掉点),又要避免老师接管得太多太长、把轨迹拽离学生自身策略从而破坏 on-policy 性质。最终目标是在不引入外部监督的前提下,同时提升蒸馏精度、缩短训练轨迹长度。
与已有工作不同的是,本文最独特的切入点是发现并利用了「老师-学生延续不对称」这个可在线观测、且无需标签的现象:在学生已经走偏的前缀上,老师的下一步倾向于停下来反思并改方向(top-1 token 往往是 But/Wait/However 这类反思词),而学生则倾向于沿着错误方向继续往下(top-1 往往是 So/Now)。这种方向级分歧恰好标记了「推理方向已坏」的位置,作者称之为 handoff trigger。这与以往基于泛化分布差异(SKD)、固定位置截断(FastOPD)或离线重写(TRD)的做法本质不同——它是对「推理状态」本身的显式信号驱动,而且是 label-free 的,任何人只要有老师-学生 logits 就能算出来。
核心方法
Relay-OPD 的直觉非常直观,就像接力赛跑:学生按自己的策略 rollout 推理轨迹,引擎在每一步实时监测前缀,一旦发现「老师会改方向、学生会继续错」的 handoff 触发点,就让老师短暂接管跑一段(teacher leg,最多 $L$ 个段落),把推理扳回正轨后把棒交还学生继续。整个 rollout 受一个 relay 预算 $(M,L)$ 约束:最多 $M$ 次接管、每次最多 $L$ 个段落,从而把干预集中在早期关键位置又不至于偏离学生策略太远。技术路线上有三个支柱:(1) 基于 top-K 支持集比较的 handoff 触发判据 $\phi(h)$;(2) 由 student leg 和 teacher leg 拼接而成的 relay 轨迹构造;(3) 把整套 relay rollout 统一在单个推测解码引擎里实现——学生当 draft 模型、老师当 target 模型,借此零额外开销地拿到触发信号并精确复现两模型的接力过程。学生最终在拼好的 relay 轨迹上,用反向 KL 单样本目标做 PPO 式更新,$K=5$、$(M,L)=(2,3)$ 是默认配置。
核心创新是把『老师-学生延续不对称』转化为无需标签的 handoff 触发判据 $\phi(h)=\mathbf 1[a_T(h)\in\mathcal R]\cdot\mathbf 1[K_S(h)\cap\mathcal R=\varnothing]$,其中 $a_T(h)=\arg\max_v\pi_T(v|h)$ 是老师 argmax 次词、$K_S(h)=\mathrm{TopK}_K(\pi_{\bar\theta}(\cdot|h))$ 是学生 top-K 支持集、$\mathcal R$ 是 Wait/But/However 等反思词及其大小写变体集合。当老师最想要的下一个词是反思词、而学生 top-K 里一个反思词都没有时,就判定学生正在走偏并触发接管。与已有方法的本质区别:相比固定截断,干预和终止位置由当前推理状态决定、而非一刀切位置;相比离线重写,纠正实时发生、不留改写痕迹;相比 token 级混合,触发依据的是『方向是否走错』而非泛化分布差异。配合 reverse-KL 单样本目标使学生选择性吸收老师纠正信号,消融证明它优于前向 KL(44.08)和 student draft(44.56),达 46.96。
方法步骤详情
算法(Algorithm 1):初始化空轨迹 $z$、接管计数 $j=0$、状态 $s=S$,循环到 EOS 或达最大长度。每步取前缀 $h=(x,z)$,学生采样 draft token $a^S\sim\pi_{\bar\theta}(\cdot|h)$,并行算老师 argmax $a_T=\arg\max_v\pi_T(v|h)$ 和学生 top-K 集 $K_S(h)$,评估判据 $\phi(h)$。若 student 状态、触发为真且 $j<M$,则进入 teacher leg:$j$ 自增、状态置 $T$,把反思词 $a_T$ 直接作为 teacher leg 首 token(不经验证),再用推测解码生成 $L$ 个段落;若 $j=M$ 终止 rollout,否则状态回 $S$ 让学生续跑。若不触发或预算耗尽则把 $a^S$ 追加进轨迹,得 relay 轨迹 $z$。训练时对每 token 用反向 KL 单样本优势 $A_t=\log\pi_T(z_t|h^z_t)-\log\pi_{\bar\theta}(z_t|h^z_t)$ 做标准 PPO 裁剪更新($\epsilon=0.2$,重要性比 $\rho_t=\pi_\theta/\pi_{\bar\theta}$)。每次更新后同步学生权重到 draft,老师冻结。默认 $K=5$、$(M,L)=(2,3)$。
技术新颖性
技术新颖性体现在三处。第一,handoff 触发的思路本身——把方向级分歧而非分布级分歧作为信号,且完全 label-free、只靠 logits。第二,relay 预算 $(M,L)$ 用段落而非固定 token 数度量 teacher leg(平均一段约 23.2 token),保证每段在『结构完整』的推理单元处结束而非断在半句话,并用 $M$ 上限把干预集中在前缀失败的早期起源。第三也是最关键的工程创新——把整个 relay 过程塞进单个推测解码引擎:通过状态机 $s_t\in\{S,T,\bot\}$ 决定每个位置的 target 策略(student 状态下 target 就是学生自己、draft 全部接受,退化为普通采样;teacher 状态下 target 是老师、执行标准推测拒绝采样),从而精确复现两模型接力(由推测采样的正确性保证),又顺带在验证中零成本拿到 $a_T(h)$ 和 $\phi(h)$。消融显示即便只换一个反思词(teacher token 仅占 0.35%)也能让准确率从 27.73 提到 34.96,证明纠正可极度局部。
实验结果
主实验(Table 1)用 Qwen3-4B-Instruct-2507 当老师、Qwen3-0.6B 和 1.7B-Non-Thinking 当学生,在 AIME24/25/26、MATH500、AMC2023、OlympiadBench、HMMT Feb26 和 HMMT Nov25 八个数学基准上评测。对 1.7B 学生,Relay-OPD 平均准确率 46.96、八项全部最佳或次佳,超过标准 OPD 的 41.23 共 +5.73%、超过最强基线 FastOPD 的 45.47 共 +1.49%,其中 AIME25/26 相对 OPD 分别 +7.29%/+7.19%;对 0.6B 学生平均 31.04,比 OPD 28.03 提升 +3.01%。训练效率上,1.7B 学生平均 rollout 长度仅 2296 token、比 OPD 的 4658 缩短 50.7%,且在更早的 step 35 达到最优(OPD step 55、FastOPD step 45);0.6B 轨迹长度 2490、比 OPD 缩短 63.9%。基线对比上,TRD 因改写痕迹反而掉点(1.7B 30.69 vs OPD 41.23),SKD 增益微弱(1.7B 42.35、0.6B 反降到 24.38)。训练动力学(Figure 6)显示 teacher token 占比从约 13% 在 20 步后稳定到 2%–3%。
查看结构化数据
| 任务 | 指标 | 本文 | 基线 | 提升 |
|---|---|---|---|---|
| AIME 2025(1.7B 学生) | 平均准确率 mean@32 | 32.81 | OPD 25.52 | +7.29 个百分点 |
| AIME 2026(1.7B 学生) | 平均准确率 mean@32 | 30.52 | OPD 23.33 | +7.19 个百分点 |
| 八个数学基准平均(1.7B) | 平均准确率 | 46.96 | OPD 41.23 / FastOPD 45.47 | 比 OPD +5.73、比 FastOPD +1.49 |
| 八个数学基准平均(0.6B) | 平均准确率 | 31.04 | OPD 28.03 / FastOPD 30.42 | 比 OPD +3.01、比 FastOPD +0.62 |
| 训练 rollout 响应长度(1.7B) | 平均 token 数 | 2296 | OPD 4658 / FastOPD 2709 | 比 OPD 缩短 50.7% |
| teacher-leg 训练目标消融 | 八基准平均 | Relay Token 46.96 | Teacher FKL 44.08 / Student Draft 44.56 | 反向 KL 单样本 +2.88/+2.40 |
局限与改进
作者在附录 F 中坦承三点局限。第一,评测局限于数学推理且只用 Qwen3 师生对,relay 机制本身虽是任务无关的,但代码生成、智能体工具调用等领域的验证留待未来,换模型族时反思词集合 $\mathcal R$ 可能需要重新调整。第二,Relay-OPD 预设老师的续写能比学生更可靠地纠正走偏前缀,因此当师生能力差距收窄时收益会衰减。第三,relay 预算 $(M,L)$ 是在 1.7B 学生上调好再直接套到 0.6B 上的,虽然敏感性分析(§4.5)显示在适度区间内鲁棒,但换新模型对可能需要重新调参。我的额外观察是:handoff 触发完全依赖一个手工设定的反思词列表,这种基于词表的离散信号对分词器很敏感,跨语言或不同 tokenizer 时可能失效;同时论文未给出 teacher leg 接管错位置(误触发)的比例分析,而这正是该类启发式方法最令人担心的失败模式。
独立分析的弱点
第一,handoff 判据用固定的反思词集合 $\mathcal R$,属于强先验:对英文数学推理的 Qwen3 tokenizer 有效,但换到中文、代码或其他模型族时,反思「信号词」可能完全不同,需要可学习的触发判据(例如用一个轻量分类头或连续向量相似度)来替代词表匹配。改进方向是让触发判据自适应、甚至随训练动态更新。第二,relay 预算 $(M,L)$ 和 top-K 都需要手调,论文虽给了敏感性曲线(Figure 7)但仍是离线选定;不同学生规模、不同任务的最优点可能漂移。可考虑用课程式或自适应预算——随着训练推进、师生差距缩小自动减小 $M$(Figure 6 的预算耗尽率下降趋势其实暗示了这种自适应的可行性)。第三,论文只在数学推理上验证,未覆盖代码、工具使用等多步决策场景,而这些恰恰是前缀失败最严重的地方;缺少对「误触发」(本不该接管却接管)的量化分析,使得该启发式方法的鲁棒性边界不够清晰。改进方向是在代码/智能体任务上扩展评测,并报告接管命中率、误接管率等诊断指标。
未来方向
作者明确提出的方向是把 relay 机制推广到数学推理以外的领域,例如代码生成与智能体工具调用,并在切换模型族时重新校准反思词集合;同时承认随着师生能力差距缩小收益会衰减,需要研究小 gap 场景下的适用性。基于本文成果可延伸的方向我认为有:把固定的反思词集合升级为可学习的、token 级或表示级的触发判据,从而摆脱对分词器和语言的依赖;把 relay 预算做成随训练自适应(参考 Figure 6 中预算耗尽率从 75%–85% 降到 50%–60% 的趋势)或基于师生熵/差距做动态调度;将 handoff 触发的思想与外部验证器、过程奖励模型结合,在有标签的领域进一步抑制误触发;以及把单次 relay 推广到多轮对话或多 agent 场景,研究接力在不同 step 粒度上的迁移性。
复现评估
复现友好度较高。作者提供了 GitHub 仓库(github.com/zju-real/Relay-OPD)和项目主页,训练完全建立在开源框架 verl 和 vLLM 0.21.0 之上,运行环境是 8 张 H100 GPU。论文给出几乎所有关键超参:最大 prompt 2048、最大响应 16384、rollout 温度 1.0、top-p 1.0、全局 batch 128、PPO mini-batch 128、PPO epoch 1、裁剪 0.2、学习率 $1\times10^{-6}$ 恒定、训练 1 epoch,Relay-OPD 默认 $K=5$、$(M,L)=(2,3)$;反思词列表、训练/推理 prompt 模板、各基线实现细节和 Algorithm 1 都齐全。训练数据是公开的 DAPO-Math-17K 英文子集,评测基准也都公开。主要门槛在于 8×H100 算力以及对 vLLM/verl 自定义解码流程的工程改造(要把 relay 状态机和推测解码打通),需要一定系统功底,但整体可复现性在同类工作中属上乘。
论文图表
三联图。(a) 用一个买水果的实际案例展示:在前缀『15 个苹果后只剩 5 RM』处触发,老师 top-1 是 But(74.4%)而学生想继续 So(50.6%),老师短暂接管指出 5 RM 买不起 9 RM 的芒果+木瓜,学生接回后改试 12 个苹果得到正确答案 15。(b) 对比标准 OPD 与 Relay-OPD 的差异:标准 OPD 学生平铺 misdirected 续写,Relay-OPD 在 trigger 处插入 teacher leg 后学生恢复。(c) 给出八基准整体性能。
这张图用最直观的方式同时讲清了动机(前缀失败)、核心机制(handoff trigger + teacher leg + 交接回学生)和最终收益,是理解整篇论文的最佳入口。
柱状图给出三个基准上的平均推理响应长度:Relay-OPD 分别为 15.6k、18.2k、14.9k,相比 FastOPD 的 18.2k、18.1k、20.8k 分别缩短 17.9%、14.2%、28.3%,且同时准确率更高(+2.39%/+4.17%/+1.14%)。
这张图直接支撑了论文『既更准又更短』的关键卖点:Relay-OPD 让学生学会在偏差累积前就改向,得到更短却更准的推理过程,是与 FastOPD 这类截断方法比较时的核心论据。