跳转至

Distilling What Matters: Confidence-Aware Selective Distillation for Large Language Models

会议: NeurIPS2026(任务包归属;缓存为 arXiv v1)
arXiv: 2609.36734
领域: 模型压缩
关键词: 知识蒸馏、置信度门控、选择性监督、认知不确定性、双向 KL 散度

一句话总结

CaRE-KD 按教师与学生的相对置信度逐 token 选择前向或反向 KL,并用 MC-dropout BALD 拒绝教师不确定而学生较确定的监督,在所测模型上改善蒸馏质量,达到 MBPP 62.4、GSM8K 73.9,但收益并非每个任务都成立,训练开销也明显增加。

研究背景与动机

生成式知识蒸馏通常把大模型的下一 token 分布当作小模型的学习目标。前向 KL 让学生覆盖教师给出的概率质量,适合传递丰富的候选答案,却也会让学生模仿教师不可靠的长尾。MiniLLM 采用反向 KL,Distillm 使用偏斜散度,ABKD 则通过参数化散度改变概率质量的分配;这些方法调整了模仿的几何形状,但通常在所有 token 上使用同一种偏好。

教师的平均能力更强,不代表在每个上下文中都比学生可靠。一个学生已经有较尖锐的分布时,教师可能对几个互相冲突的续写都分配较高概率;强制匹配会冲淡学生已有的判断。反过来,学生不确定而教师有清晰判断时,一味追求反向 KL 的模式选择又可能丢失有用的覆盖。论文把前一种失配称为 fidelity trap,不过“分布尖锐”并不自动意味着“判断正确”,这是理解方法时必须保留的边界。

本文不预先清洗整个训练集,而是在训练中分别回答两个问题:这个 token 应采用哪种蒸馏方向,以及这段监督是否值得接受。核心 idea:用相对熵置信度决定逐 token 的学习几何,再用教师与学生的相对认知不确定性决定是否更新,把“怎样模仿”和“何时信任教师”分开。

方法详解

整体框架

CaRE-KD 不改变学生的网络结构,而是改变训练目标与更新选择。教师和学生在相同前缀上输出下一 token 分布;“置信度门控散度”先生成 token 级损失,“Revival 认知拒绝”另用随机前向传播得到序列级可靠性信号,决定是否保留该监督。

在学生生成输出(SGO)训练中,续写来自学生,教师对这些相同的 on-policy token 上下文打分;非 SGO 则使用真实回答的 teacher-forced 上下文。门控与蒸馏必须共享同一个前缀,否则比较的是不同条件下的置信度,不能解释为当前 token 上谁更值得信任。

下面是训练依赖关系,而非部署时必须运行的模块。BALD 估计与 token 损失可以分开计算,图中的串接表示最终是否允许这份损失驱动更新;部署时只有训练后的学生,不需要教师、随机采样或拒绝规则。

%%{init: {'flowchart': {'rankSpacing': 24, 'nodeSpacing': 28, 'padding': 6, 'wrappingWidth': 400}}}%%
flowchart TD
    A["训练前缀<br/>教师与学生分布"] --> B["置信度门控散度"]
    B --> C["Revival 认知拒绝"]
    A -. "MC-dropout 与历史分位数" .-> C
    C -->|保留监督| D["更新学生参数"]
    C -->|拒绝监督| E["跳过相应更新"]
    D -. "训练完成" .-> F["推理:仅学生生成"]

关键设计

1. 置信度门控散度:按当前位置的相对置信度选择模仿方向

单次前向传播的置信度由归一化熵定义:概率越集中,置信度越高。它不需要额外的置信度预测器,也不是模型用自然语言报出的“把握”。教师置信度不低于学生时,硬门控采用前向 KL,让学生覆盖教师分布;教师更不确定时,采用反向 KL,减弱学生对教师长尾的追随。比较是相对的:即使教师在绝对尺度上相当确定,只要学生更确定,也可能进入反向分支。

令 \(p\) 为教师分布,\(q\) 为学生分布,\(H\) 为 Shannon 熵,\(\mathcal{V}\) 为词表,\(m\) 为软门控的边际偏移。正文式 (1)–(3) 的核心机制为:

\[ c(p)=1-\frac{H(p)}{\log|\mathcal{V}|},\qquad \mathcal{L}_{\mathrm{CARE}}=g\,\mathrm{KL}(p\|q)+(1-g)\,\mathrm{KL}(q\|p), \qquad g_{\mathrm{hard}}=\mathbb{I}[c(p)\geq c(q)],\qquad g_{\mathrm{soft}}=\sigma(c(p)-c(q)-m). \]

