蒸馏小模型老是跑偏?这篇论文搞了个「接棒」机制,训练轨迹砍半还提点 5.73%
蒸馏过小模型的人,大概率见过这个场景:学生模型在某个推理步骤走歪了,顺着这个错误方向一路狂奔,生成了一长串毫无意义的内容。整条训练轨迹废掉,算力白烧,梯度白算。
做 on-policy distillation(OPD)的人管这个叫「前缀失败」——一旦 prefix 歪了,后面全白给。学生歪的时候,老师给的监督信号也跟着失效,因为老师的 loss 是基于学生自己生成的 token 算的,学生都跑到沟里了,老师在沟外喊话也没用。
浙大团队今天放出来的这篇新论文,解决的就是这个痛点。思路朴素,但效果很硬。
他们把蒸馏做成了接力赛。
学生跑着跑着跑偏了,老师不站在终点喊加油,直接进场接棒,带着跑几步纠正方向,再把棒交回给学生继续跑。论文里叫 Relay-OPD(Trajectory-Relayed On-Policy Distillation),「接力式轨迹蒸馏」。
接棒时机怎么定?论文发现了一个 pattern:在失败前缀上,老师和学生的续写方向会出现明显的「分歧不对称」——老师倾向于纠正方向重新来,学生则沿着原路接着歪。这个差异可以被直接当成一个无监督的触发信号。不需要额外标注,不用人工设阈值,模型自己就能判断「这里该接棒了」。
具体执行上,Relay-OPD 的做法是:训练过程中,检测到触发点时让老师短暂接管,生成一段「老师 leg」,然后把控制权交回学生,学生在完整的接力轨迹上做优化。
因为接棒只发生在关键早期位置,而且老师只接管一小段,整体上学生的策略分布不会偏离太远——论文管这个叫「有限接力预算」。效果也直观看得见:训练轨迹长度直接砍掉 50% 以上。以前跑 1000 个 token 才到终点,现在 500 不到就完事了,而且质量更高。
实验配置是 Qwen3-4B-Instruct-2507 做老师,Qwen3-0.6B 和 1.7B 的 Non-Thinking 版本做学生,在 8 个数学推理 benchmark 上测。结果每项都是第一或第二。1.7B 学生平均比标准 OPD 高出 5.73%,比之前最强的 FastOPD 基线也高出 1.49%。0.6B 版本同样有 consistent gains。
用同样的老师去蒸馏一个小模型,换这套接力方案,训练时间更短、算力更省、最终效果还更好。尤其是数学推理类的蒸馏任务,这个 5.73% 的差距在实际部署里能明显感知到。
我比较喜欢的点是,这个方法没有引入复杂的辅助网络、没有用额外的 reward model、不需要人工标注好坏轨迹。它纯粹利用了老师和学生在错误前缀上的行为差异来做触发判断,算是在系统层面打了个「不对称差」。
当然也有没讲透的地方。论文里触发点的检测具体怎么做、误触发率和漏触发率如何,目前只有实验结果支撑,没有独立的分析模块。如果是生产环境下的蒸馏 pipeline,大概率还得自己调一调接棒灵敏度。
项目代码和论文都开源了,链接在下面。现在手里正好在跑蒸馏任务的话,值得拉下来试试——至少能把那 50% 的白给轨迹省下来。
相关链接
- 论文:https://arxiv.org/abs/2607.26057
- 项目页:https://zju-real.github.io/Relay-OPD
- 代码:https://github.com/zju-real/Relay-OPD