模型蒸馏最近在技术社区里讨论得非常多。很多人把它当成一种神奇的提效手段,觉得只要把大模型的输出拿回来蒸一遍,小模型就能立刻逼近前沿水平。实际上,蒸馏是深度学习中一套成熟的迁移学习方法,核心思路很直接:用一个更强的教师模型,把知识迁移到一个更小的学生模型上,让学生模型在参数更少、推理更便宜的情况下,尽量接近教师模型的能力。硅谷技术社区讨论开放模型逼近前沿时,蒸馏几乎绕不开,因为过去一段时间里,不少开放权重模型能够快速追上来,靠的并不是从零复现一个超大模型的预训练,而是把已有前沿模型的高质量输出高效转化成训练资源。这篇文章不谈观点争议,只从原理、实验、边界和排查几个角度,把蒸馏这件事讲清楚。
1. 先弄清“蒸馏”到底在蒸什么
1.1 教师模型与学生模型:一个类比先讲清楚
蒸馏不是新概念。2015 年 Hinton 等人发表的《Distilling the Knowledge in a Neural Network》把它带到主流视野,但背后的直觉更早就有。你可以把教师模型想象成一个经验丰富的专家,学生模型是一个刚入行的新人。专家给新人讲题时,不会只告诉他对错,还会解释为什么选 A、B 哪里像但不对、C 的表述有什么隐患。新人听到的不只是标准答案,还有判断过程中的边缘信息,所以学得比光看正确答案更快更稳。
落到模型上,教师模型输出的是每个类别的概率分布。比如一张图片,教师模型判断结果是“猫”的概率 0.8、“狗”0.15、“狐狸”0.05。如果只用硬标签训练,学生只知道“这张图是猫”;如果使用教师输出的概率分布,学生还会知道“猫和狗在视觉特征上比较接近,和狐狸差异更大”。这类软信息,就是蒸馏要迁移的核心。
这也是“知识蒸馏”和普通监督学习最明显的区别。普通训练把正确类别当成唯一标准,蒸馏则把教师模型的内部判断习惯也带给了学生。学生不只是在背答案,更像是在模仿教师做判断的方式。
1.2 温度、软标签和损失函数
为了让软标签里的边缘信息更明显,蒸馏通常会给 softmax 加一个温度参数。标准 softmax 相当于温度 T=1,也就是直接把 logits 做归一化。当 T 增大,概率分布变得更平滑,类别之间的相对关系会更清楚;T 太小,分布就接近 one-hot,软标签的信息基本被抹掉。
典型的知识蒸馏损失函数由两部分组成:一部分是学生模型直接学硬标签的交叉熵,另一部分是学生和教师软标签之间的 KL 散度。实践中用 alpha 控制两者权重。下面是一个最简单的蒸馏损失函数示例:
import torch import torch.nn.functional as F def kd_loss(student_logits, teacher_logits, labels, T=3.0, alpha=0.7): # 软标签部分:KL散度 # 乘 T*T 是为了抵消温度缩放带来的梯度量级变化 s_soft = F.log_softmax(student_logits / T, dim=-1) t_soft = F.softmax(teacher_logits / T, dim=-1) kd = F.kl_div(s_soft, t_soft, reduction='batchmean') * (T * T) # 硬标签部分:正常交叉熵 ce = F.cross_entropy(student_logits, labels) return alpha * ce + (1 - alpha) * kd这个函数里的关键点有三个:
- 教师 logits 要除以同一个 T,再做 softmax。
- 学生 logits 也要除以 T,然后取 log_softmax,才能和教师分布计算 KL 散度。
- kd 部分乘以 T 的平方,是因为温度缩放了梯度量级,不乘回来会导致软标签部分的损失被压得太低。
alpha 的取值决定了学生更听教师的软标签,还是更听真实硬标签。alpha 偏大,学生更接近普通训练,蒸馏痕迹变弱;alpha 偏小,学生容易过度模仿教师的判断偏好,甚至把教师的错误也学过来。通常我会从 alpha=0.7、T=3.0 开始调,先跑通,再看验证集效果调整。
2. 开放模型逼近前沿,蒸馏到底起了什么作用
2.1 三种常见的蒸馏路径
讨论模型蒸馏时,首先要区分路径。不同路径对教师模型可见度、训练成本和适用场景的要求完全不同。
第一种是 logits 蒸馏,也就是白盒蒸馏。这种方式需要拿到教师模型的输出 logits,通常要求教师权重开放,或者在同一个框架里能直接做前向计算。它适合分类、排序、结构化理解这类有明确输出空间的任务,也是学术论文里最常见的蒸馏方式。
第二种是特征蒸馏。它在 logits 之外,还对齐教师模型中间层的表征,让学生模型逐步学习教师抽取特征的方式。这种方式训练更复杂,但很多场景下效果比只蒸馏输出层更好,尤其是在视觉模型和小型语言模型上。
第三种是数据蒸馏,也叫黑盒蒸馏或输出蒸馏。它不需要教师权重,只要能稳定调用教师模型,拿到教师生成的文本、答案、推理过程,再把这些高质量输出整理成学生模型的训练数据。开放模型逼近前沿时,数据蒸馏是讨论最多的一条路径。
为什么数据蒸馏这么关键?因为教师模型往往就是目前能力最强的闭源或开源模型之一,学生模型只需要通过二次训练吸收教师的回答风格、推理链和知识覆盖。整个过程相当于把“模型能力”通过数据搬运到更小或更开放的模型上,而不必从零复现一次巨型预训练。
2.2 数据蒸馏为什么更容易被反复讨论
开放模型和前端的差距,大致可以分为三块:基础能力、指令遵循、推理深度。数据蒸馏在这三块上都能起作用,只是做法不同。
基础能力方面,可以用教师模型生成大量答案,然后用筛选后的高质量子集继续预训练或增量训练。指令遵循方面,教师可以生成多轮对话、格式化和工具调用样例。推理深度方面,常见做法是让教师生成带思维链的解题过程,学生通过学习这些中间步骤来提升推理能力。
这里最值得注意的不是“能不能蒸”,而是“蒸出来的数据干不干净”。如果教师模型本身输出了错误答案,学生模型在训练时并不会自动识别错误,反而会把这个错误当成标准答案记住。所以数据蒸馏通常要搭配一套筛选流程,比如多个教师投票、人工抽检、规则校验、奖励模型排序等。单纯把教师输出全部灌给学生,风险很高。
2.3 成本账:为什么蒸馏是高效的杠杆
从零训练一个前沿模型,需要的数据量、算力和时间都非常大,绝大多数团队不具备这个条件。蒸馏把成本结构改变了。
| 路径 | 需要什么 | 主要成本 | 适合场景 |
|---|---|---|---|
| 从零预训练 | 海量文本、大规模算力集群、超长训练周期 | 极高 | 真正做底层基础模型的团队 |
| logits 蒸馏 | 可访问教师 logits、学生模型、GPU 训练资源 | 中 | 分类、排序、结构化任务 |
| 数据蒸馏 | 教师模型调用、数据筛选流水线、微调算力 | 中低 | 指令跟随、对话、特定领域能力增强 |
从成本角度看,数据蒸馏的低门槛在于它把最贵的部分外包给了教师模型。教师模型已经完成大规模预训练和后续对齐,学生要做的只是“学习教师写好的答案”。这也是为什么很多开放模型能够以较小参数量逼近前沿能力:不一定每个能力都是自己从零长出来的,有一部分是通过蒸馏和合成数据补上的。
3. 自己动手跑一次蒸馏实验
3.1 环境准备
如果你只是想理解蒸馏的运行机制,不需要一上来就碰大模型。先用小模型跑通流程,比直接上几十亿参数要快得多,也更容易调试。
我建议的环境配置:
- Python 3.10 或更高版本。
- PyTorch 2.x,安装了 CUDA 版本。
- transformers、datasets 库。
- 一块显存不低于 8GB 的 GPU,如果只用小模型,显卡压力很小。
我这里以一个简单的文本分类场景为例。教师模型用较大的预训练模型,学生模型用一个小型模型,数据集用常见的句子分类数据集。你可以换自己的数据集,但最好先确保数据集是一条条文本加一个标签的结构。
另外,准备环境时最容易踩的坑是依赖版本。transformers 版本太旧,可能不支持某些模型的加载方式;PyTorch 和 CUDA 版本不匹配,程序会在第一次 forward 时报底层错误。建议先跑一个最小样例,确认模型能加载、数据能过 dataloader,再开始正式训练。
3.2 完整训练流程
先加载教师模型和学生模型,教师模型必须切成 eval 模式,并且整个训练过程不更新参数。教师模型一旦处于训练模式,dropout 和归一化层会随机扰动输出,蒸馏目标就不稳定了。
teacher.eval() student.train() optimizer = torch.optim.AdamW(student.parameters(), lr=2e-5) for batch in dataloader: input_ids = batch["input_ids"].to(device) attention_mask = batch["attention_mask"].to(device) labels = batch["labels"].to(device) with torch.no_grad(): teacher_logits = teacher( input_ids=input_ids, attention_mask=attention_mask ).logits student_logits = student( input_ids=input_ids, attention_mask=attention_mask ).logits loss = kd_loss( student_logits, teacher_logits, labels, T=3.0, alpha=0.7 ) optimizer.zero_grad() loss.backward() optimizer.step()如果数据集不大,建议先把教师模型的 logits 预计算并保存到磁盘,训练时直接读取。这样可以减少一半的前向计算量,也能避免教师模型和目标数据在分布式训练中被重复搬运。实际项目中,我还习惯每隔几百步保存一次学生模型 checkpoint,防止中途断掉。
3.3 怎么判断蒸馏是否有效
训练结束后,不要只看训练集 loss。我一般会同时跑三个模型做对比:
- 只用硬标签训练的学生模型。
- 用蒸馏方式训练的学生模型。
- 教师模型本身。
然后放在同一个验证集上,比较准确率、F1 或你关心的任务指标。蒸馏有效的表现,是学生模型在验证集上的指标高于“只用硬标签训练”的对照组,而不是学生模型的 loss 比教师更低。学生模型的容量有限,不可能在所有维度上超过教师,但它的优势是推理更快、参数量更小。
另外,除了指标,还要看错误分布。比如分错的类别是更接近教师模型的偏好,还是完全随机的。如果学生模型和教师的错误高度一致,说明它确实学到了教师的判断模式;如果错误完全杂乱,说明蒸馏目标没有有效传递,可能需要调温度或者 alpha。
4. 蒸馏模型的边界:哪些坑不能只看指标
4.1 学生模型不会超过教师,错误也会被放大
蒸馏的本质是知识迁移,不是无中生有。学生模型的能力上限通常不会超过教师模型,尤其在同一分布的数据上。教师如果对某个领域根本不了解,学生通过蒸馏也学不到这个领域的正确知识,反而可能把教师信誓旦旦的错误答案当成标准内容记住。
这一点在数据蒸馏场景里尤其明显。教师生成的回答如果包含幻觉,学生学会了之后,等于把幻觉固化进了权重。后续无论怎么提示,学生都可能复现这条错误路径。所以蒸馏项目必须配备数据质检环节,不能只看生成量。
4.2 “蒸馏一本书”这种说法要分清
最近经常看到“蒸馏一本书”“把知识库蒸馏成一个 skill”之类的说法。严格来说,这跟模型蒸馏不是一回事。模型蒸馏发生在神经网络训练阶段,转移的是模型内部判断逻辑和概率分布。把一本书整理成知识库、再让模型通过检索或微调来学习,本质上是文档处理、RAG 和指令微调的组合,并不涉及教师模型的 logits 传递。
我会把这两件事分开看。前者适合做领域知识增强,后者适合做模型能力压缩。如果你把书里的内容直接当训练文本微调小模型,效果不一定差,但它不是知识蒸馏的标准做法,也不能简单套用蒸馏损失函数。
4.3 数据授权、模型协议和合规边界
这是最容易被忽略、但实际影响最大的一块。用教师模型 API 的输出做蒸馏训练,需要先确认平台的用户协议是否允许。很多模型的协议会明确限制用输出训练竞争模型,有的是完全禁止,有的是允许但要求注明来源。开源权重模型的权重许可协议也有差异,有些要求衍生模型继续开放相同协议。
我的建议是,在启动任何蒸馏项目之前,先列一份清单:
- 教师模型的权重或 API 输出是否允许用于训练。
- 训练数据是否包含受版权保护的文本、代码或图片。
- 蒸馏后的学生模型发布时,需要遵循什么许可证。
- 是否需要在模型卡片里注明使用了教师模型输出。
这些不是形式问题。如果中间数据来源不清晰,后续发布模型时很容易陷入授权纠纷。合规成本应该算进蒸馏项目的初始成本里。
4.4 什么时候不该用蒸馏
蒸馏不是万能的。如果你需要让模型掌握大量长尾知识,蒸馏可能不是最优解,因为你还要先把这些知识变成教师能生成的高质量答案,成本很高,不如检索增强来得直接。如果你的学生模型参数极小,容量严重不足,强行蒸馏只会得到一个什么都沾一点、什么都不精的模型。如果你需要快速更新模型,每次教师模型升级都要重新蒸馏一遍,维护成本也会很高。
我一般会先问三个问题:教师能稳定拿到吗?学生容量够吗?数据质量能管控吗?三个问题里有一个不满足,就先不要急着上蒸馏。
5. 常见问题与排查顺序
5.1 损失不下降或者震荡
先看几个最容易出错的地方。
第一,教师模型有没有切成 eval 模式。第二,教师 logits 和学生 logits 是否来自同一个 tokenizer 和同一个 label 空间。如果教师用了不同的分词方式,或者类别 id 没有对齐,蒸馏损失会一直乱跳。第三,学习率是不是太大。蒸馏任务的目标分布比硬标签更平滑,学习率可以比普通微调稍低一些。
我还习惯做一个消融验证:先把蒸馏损失里的 alpha 调成 1.0,只保留交叉熵,看学生能不能正常过拟合训练集。如果走不通,说明问题不在蒸馏逻辑,而在数据加载、模型结构或优化器配置。
5.2 蒸馏后的指标比普通训练还差
出现这种情况,不要急着怀疑蒸馏没用。先检查温度和 alpha 的取值。温度太高,软标签变得过于均匀,学生基本学不到类别间差异;alpha 太小,硬标签信号被严重削弱,学生可能只模仿教师,反而失去了对真实标签的敏感度。
我建议做一组小规模搜索:
- T 取 1、2、3、5。
- alpha 取 0.3、0.5、0.7、0.9。
每组只跑少量步数,观察验证集趋势,再选择稳定的一档做全量训练。不要一上来就追求跑满所有 epoch,蒸馏实验的调试成本主要在参数组合。
另外,如果学生模型容量和教师差距过大,比如教师是 70B、学生只有 0.5B,单靠 logits 蒸馏提升很有限。这种情况更适合数据蒸馏,也就是先生成教师的高质量回答,再在小模型上做指令微调,而不是直接对齐 logits。
5.3 数据蒸馏场景:并发、重试和中断
做数据蒸馏时,很多人会把精力放在模型训练上,却忽略了数据生产环节。调用教师模型生成数据时,最容易遇到的问题是接口限流、超时、返回格式错误和任务中断。
我的经验是先跑单条,再跑批量。单条确认输入提示词、输出字段、保存格式都正常之后,再开较小规模的并发。不要一上来就把并发拉到最高,因为接口限流和超时往往到批量阶段才会暴露。批量任务最好做成可断点续跑的结构,也就是每次调用成功就立即保存一条结果,并记录已经处理的 id。任务中断后,只需要从断点继续,不需要重新跑完整批数据。
输出文件命名也要提前规划。每个样本最好带上任务 id、教师模型版本、生成时间和参数信息,避免模型升级后无法追溯数据来源。这些细节看似不重要,真正到了模型发布和问题回溯时,缺失的元数据就是最短的短板。
6. 落地建议:从什么时候开始考虑蒸馏
如果你想在真实项目里引入蒸馏,我给出的顺序是这样的。
第一步,先定义清楚要解决什么问题。是模型太大推理太贵,还是想让小模型具备更强的指令跟随能力,或者单纯想让开放模型在某个领域更接近前沿表现。问题不同,蒸馏路径就不同。
第二步,做一个最小验证。用一个小教师、一个小学生、一个小数据集,把损失函数、训练脚本和评估流程跑通。这一步能帮你快速发现工具链问题,也能让你对温度和 alpha 有直觉。
第三步,再放大到真实数据。放到真实数据后,优先盯住数据质量和教师输出的一致性,而不是盲目增加训练步数。先小规模看效果,确认有效再扩大数据量。
第四步,把蒸馏流程固化成可复用的流水线。包括数据生产、质检、训练、评测、发布五个环节。只要中间任何一环依赖人工手动处理,整个流程就谈不上稳定。
说到底,蒸馏只是把已有知识高效搬运的手段,并不是一个神奇的开关。能不能逼近前沿,最终看的还是数据质量、模型容量和评测闭环是否完整。硅谷那边讨论开放模型追赶时,真正关注的也不是某一次蒸馏实验的指标,而是这条技术路径能不能持续降本、能不能稳定复现。把它当成系统工程来做,比把它当成一个 trick 来用,价值要大得多。