硬门控确实在两个方向间选一个;软门控则是连续加权,不能把每个 token 都描述为完全落入单一分支。边际只移动 sigmoid 的中心,不改变斜率,因而较小边际也不会把软门控变成硬门控。没有额外斜率参数时,有限的置信度差也不必使 sigmoid 接近两端。

前向 KL 惩罚遗漏教师概率质量,反向 KL 更关注学生正在放置概率的位置,所以后者能减少教师长尾对学生的影响。但反向 KL 仍然是在向教师分布靠近,并不等于保存原学生分布,更不能保证学生原来的高置信判断正确。真正停止这一监督的动作由第二个设计完成。

2. Revival 认知拒绝:教师随机预测不稳定而学生稳定时停止接受监督

单次熵高,可能只是问题有多个合理答案;它不能独立识别模型参数层面的不确定性。Revival 为教师与学生分别运行多次 MC dropout,先计算每个 token 的 BALD,再跨 token 平均成序列级分数。BALD 是“平均预测分布的熵”减去“各次预测熵的平均”:如果每次预测都很尖锐,但不同次选中不同答案,第一项仍高、第二项低,就能检测到随机模型实例之间的分歧。

令 \(\omega\) 表示随机 dropout 实例,\(\mathcal{I}_{T}\) 与 \(\mathcal{I}_{S}\) 为教师和学生的序列级 BALD,\(Q_{\tau}\) 为各自历史分数的运行分位数。正文式 (4)–(5) 写为:

\[ \mathcal{I}_{\mathrm{BALD}}(\mathbf{x})= H\!\left[\mathbb{E}_{\omega}p(y\mid\mathbf{x},\omega)\right] -\mathbb{E}_{\omega}\!\left[H(p(y\mid\mathbf{x},\omega))\right], \qquad M_{\mathrm{Revival}}=\mathbb{I}\!\left[ \mathcal{I}_{T}>Q_{\tau}(\mathcal{I}_{T})\ \land\ \mathcal{I}_{S}<Q_{\tau}(\mathcal{I}_{S})\right]. \]

这不是直接比较教师 BALD 是否大于学生 BALD,而是让两者分别与自己的历史分布比较。教师落入自身较不确定的区域,且学生落入自身较确定的区域,才触发拒绝;双方都不确定并不自动拒绝。这一设计避免仅因为样本困难就删掉所有监督,也减轻教师与学生 BALD 绝对尺度不一致的问题。

附录 §8.1 明确说明,即使原模型部署配置关闭 dropout,估计时仍强制开启 dropout 层,把 attention dropout 设为 0.1,计算后恢复配置。默认进行 3 次随机前向传播。因此它是一个有额外计算成本的可靠性估计,不是单次熵门控的免费副产品;教师参数在蒸馏中保持固定。

粒度需要谨慎阅读:§2.2 和附录 §8.1 定义的是序列分数,但正文式 (6) 把拒绝掩码放在整个 batch 求和之外,讨论部分又称“跳过 batch 的反向传播”。缓存没有明确给出如何把多个序列掩码聚合成一个 batch 决策。因此可以确认其意图是拒绝相应监督、避免相关更新,却不能自行把式 (6) 修成某个逐样本加权平均并声称是作者实现。

论文用逐渐增强的拒绝调度,让早期广泛学习、后期更少接受教师不可靠的区域。原文同时把 \(\tau\) 用作运行分位数与“目标跳过比例”,但二者不是一般情况下天然相等的数量:联合条件的触发率取决于两个分数的分布及相关性。报告的 50–70% 优势区间应视为作者实验中的调度设置,不应由式 (5) 直接推导为精确的实际拒绝比例。

损失函数 / 训练策略

在未拒绝监督时,优化门控后的蒸馏目标;拒绝时跳过相应更新。硬门控是 detached 阶跃函数,其门控梯度为零,只保留逐 token 的分支选择。软门控若未从计算图中分离,除了两种 KL 的学习梯度,还多出一个门控选择项。正文命题 3.1 与附录 §9.1 给出:

\[ \nabla_{\theta}\mathcal{L}_{\mathrm{CARE}} =g\nabla_{\theta}\mathcal{L}_{\mathrm{FKL}} +(1-g)\nabla_{\theta}\mathcal{L}_{\mathrm{RKL}} +(\mathcal{L}_{\mathrm{FKL}}-\mathcal{L}_{\mathrm{RKL}})\nabla_{\theta}g, \qquad \nabla_{\theta}g=-g(1-g)\nabla_{\theta}c(q_{\theta}). \]

这里 \(\theta\) 是学生参数。在前向 KL 大于反向 KL 的情形,额外项结合梯度下降可推动学生置信度增加;不能把这种作用说成硬门控也具备的“置信度梯度”。实验的最高峰值反而来自硬门控,软门控更偏向跨配置的稳定性。因此理论解释与最优实验配置必须分别陈述。

