掩码蒸馏:将思维链内化到语言模型中

arXiv cs.AI 论文

摘要

掩码蒸馏是一种知识蒸馏框架,它训练学生大语言模型仅预测解决方案令牌,同时推理教师提供反馈,旨在将思维链计算内化到模型参数中。该方法显示任务相关的成功,在GSM8K上有效,但对于像Countdown这样更困难的任务需要小型脚手架。

arXiv:2607.22629v1 公告类型:新 摘要:大型推理模型(LRMs)在推理时生成冗长、显式的中间步骤链,然后产生最终答案。这些中间轨迹主导了延迟、内存使用和服务成本,尽管最终答案的正确性与轨迹的正确性没有因果关系,并且轨迹长度并不是问题复杂度的可靠指标。这提出了一个自然的问题:能否将这些中间令牌中表达的计算内化到语言模型的参数中,使其直接(或以更短的中间轨迹)产生答案?我们引入了 \textit{掩码蒸馏},这是一种知识蒸馏框架,其中学生大语言模型(LLM)被训练为仅预测基于问题条件的解决方案令牌,而推理教师在基于问题及其自身CoT轨迹的条件下,对学生的回答提供反馈。我们在这两种设置中实现了该框架:(i) \textit{自蒸馏}设置,其中同一模型在思考模式下作为教师,在非思考模式下作为学生;(ii) \textit{双模型}设置,其中较大的推理教师在解决方案令牌上监督一个单独的较小的非思考学生。通过将中间令牌视为推理模型用来拟合解决方案令牌的脚手架,我们另外变化学生被监督的中间令牌脚手架的长度,在全内化(学生仅输出解决方案)和无内化(学生在答案前输出完整轨迹)之间插值。我们通过在两个推理领域(GSM8K(小学数学)和Countdown(数字谜题搜索任务))上进行受控实验来评估该框架。
查看原文
查看缓存全文

缓存时间: 2026/07/28 06:26

# 掩码蒸馏:将语言模型中的思维链内化

