之前帮团队做小模型落地方案时,我们一直用“教师模型生成的答案”来蒸馏,训练出的模型在简单任务上表现尚可,但一到数学推理、多跳问答、代码调试这类复杂场景,效果就明显缩水。后来把研究方向从“教师给了什么答案”转向“教师怎么得到答案”,问题才慢慢打开。这篇文章想围绕一个核心观点展开:大模型知识蒸馏研究中,教师模型的推理习惯,往往比教师给出的分数更重要。
本文适合正在做模型压缩、小模型训练、大模型数据生产的算法工程师,也适合想弄懂知识蒸馏原理的初学者。读完你会理解传统蒸馏为什么在大模型场景下不够用、推理习惯到底指什么、如何把推理过程蒸馏给学生模型,并拿到一套可落地的实验流程与代码骨架。
1. 知识蒸馏的背景:从模型压缩到能力复制
1.1 知识蒸馏解决什么问题
知识蒸馏的核心思想很简单:用一个能力更强的大模型作为教师,去指导一个小模型学习,让这个小模型在参数量小得多的前提下,尽量逼近教师的输出效果。
传统上,知识蒸馏被看作一种模型压缩手段。比如在图像分类任务中,一个 ResNet-50 可能不如一个大模型精度高,但通过蒸馏,可以让 ResNet-18 学到 ResNet-50 的“判断倾向”,获得比单独训练更高的精度。这里的“知识”被定义为教师模型输出层上的概率分布。
但随着大语言模型兴起,蒸馏的定位发生了变化。大模型可以处理复杂指令、多步推理、代码生成等任务,但其推理成本高、部署门槛高,很多业务场景根本承载不起。于是团队开始用大模型生成数据,再训练一个小模型来承接这些能力。这种模式已经不只是“压缩”,而是一种“能力复制”。
1.2 大模型时代的蒸馏动机
在大模型知识蒸馏的实践中,最常见的做法是:调用教师模型,给一批问题生成答案,再用这些“问题—答案”对去微调学生模型。这个流程看似直接有效,却隐藏着一个问题——教师模型在回答过程中经历了大量中间推理,而“问题—答案”对把这些中间过程全部省略了。
当任务简单时,省略中间过程问题不大。学生模型可以直接记住问题到答案的映射。但当任务复杂、需要多步推理时,这种跳跃式学习就很难奏效。小模型本身容量有限,它需要一个逐步拆解问题的“脚手架”,而不是直接面对一个高维映射。
这正是“教师模型推理习惯比分数更重要”这一观点的现实背景。教师模型最宝贵的产出,往往不是最终答案,而是它为了得到答案所走的那条推理路径。
2. 传统知识蒸馏:让模型学会教师的“分数”
2.1 Hinton 式蒸馏的基本原理
理解推理习惯为什么重要,需要先回顾传统蒸馏是怎么做的。Hinton 等人在 2015 年提出知识蒸馏时,核心思路是让学生模型去匹配教师模型的软标签。
普通分类任务中,模型输出的 logits 经过 softmax 后变成一个概率分布,比如一张图片有 70% 的概率是猫、20% 的概率是狗、10% 的概率是鸟。传统训练只关心最终正确类别,而蒸馏会额外让学生模型去模仿教师模型的完整概率分布。
为了让分布中的“暗知识”更明显,蒸馏引入了温度参数 T:
p_i = exp(z_i / T) / sum_j exp(z_j / T)T 越大,概率分布越平滑,类别间的相对关系越突出。训练时,学生模型一方面计算与真实标签的交叉熵,另一方面计算与教师软标签的 KL 散度。这样,学生不仅能学会“正确答案是猫”,还能学会“在教师眼里,猫和狗比猫和鸟更接近”。
2.2 只看分数的局限
这种基于软标签的蒸馏在分类任务中效果很好,但它有一个隐含假设:知识可以被压缩到输出层的概率分布中。这个假设在多步推理任务中并不成立。
首先,软标签只编码了教师对最终答案的置信度。教师模型在推理过程中可能多次调整思路,可能经历了“尝试错误—发现矛盾—重新计算”的过程,这些信息在最终概率分布中几乎没有体现。学生模型看到的只是一个终点,而不是完整路线。
其次,复杂任务的输出空间非常大。比如数学题的答案是数字,但得到这个数字的过程可能有十几种不同路径。教师模型选择的路径、采用的中间公式、对计算结果的校验方式,才是真正有价值的知识。如果只给分数,学生模型就必须自己重新发明这些推理策略,这对小模型来说负担太重。
最后,错误信息也被过滤了。教师模型在推理中可能发现某个中间结果不合理,从而回退重算,这种能力在“分数”中完全不可见。学生模型如果只学最终答案,遇到类似中间错误时,不知道该如何自我纠正。
3. 核心观点:教师模型的推理习惯比分数更重要
3.1 什么是推理习惯
本文所说的“推理习惯”,不是指教师模型的某一个输出,而是指教师在处理任务时的整体行为模式。它至少包含几个层面:
- 文本层面的思维链:教师生成问题解析、分步计算、逻辑判断等中间文本。
- 中间层特征表示:教师模型每一层 Transformer 对输入信息的编码方式。
- 注意力分布:模型在每一步关注了哪些历史信息或上下文。
- 错误纠正模式:模型发现中间结果不合理后如何调整策略。
在实际蒸馏中,最容易利用的是第一类,也就是思维链文本。因为文本是显式、可读、可直接作为训练语料的。中间层特征表示则更适合结构相近的模型之间对齐。注意力分布对齐实现难度更高,但在可解释性研究中有不少探索。
3.2 分数与推理习惯的信息量差异
那么,为什么推理习惯比分数更重要?最直接的原因是信息量差异。
假设教师模型处理一道数学应用题。最终答案只是一个 token 序列或一个数值,能提供给学生的监督信息非常有限。而教师的推理轨迹可能包含 200 到 500 个 token,其中有问题抽象、条件拆解、计算步骤、结果校验。这些中间 token 把一个大问题分解成了若干个小问题,每一个小问题都成为学生模型的学习目标。
从监督信号的角度看,传统蒸馏只在最终输出上提供反馈,属于稀疏监督;而基于推理轨迹的蒸馏在每一步都提供反馈,属于密集监督。密集监督显著降低了学生模型的学习难度,因为它不需要一次性学会复杂映射,只需要学会每一步的小映射。
更关键的是,推理轨迹让“错误模式”变得可控。教师如果每一步都输出中间结果,学生训练时就能看到教师如何在某一步修正偏差。这种能力很难用最终分数传递。
3.3 复杂任务上的表现差异
在简单的分类、短文本匹配任务上,软标签蒸馏和推理轨迹蒸馏的差距可能不明显。但在需要多步推理的任务上,比如数学应用题、多跳问答、逻辑推理、代码生成,两者的差距会迅速拉大。
原因在于复杂任务的中间状态非常多。学生模型不仅要学会“输入到输出”的映射,还要学会“如何在中间状态之间转移”。如果缺少中间状态的监督,学生模型容易学到表面相关性。比如它可能记住了某些题型的关键词,却无法真正理解推导逻辑,一旦题目换一种表达方式,准确率就会明显下降。
反过来,如果学生模型学习了教师的推理过程,它更像是学会了“解题方法”而不是“背答案”。面对变体问题时,学生可以按照学到的推理框架重新推导,鲁棒性会好很多。这也是当前很多大模型知识蒸馏工作开始关注 CoT(Chain-of-Thought,思维链)蒸馏的原因。
4. 让推理习惯参与蒸馏:主流方法
4.1 思维链文本蒸馏
思维链蒸馏是目前最直观、落地成本最低的一种方式。流程大致是:用教师模型对每个问题生成一段包含逐步推理的答案,然后把“问题 + 推理过程 + 最终答案”作为训练样本,用标准的语言建模目标训练学生模型。
这种方法的优点是数据格式简单,不需要修改模型结构。学生模型只需要拥有生成文本的能力,就能学习教师的推理过程。对于 7B、3B 甚至更小的模型,只要训练数据里的推理过程足够清晰,学生通常都能体现出明显的推理能力提升。
代表性研究思路包括 Distilling Step-by-Step 等。这类工作的共同点是把教师的推理过程作为额外的监督信号,而不是仅仅把最终答案当作标签。需要说明的是,不同实现细节差别很大,实际效果需要结合具体任务验证。
4.2 中间层特征对齐
思维链文本是显式知识,但教师模型内部还有大量隐式知识,分布在每一层的 hidden state 中。如果学生模型和教师模型结构相近,或者使用了相同的 tokenizer,可以考虑做中间层特征对齐。
具体做法是:让教师和学生处理同样的输入,然后取出某一层或若干层的输出向量,通过一个投影层将学生向量映射到教师的向量空间,再计算 MSE 或余弦相似度损失。
这种对齐方式的优势是知识传递更完整,学生不仅知道教师“说了什么”,还知道教师“在想什么”。但它的工程成本更高。首先是层与层之间的对应关系需要设计;其次是学生模型的维度通常比教师小,需要额外引入投影层;最后是如果教师是 API 服务,根本无法拿到中间层输出,只能放弃这种方式。
4.3 多路径采样与最佳路径筛选
推理习惯并不是越多样越好。教师模型也可能生成错误推理、重复推理或者幻觉推理。因此在蒸馏之前,需要做推理路径的质量控制。
一个比较实用的方案是 best-of-N 采样。对于同一个问题,让教师模型用稍高的温度采样 N 条推理路径,然后按照规则筛选出质量最高的一条。筛选规则可以包括:
- 最终答案是否正确;
- 推理过程中是否包含提取出的关键步骤;
- 是否存在重复片段;
- 推理长度是否合理。
这种“先生成,后筛选”的方式,本质上是在构建一份高质量的推理轨迹数据集。从实际经验看,数据质量对蒸馏效果的影响往往大于数据数量。一万条经过筛选的高质量推理轨迹,效果通常会好于十万条未经筛选的原始生成结果。
5. 完整实验流程与代码示例
5.1 整体流程
下面以一个常见的场景为例:用一个大语言模型作为教师,蒸馏出一个参数量较小的学生模型,让它在数学推理任务上具备类似教师的推理能力。整体流程分七步:
- 准备评测数据集和训练问题集。
- 让教师模型对训练问题生成推理轨迹。
- 过滤低质量推理轨迹。
- 构造“问题—推理过程—答案”训练样本。
- 加载学生模型,准备训练环境。
- 组合蒸馏损失进行训练。
- 在评测集上验证学生模型的推理能力。
5.2 生成与筛选推理轨迹
首先让教师模型生成推理轨迹。下面的代码是核心片段,以 Hugging Face transformers 为例,实际使用时需要根据教师模型类型调整 prompt 模板。
from transformers import AutoModelForCausalLM, AutoTokenizer model_name = "teacher-model-path" tokenizer = AutoTokenizer.from_pretrained(model_name) model = AutoModelForCausalLM.from_pretrained(model_name, device_map="auto") def build_prompt(question: str) -> str: return f"请一步步推理并回答问题:\n{question}\n\n推理过程:" def generate_cot(question: str, max_new_tokens: int = 512) -> str: prompt = build_prompt(question) inputs = tokenizer(prompt, return_tensors="pt").to(model.device) outputs = model.generate( **inputs, max_new_tokens=max_new_tokens, do_sample=True, temperature=0.7, top_p=0.9, num_return_sequences=1, ) input_len = inputs["input_ids"].shape[1] generated = outputs[0][input_len:] return tokenizer.decode(generated, skip_special_tokens=True)生成之后,需要对推理轨迹做一次筛选。下面是一个通用的过滤函数,重点检查答案正确性、推理长度和重复度。
def extract_answer(prediction: str) -> str: # 简单示意:取最后一个等号后面的内容作为答案 if "=" in prediction: return prediction.strip().split("=")[-1].strip() return prediction.strip() def has_repetition(text: str, threshold: int = 5) -> bool: # 检测是否出现大量连续重复片段 words = text.split() for i in range(len(words) - threshold): if len(set(words[i:i + threshold])) == 1: return True return False def filter_cot_samples(samples: list[dict]) -> list[dict]: result = [] for item in samples: answer = extract_answer(item["prediction"]) if answer != item["ground_truth"]: continue if len(item["prediction"]) < 20: continue if has_repetition(item["prediction"]): continue result.append(item) return result需要说明的是,实际项目中的答案提取不能只依赖等号,最好的做法是在生成前要求教师按固定格式输出,例如最后一行写“答案是:XXX”,然后用更稳定的规则提取。
5.3 构造训练样本
过滤完成后,把数据整理成统一格式。一个典型的训练样本如下:
{ "question": "一个长方形长 8 厘米,宽 5 厘米,求周长。", "reasoning": "长方形的周长等于两倍的长加宽。\n长加宽等于 8 + 5 = 13 厘米。\n两倍为 13 * 2 = 26 厘米。", "answer": "26 厘米" }训练时,把 question、reasoning、answer 拼接成一个完整的文本序列,作为学生模型的 target。这个拼接过程可以在数据预处理阶段完成,也可以在训练循环中动态完成。
5.4 核心训练代码
学生模型的训练目标是两部分的组合:一部分是传统的 logits 蒸馏损失,另一部分是思维链序列的语言建模损失。下面的代码是核心思路,实际运行时需要根据模型和框架调整。
import torch import torch.nn.functional as F def kd_loss(student_logits, teacher_logits, temperature=4.0): """ 软化 logits 后的 KL 散度损失。 temperature 越大,分布越平滑。 """ student_log_probs = F.log_softmax(student_logits / temperature, dim=-1) teacher_probs = F.softmax(teacher_logits / temperature, dim=-1) loss = F.kl_div(student_log_probs, teacher_probs, reduction="batchmean") return loss * (temperature ** 2) def cot_seq_loss(student_logits, target_ids): """ 思维链文本的交叉熵损失。 target_ids 中需要 mask 的位置可以设为 -100。 """ vocab_size = student_logits.size(-1) return F.cross_entropy( student_logits.view(-1, vocab_size), target_ids.view(-1), ignore_index=-100, ) # 训练循环中的关键计算(示意) # alpha 和 beta 是损失权重,需要根据实验调整 total_loss = alpha * kd_loss(student_logits, teacher_logits) \ + beta * cot_seq_loss(student_logits, target_ids)在实现中有几个细节需要重点注意:
- teacher_logits 如果很大,建议提前缓存到磁盘,避免每次训练都重复前向推理。
- 如果学生模型和教师模型的词表不一致,logits 蒸馏无法直接使用,此时可以只用思维链文本损失。
- ignore_index 要确保 prompt 部分的 token 不会被计算损失,学生只需要学习推理过程和答案部分。
5.5 验证与指标
训练完成后,不能只看最终答案准确率,还需要评估推理质量。建议同时关注以下几类指标:
- 答案准确率:学生模型生成结果中,最终答案正确的比例。
- 推理过程完整度:生成结果中是否包含关键推理步骤。
- 格式符合率:是否按照训练时的格式输出。
- 人类抽样评估:随机抽 50 到 100 条,人工判断推理逻辑是否成立。
评估时建议使用与训练时不同的 prompt 模板,避免学生模型只是记住了模板格式。比如训练时用“请一步步推理”,评测时改成“请解决以下问题”,观察推理能力是否真正迁移。
6. 常见问题与排查思路
| 问题现象 | 常见原因 | 解决思路 |
|---|---|---|
| 学生模型推理时频繁重复 | 教师生成的推理轨迹存在重复,或训练数据覆盖不足 | 加强数据过滤,增加高质量数据,调整解码参数 |
| 蒸馏后答案准确率不升反降 | 只学了答案,没学推理过程,或教师推理质量差 | 切换为 CoT 蒸馏,先筛选推理轨迹 |
| 学生模型输出格式混乱 | 训练数据格式不统一 | 统一教师生成格式,预处理时做格式规整 |
| 训练损失下降但评测效果差 | 过拟合训练数据,缺少多样化问题 | 增加数据多样性,加入正则化或早停 |
| 教师模型是 API,无法拿到 logits | 接口不开放中间层和 logits | 只用推理文本作为训练目标,不计算 KD loss |
| 学生模型容量太小,学不会 | 学生参数量与任务复杂度不匹配 | 适当增大模型容量,或把复杂任务拆分成子任务蒸馏 |
下面展开两个最常见的问题。
第一个问题是“推理轨迹本身质量差”。很多情况下,教师模型生成的推理过程看起来通顺,但中间步骤有隐藏错误,只是最后答案碰巧对了。这种数据进入训练集后,会让学生学会错误的推理方式。解决方法是在筛选阶段增加规则校验,比如让推理过程中必须出现某些关键公式或关键实体,或者对答案正确但推理质量存疑的样本进行过滤。
第二个问题是“学生模型学到了推理格式但没有学到推理能力”。这种情况通常表现为学生可以按照“第一步、第二步”的格式输出,但每一步的内容逻辑不连贯。根因是推理轨迹里的中间步骤缺乏足够约束,学生只是记住了模板。缓解办法是增加思维链数据的多样性,同时在训练损失中加大对推理过程 token 的权重。
7. 最佳实践与工程建议
7.1 数据质量优先于数据量
大模型知识蒸馏中,数据质量的重要性被反复验证。建议从几千条精心筛选的推理轨迹开始,而不是一开始就追求十万条数据。先验证小规模数据上学生模型是否具备推理能力,再逐步扩大数据规模。这样可以减少无效训练成本,也更容易定位问题。
7.2 控制教师模型的采样参数
教师模型生成推理轨迹时,建议使用 0.7 左右的温度并开启 top-p 采样。温度过低会导致生成内容过于保守,推理路径单一;温度过高则容易引入幻觉。best-of-N 采样时,N 一般取 4 到 8 比较合适,避免采样过多带来的成本压力。
7.3 分开考虑“答案蒸馏”和“过程蒸馏”
对简单任务,答案蒸馏成本低、收益明确。对复杂任务,建议优先加入过程蒸馏。如果计算资源有限,可以在同一个 batch 中混合两种样本:一部分样本只提供标准答案,一部分样本提供完整推理轨迹,然后通过损失权重控制两类样本的贡献比例。
7.4 缓存教师模型的推理结果
教师模型的一次推理成本远高于学生模型训练的一个 step。建议把教师模型的生成结果按问题 ID 缓存到本地,格式可以是 JSONL。这样多次实验不需要重复调用教师模型,能节省大量成本。
7.5 评估体系要跟上
只靠答案准确率评估蒸馏效果,很容易高估或低估模型能力。尤其是推理任务,可能出现答案正确但推理错误的情况。建议建立多层次评估体系:自动指标负责批量筛选,人工评估负责最终把关。
7.6 参数高效蒸馏降低迭代成本
如果学生模型本身也是亿级参数模型,建议使用 LoRA 等参数高效微调方法,先训练 adapter,再决定是否合并回主模型。这样可以在同一份推理轨迹数据上快速尝试不同的损失权重和数据组合,提升实验迭代效率。
8. 结语与下一步学习方向
本文从大模型知识蒸馏的实际问题出发,介绍了传统蒸馏的原理与局限,重点解释了为什么教师模型的推理习惯比最终分数更重要,并给出了一套基于思维链蒸馏的实验流程。核心收获可以归纳为三点:第一,推理轨迹本质上是一种密集监督信号,能显著降低学生模型的学习难度;第二,数据质量必须放在首位,教师模型生成的数据也要经过筛选;第三,评估不能只看答案准确率,还要关注推理过程的有效性。
如果你准备在自己的项目中落地蒸馏,我建议不要一上来就追求大规模数据。先找几百道有代表性的复杂问题,让教师模型生成推理轨迹,人工检查其中 20 到 30 条,感受一下数据质量,再训练一个小模型观察效果。这一步虽然简单,却能帮你少走很多弯路。
下一步可以继续研究中间层特征对齐、多教师蒸馏、以及推理轨迹的自动化质量评估。这些方向都能和本文介绍的思维链蒸馏结合起来,帮助你把大模型的能力更完整地迁移到小模型上。如果这篇文章对你有帮助,欢迎收藏备用,也欢迎在实际实验后回来交流你的蒸馏效果。