命题 3.2 的差异在于不同 token 可以使用不同分支,一个全局固定的偏斜参数通常无法复现这种逐位置行为;这并不证明静态散度在每个任务都更差。定理 3.5 更有明确前提:优化始终停留在固定分支,并且学生收敛到教师,学生熵才向教师熵收敛。它不保证切换分支时全局收敛、熵单调变化,或相对于真实答案的概率校准。

熵替代 BALD 的定理也有两个前提:每个随机前向分布足够尖锐,以及单次分布近似随机预测的均值。多种合理续写、很高的偶然不确定性或明显随机预测差异都可能削弱这些近似。全部报告结果使用完整 MC-dropout BALD,不能把便宜的单次熵拒绝当作已经验证等效的替换方案。

附录 §10 的训练配置为 LoRA rank 16、学习率 \(5\times10^{-5}\)。Dolly 的小于 1B 学生 batch size 32、最多 20 epochs;大于 1B 学生 batch size 8、10 epochs,按验证 ROUGE-L 选 checkpoint。UltraChat、WizardCoder、MetaMathQA 分别训练 3、2、2 epochs;BALD 默认 3 次随机前向、温度 3.0。指令评测使用温度 0.8、top-p 0.95、最多 512 tokens,代码与数学使用贪心解码、最多 1024 tokens,实验在单张 A100 80GB 上进行。

实验关键数据

主实验

实验覆盖 8 个教师—学生对、11 个评测集。下表摘取正文表 1 的 SGO 指令任务平均值,单元格为“ROUGE-L / LLM 评审事实性”,后者是 GPT-5-Mini 对参考答案一致性的 0–100 评分;结果按原文平均 5 个种子。这里的 SRKL 是 Skewed Reverse KL。

教师 → 学生 FKL SRKL CaRE-Div
GPT2-XL → GPT2-base 17.0 / 16.7 19.3 / 19.5 19.6 / 21.0
GPT2-XL → GPT2-large 16.1 / 15.7 17.9 / 16.9 18.1 / 19.1
OPT-2.7B → OPT-125M 16.1 / 15.5 19.0 / 20.2 19.3 / 18.5
Gemma-2-9B-IT → Gemma-2-2B-IT 16.3 / 18.2 16.5 / 20.6 17.1 / 21.6
OpenLLaMA2-7B → OpenLLaMA2-3B 22.3 / 27.9 26.4 / 35.3 26.2 / 33.8

OPT 的 ROUGE-L 更高,但事实性低于 SRKL;OpenLLaMA 的两个平均指标均未胜过 SRKL。表 1 标题称 3 个模型对的平均 ROUGE-L 最优,后续正文却称 4 个;逐行数字支持后者,本笔记保留数字并指出文字冲突,不据宣传措辞宣称全面领先。

正文表 3 的领域结果如下。代码指标是执行测试的 pass@1,数学指标按 §4 与 §10.2 为 exact-match accuracy;表 3 标题将所有单元格统称 ROUGE-L / LLM、数学列头又标 pass@1,均与指标说明不一致。

数据集 指标 原学生 Distillm2 ABKD GKD Distillm CaRE-KD
HumanEval pass@1 32.3 43.3 42.7 43.1 42.9 43.3
MBPP pass@1 58.5 60.3 59.8 60.2 60.0 62.4
GSM8K exact-match accuracy 69.9 72.2 71.0 71.4 71.7 73.9
CollegeMath exact-match accuracy 37.1 44.2 44.3 43.9 44.1 46.1

MBPP 相对表内最佳基线提升 2.1 个百分点,GSM8K 提升 1.7,CollegeMath 提升 1.8;HumanEval 只是并列最佳。原文正文对 CollegeMath 的对比标签并不完全一致,因此以上提升按表 3 数字计算,而不是混用不同基线名称。

消融实验

正文表 2 单独报告加入 Revival 后的 ROUGE-L 变化,即“有 Revival − 无 Revival”。下表只保留跨四个指令集的平均变化,以及必要的反例,不把消融差值当作绝对成绩。

学生 损失 平均变化 单任务反例 / 说明
GPT2-base CaRE +0.90 四个任务均为正
GPT2-large CaRE +0.46 Dolly −0.04
OPT-125M CaRE +1.42 Super-Natural Instructions +3.62
Gemma-2-2B-IT CaRE +2.90 Vicuna −0.40
OpenLLaMA2-3B CaRE +0.73 四个任务均为正
Gemma-2-2B-IT FKL −3.91 过滤并非通用增益
GPT2-large RKL −0.28 Self-Instruct −1.78