**来源:** [https://arxiv.org/html/2607.22629](https://arxiv.org/html/2607.22629)

**Durgesh Kawlar** SCAI, 亚利桑那州立大学 dkalwar@asu\.edu & **Vardhan Palod*** SCAI, 亚利桑那州立大学 vpalod@asu\.edu & **Subbarao Kambhampati** SCAI, 亚利桑那州立大学 rao@asu\.edu
*同等贡献 - 联合第一作者*
已被 FoGen 2026: 深度生成模型基础:理解记忆、泛化与推理研讨会(ICML 2026 研讨会)接收。

###### 摘要

大型推理模型(LRM)在推理时生成最终答案之前,会产生长而显式的中间步骤链。这些中间轨迹主导了延迟、内存使用和服务成本,尽管最终答案的正确性与轨迹的正确性并无因果关系,且轨迹长度也并非问题复杂度的可靠指标。这引出了一个自然的问题:这些中间令牌所表达的计算能否被内化到语言模型的参数中,使其能够直接(或通过更短的中间轨迹)生成答案?我们提出了**掩码蒸馏**,这是一种知识蒸馏框架。在该框架中,学生LLM被训练为仅根据问题预测解决方案令牌,而推理教师则在根据问题及自身的CoT轨迹进行条件化后,对学生的响应提供反馈。我们在两种设置中实例化这个框架:(i) **自蒸馏**设置,其中同一模型既充当思考模式下的教师,又充当非思考模式下的学生;(ii) **双模型**设置,其中更大的推理教师监督一个独立的、更小的非思考学生,关注于解决方案令牌。通过将中间令牌视为推理模型用来拟合解决方案令牌的脚手架,我们进一步改变了学生所受监督的中间令牌脚手架长度,在完全内化(学生仅输出解决方案)和无内化(学生在答案前输出完整轨迹)之间进行插值。我们通过两个推理领域的控制实验来评估该框架:GSM8K(小学数学)和 Countdown(一种数字谜题搜索任务)。我们的结果表明,完全内化的成功与否严重依赖于具体任务,并追踪学生在预训练期间对该任务的先前暴露情况:它在GSM8K上有效,但在Countdown上如果没有推理时的脚手架则会失败。然而,为学生模型提供一个小的脚手架可以缩小Countdown上与自蒸馏的差距。任务性能从41.7%(完全掩码)提高到86.2%,基本上匹配了教师模型(87.3%),且推理令牌比非掩码变体减少了约1.3倍。

## 1 引言

经过强化学习后训练的大型推理模型(LRM)通过生成较长的中间令牌链再产生最终答案,在数学、代码和规划基准上取得了强劲的性能。虽然这些轨迹提高了任务性能,但它们也显著增加了推理成本。LRM将其大部分生成预算花在了中间令牌上,而不是答案上,而推理延迟、KV缓存占用和能耗都大致随轨迹长度线性增长。在部署规模上,这是服务它们的主要成本。一系列广泛的工作旨在通过在后训练RL阶段添加长度控制目标来降低LRM的这种成本\[2, CoT-Valve\[17\], O1-Pruner\[16\], Kimi-1.5风格的长度惩罚奖励\[25\], GFPO\[23\],\[3\]\]。这些方法奖励更短的轨迹,但模型在推理时仍然生成显式轨迹,且节省量受限于轨迹能短到何种程度。虽然中间轨迹能提升性能是众所周知的,但其原因尚不明确。我们最近的工作表明,轨迹正确性与最终答案正确性之间没有因果关系\[26,13,4\],并且轨迹长度与正在解决的问题实例的计算复杂度无关\[19\]。这些发现提出了一个问题:如果模型的推理与其解决方案没有因果关系,那么模型是否有必要显式生成这些令牌?这些中间令牌所表达的计算能否被内化到语言模型的参数中?最近一系列与自蒸馏相关的工作表明,可以使用知识蒸馏来训练语言模型内化额外的上下文\[12,21\]。从一个更大、能力更强的教师蒸馏到一个更小的学生这一总体思想本身已经确立\[9,1\];\[12\]和\[21\]更进一步,表明即使教师和学生是同一个模型,教师对额外上下文(专家演示或环境反馈)的访问也可以转移到学生的参数中。在LRM出现之前,有一些工作研究了思维链(CoT)令牌是否可以内化到模型参数中。\[24\]引入了上下文蒸馏,并展示了T5-small可以从其自身的CoT提示版本蒸馏成一个直接回答模型,从而解决算术问题。\[7\]通过隐式思维链知识蒸馏(ICoT-KD)推广了这一点,而\[8\]后来提出了逐步内化(ICoT-SI),它逐阶段地移除中间令牌,以便教师的推理逐渐被压缩到学生的隐藏状态中。这些方法的一个共同模式是,只要训练设置得当,分布内精度通常可以匹配,同时提高推理效率。然而,额外训练成本与推理节省之间的权衡尚未得到充分研究。更重要的是,由此产生的学生模型能否泛化也是一个未解决的问题。基于\[10,21,12\]的工作,我们使用知识蒸馏将中间令牌携带的信息内化到学生参数中。我们提出**掩码蒸馏**,这是一种知识蒸馏框架,其中学生模型被训练为仅根据问题预测解决方案令牌,而教师则在根据问题及其CoT轨迹进行条件化后,对学生的响应提供反馈。我们在两种设置中实例化这个框架:(i) **自蒸馏**设置,其中同一模型既充当思考模式下的教师,又充当非思考模式下的学生;(ii) **双模型**设置,其中推理模型监督一个独立的非思考模型,关注于解决方案令牌。通过将中间令牌视为推理模型用来拟合解决方案令牌的脚手架,我们进一步研究了一种**α-后缀掩码蒸馏**变体,其中学生内化教师中间轨迹的前$(1-\alpha)$部分(*前缀*),并被训练在推理时生成剩余的$\alpha$比例部分(*后缀*),连同样解决方案令牌。在本研究中,我们探讨以下研究问题:
- • **RQ1.** 通过掩码蒸馏,学生模型能否内化教师推理令牌中存在的信息,并以更高的推理效率实现相似的任务性能?
- • **RQ2.** 如果掩码蒸馏学生没有达到与教师相似的任务性能,提供额外的脚手架(通过在推理时发出的中间令牌数量来衡量)是否会提高任务性能,以及权衡推理成本的关系如何?
- • **RQ3.** 中间令牌的内化是否会导致分布外(OOD)任务性能下降,提供额外的脚手架是否有助于学生保持OOD性能?
- • **RQ4.** 训练目标重要吗?具体来说,我们能否为学生提供不同的脚手架,能否使用监督微调在教师轨迹上实现类似的结果?

我们在两个领域进行了广泛的实验:数学(GSM8K,以MATH-500和AIME-25作为分布外拆分)和Countdown(以目标范围偏移和搜索深度偏移作为分布外拆分),分别在**自蒸馏**和**双模型**设置下进行。通过扫描后缀脚手架参数$\alpha \in \{0, 0.3, 0.5, 0.7, 1.0\}$,我们的结果表明,完全内化中间令牌中信息的能力依赖于具体任务:它在基础学生有先验领域暴露的GSM8K上成功,但在Countdown上失败。然而,我们的结果表明,$\alpha$-后缀脚手架以适度的成本弥补了这一差距;例如,在Countdown上对自蒸馏应用$\alpha=0.3$,将准确率从41.7%(完全掩码)提高到86.2%,基本上匹配了教师(87.3%),同时推理令牌相对于非掩码变体减少了约1.3倍。在分布外任务中,在各种脚手架范式下训练的模型在目标范围偏移下能够干净地迁移,并且在自蒸馏设置下,它们在搜索深度偏移下比非掩码变体泛化得更好。这些发现将后缀脚手架确立为一个可控的轴线,沿该轴线可以调整准确率与推理成本的权衡,最佳操作点取决于任务和师生能力差距。

本文的其余部分组织如下。第2节提供了知识蒸馏技术的背景。第3节介绍了掩码蒸馏框架。第4节详细介绍了实验设置。第5节报告了GSM8K和Countdown上的结果,并讨论了它们对脚手架-泛化权衡的影响。

## 2 背景

### 2.1 知识蒸馏

传统上,知识蒸馏是一个框架,其中存在一个学生-教师对,学生模型通过最小化其输出分布之间的散度来训练以模仿教师的行为。该框架最早由\[10\]引入,用于将知识从集成或大型高度正则化模型迁移到更小、更简洁的模型中。在自回归模型的背景下,知识蒸馏已被广泛研究,用于训练一个较小的语言模型模仿较大教师LLM的输出。令 $\theta$ 表示学生模型的参数,$p_S^\theta$ 表示学生模型的策略,它关于 $\theta$ 可微。令输入-输出对的数据集为 $(X,Y)$。对于散度 $D$,我们将 $p_T$ 和 $p_S$ 在令牌级别分布上的差异定义为 D(p_T \| p_S^\theta)(y|x) := \frac{1}{L_y} \sum_{n=1}^{L_y} D(p_T(\cdot | y_{<n}, x) \| p_S^\theta(\cdot | y_{<n}, x)),其中 $L_y$ 是输出序列 $y$ 的长度。基于令牌级散度的蒸馏目标旨在通过优化 $\min_\theta \mathbb{E}_{(x,y)\sim\mathcal{D}} [D(p_T \| p_S^\theta)(y|x)]$ 来匹配分布。

### 3 掩码蒸馏

本节介绍了掩码蒸馏框架及其核心目标,并描述了我们将教师推理轨迹中的中间令牌内化到学生参数中的方法。

**教师轨迹分解。** 一个被提示使用固定格式(即,先在其 `<think>` 和 `</think>` 标签内输出推理轨迹,然后在 `<answer>` 和 `</answer>` 标签之间输出解决方案)的推理模型,其响应 $y$ 具有结构:$y = ITs \oplus STs$,其中 $ITs = y_{1:\ell}$ 是中间令牌,$STs = y_{\ell+1:L_y}$ 是解决方案令牌。我们定义推理模型的条件响应分布为 $\pi^T(\cdot | x, ITs)$,即给定问题 $x$ 和中间令牌 $ITs$ 的分布。教师轨迹由 $\pi^T$ 一次性生成,或者通过蒙特卡洛树搜索(MCTS)[20] 或束搜索[22] 进行搜索。

**推理时的非思考学生。** 在推理时,学生的目标是仅根据问题 $x$ 直接输出解决方案令牌 $STs$,而不生成任何中间令牌。因此,学生的分布定义为 $\pi^S_\theta(\cdot | x)$,并在推理时从该分布中采样解决方案令牌。

**掩码蒸馏。</b> 我们的掩码蒸馏框架如图1所示。在一个两阶段过程中:阶段1,数据收集(顶部):教师模型 $π^T$ 在思考模式下被提示,对问题 $x$ 给出条件。它生成一个完整响应 $y \sim π^T(\cdot | x)$,该响应被结构化为 ${ITs}</think><answer>{STs}</answer>$,从中我们提取中间令牌(IT)和解决方案令牌(ST)。
阶段2,训练(底部):非思考学生 $π^S_θ$ 仅以 $x$ 为条件,而冻结的教师以 $x$ 以及阶段1中 IT 的前 $(1-\alpha)$ 比例为条件。参数 $\alpha$ 控制在推理时学生被监督发出多少中间轨迹:学生被训练去重现教师中间轨迹的最后 $\alpha$ 比例部分以及解决方案令牌。学生的下一个令牌分布 $P$ 通过在线策略反向-KL损失 $\mathbb{E}_{y\sim P}[\log P/Q]$ 与教师的分布 $Q$ 进行匹配,该损失在学生采样的令牌上计算;梯度仅更新学生。当 $\alpha=0$ 时,该变体简化为完全掩码蒸馏;当 $\alpha=1$ 时,则简化为非掩码蒸馏。当教师和学生是同一个模型时,这是**自蒸馏**设置;当教师是更大的思考模型而学生是更小的非思考模型时,这是**双模型**设置。

我们提出了一个知识蒸馏框架,我们称之为**掩码蒸馏**,它将推理模型 $π^T$(教师模型)的中间令牌(IT)生成过程(即所谓的推理过程)内化到非思考学生模型 $π^S$ 中。目的是将教师的推理推入学生的参数中,这样在推理时,学生仅根据输入问题直接生成解决方案令牌,而不发出任何自己的中间令牌。这使得推理更加高效和廉价,因为中间计算已经在蒸馏过程中被内化到学生的内部表示中。具体来说,学生被训练成模仿教师在输入问题 $x$ 和教师自身生成的中间令牌上的条件分布,即 $π^T(\cdot | x, ITs)$。我们考虑两种**掩码蒸馏**设置:**自蒸馏**和**双模型**设置。在**自蒸馏**设置中,同一模型同时充当教师和学生:教师以思考模式运行并生成中间令牌,而学生以非思考模式运行并被训练直接生成最终解决方案。在**双模型**设置中,更大的推理模型作为教师,监督一个独立的、更小的非思考学生生成解决方案令牌。除非另有说明,我们使用“**掩码蒸馏**”一词泛泛地指代该框架,仅当必要时才明确区分两种设置。

图1显示了掩码蒸馏框架的概览。在**阶段1**中,我们通过从教师模型采样一个响应并提取 `<think>` 和 `</think>` 标签之间的片段作为 IT,为每个问题 $x \in \mathcal{D}$ 收集中间令牌;在**阶段2**中,我们通过最小化学生分布 $π^S_θ(\cdot | x)$ 与教师分布 $π^T(\cdot | x, ITs)$ 之间的散度来训练学生模型,其中教师以 $x$ 连同其在阶段1生成的 IT 为条件,而学生仅以 $x$ 为条件。掩码蒸馏的目标是学生和教师下一个令牌分布之间的反向-KL散度,该散度在学生采样的响应上计算。损失定义为
$\mathcal{L}_{\mathrm{MD}} = D_{KL}(\pi^S_\theta(\cdot | x, y_{<n}) \| \pi^T(\cdot | x, ITs, y_{<n}))$ 在从 $π^S$ 采样的令牌 $y_n$ 上。

**$\alpha$-后缀掩码蒸馏。** 为了探索从非掩码蒸馏($\alpha=1$)到完全掩码蒸馏($\alpha=0$)的频谱,并回答研究问题 **RQ2** 和 **RQ3**,我们引入了 $\alpha$-后缀掩码蒸馏。该变体在训练时控制学生在推理时被监督发出的中间令牌数量。令 $ITs = [\text{IT}_1, \text{IT}_2, \ldots, \text{IT}_L]$ 为教师生成的中间令牌的完整序列。我们将中间令牌序列分成两个部分:前缀 $ITs_{\text{pre}} = [\text{IT}_1, \ldots, \text{IT}_{\lfloor (1-\alpha) \cdot L \rfloor}]$ 和后缀 $ITs_{\text{suf}} = [\text{IT}_{\lfloor (1-\alpha) \cdot L \rfloor + 1}, \ldots, \text{IT}_L]$。在训练时,教师以前缀 IT 为条件:$π^T(\cdot | x, ITs_{\text{pre}})$,而学生被训练在推理时生成后缀 IT 以及解决方案令牌。因此,教师的条件分布变为 $π^T(\cdot | x, ITs_{\text{pre}})$,其中 $ITs_{\text{pre}}$ 是教师轨迹的前 $(1-\alpha)$ 比例部分。学生的生成过程变为 $y \sim π^S_θ(\cdot | x)$,预测后缀 IT 和解决方案令牌。损失现在变为
$\mathcal{L}_{\mathrm{$\alpha$-MD}} = D_{KL}(\pi^S_\theta(\cdot | x, y_{<n}) \| \pi^T(\cdot | x, ITs_{\text{pre}}, y_{<n}))$。$\alpha$ 参数作为一个旋钮,控制学生需要显式生成的中间轨迹比例,允许我们在完全内化($\alpha=0$,学生仅生成解决方案令牌)和完全显式推理($\alpha=1$,学生生成完整的中间轨迹)之间插值。我们的实验设置包括提取教师轨迹中的中间令牌,然后用不同的 $\alpha$ 值(0, 0.3, 0.5, 0.7, 1.0)训练学生。对于 $\alpha=0$,学生仅使用前缀(空)进行训练,使其能够完全内化推理过程。对于 $\alpha=1$,学生在完整轨迹上进行蒸馏。

**Inference with the non-thinking student.** At inference time, the student model $\pi_\theta^S$ directly rolls out its response $y \sim \pi_\theta^S(\cdot | x)$ without generating any intermediate tokens. The student generates the suffix ITs (for $\alpha > 0$) and the solution tokens. The inference cost is thus determined by the number of tokens generated, which scales with $\alpha$. For $\alpha=0$, the student generates only the solution tokens, leading to maximal inference efficiency.

**使用非思考学生的推理。** 在推理时,学生模型 $\pi_\theta^S$ 直接展开其响应 $y \sim \pi_\theta^S(\cdot | x)$,而不生成任何中间令牌。学生生成后缀IT(对于 $\alpha > 0$)和解决方案令牌。因此,推理成本由生成的令牌数量决定,该数量随 $\alpha$ 缩放。对于 $\alpha=0$,学生仅生成解决方案令牌,从而实现最大的推理效率。

相似文章

揭秘同策略蒸馏:其益处、危害及原因

Hugging Face Daily Papers

本文介绍了一种无需训练的框架,用于分析推理模型在逐token级别上的蒸馏信号。研究揭示,蒸馏引导在错误推理路径上更为有效,且其效果取决于学生模型的能力及任务上下文。

基于轨迹的在策略蒸馏用于掩码扩散语言模型

arXiv cs.CL

一篇论文提出了基于轨迹的在策略蒸馏(TOPD),一种教师监督框架,用于将推理能力迁移到掩码扩散语言模型,无需奖励估计,在显著的计算加速下实现了与经过RL训练的模型相当的准确率。

通过混合层蒸馏和关键信息的逐步注意力改进小模型的推理能力

arXiv cs.CL

本文提出一种新颖的思维链蒸馏框架,通过混合层模块的动态层对齐,将教师模型对关键信息的逐步注意力转移到学生模型中。该方法通过明确指导学生模型在推理过程中逐步聚焦关键信息,在数学和常识推理基准测试中实现了一致的性能提升。

基于理据引导的知识蒸馏跨语言立场检测

arXiv cs.CL

本文提出了一种基于理据引导的知识蒸馏框架,用于跨语言立场检测,利用大型语言模型的思维链提示,通过双路径蒸馏和对比学习训练一个紧凑的学生模型。