CaRE 的平均变化均为正,但仍有两个负单元格;静态损失搭配 Revival 甚至可能平均退步。原文对 CaRE 的差值给出单样本检验统计量 3.98、显著性小于 0.001;这支持其所测配置的平均互补性,而不是“任何损失加过滤都更好”。图 3、8 的数值曲线未包含在文本缓存中,因此这里只引用正文对硬/软门控、调度和采样数的趋势,不补造逐点消融成绩。

关键发现

  • 高压缩配置收益较清楚:GPT2-base 相对 FKL 的平均 ROUGE-L 增益为 2.6,OPT-125M 为 3.2;相对更强 SRKL 则都只有 0.3。
  • 事实性证据是终端输出改善,不是已经证明 BALD 能识别幻觉。§4.2 的拒绝区教师事实性为 25.2,对照为 28.6,但 Wilcoxon 显著性为 0.06,另一项区域检验为 0.26,均未显著。
  • 训练成本不能称为普遍温和:§5 在默认 3 次采样下报告指令任务约增加 42%,chat 106%、代码 170%、数学 148%。Dolly 单 epoch 从 5520 秒增至 7860 秒;这不增加部署推理模块,却显著增加离线训练预算。

亮点与洞察

  • 相对置信度比固定的“教师一定正确”假设更细粒度。可迁移的不是简单奖励尖锐分布,而是让监督方向随具体上下文改变,并保留明确的拒绝出口。
  • 门控与拒绝承担不同职责:前者仍向教师学习,只改变覆盖与模式选择;后者才让某些监督停止影响学生。两层选择的互补性比再调一个全局散度参数更有解释力。
  • 历史分位数比较不要求两个模型的 BALD 数值同尺度。实际复用时仍须记录分位数窗口、真实触发率及保留样本分布,不能只公布“目标 skip”参数。

局限与展望

  • 置信度不是正确率,低 BALD 也可能来自稳定地犯错。需要独立正确性验证或更可信的校准实验,才能判断被保护的学生先验究竟是否正确。
  • 附录表 9 的默认 CaRE-KD ECE 为 0.34,FKL 为 0.22;软化配置为 0.28,并非正文所说的“相同 ECE”。单参考 token 匹配确有局限,但现有数字不支持“真实标签校准最佳”。
  • 序列级定义与 batch 级式 (6) 不一致,分位数参数与目标跳过比例的对应也不充分明确。复现应先核对实现中的掩码聚合、padding 处理与实际拒绝调度,不能由本笔记替作者补齐。
  • MC dropout 被强制打开后的分歧非零,不代表该分歧已校准为可靠的认知不确定性。附录表 4 仅验证 GPT2-XL 和 OPT-2.7B 两个教师,不能外推到所有教师模型。
  • 单次熵替代 BALD 只有条件性理论说明,所有实验仍使用完整 BALD。需要在相同计算预算下比较该替代、增加训练步数与完整 CaRE-KD,才能确定额外前向传播是否值得。
  • 附录表 10 的难度与多样性差异不显著,只能说明该 OpenLLaMA-Dolly 检查没有检出差异,不等于证明所有任务无选择偏差;chat 的 GKD、Distillm 对比也仍缺失。

相关工作与启发

  • vs MiniLLM:MiniLLM 用反向 KL 抑制教师长尾;CaRE-KD 在 token 层面选择前向或反向几何,并在可靠性不足时拒绝更新。代价是更复杂的控制逻辑与随机前向计算。
  • vs Distillm / ABKD:偏斜 KL 与参数化散度使用全局几何参数;本文使用随上下文变化的门控。OpenLLaMA 的平均结果仍略低于 Distillm,说明动态选择不是对静态方法的无条件替代。
  • vs GKD:on-policy 数据解决训练与学生生成上下文失配,本文解决给定上下文中如何信任教师。两者可以结合,SGO 下门控与教师打分必须对齐学生的同一生成轨迹。
  • 研究启发:可测试“学生自信但错误”区域中的拒绝副作用,并用代码单元测试或数学答案验证作为额外信任信号;同时按计算预算而非 epoch 数匹配基线。这个方向是笔记提出的实验建议,不是论文已经验证的结论。

评分

  • 新颖性: 4/5;逐 token 几何选择与序列可靠性拒绝形成清晰组合。
  • 实验充分度: 3/5;模型与任务覆盖较广,但等预算比较、拒绝机制事实性验证及部分基线仍不足。
  • 写作质量: 3/5;机制解释完整,但掩码粒度、指标标签和若干正文数字解读存在冲突。
  • 价值: 4/5;适合离线高压缩蒸馏,但需同时评估教师可靠性与额外训练成本。