最近不少团队都在问我同一个问题:模型太多、显卡太少、线上资源越来越紧张,大模型效果是好,可真要部署到端侧或者高并发服务里,又跑不动。然后大家就都盯上了知识蒸馏。知识蒸馏说白了就是让一个能力更强的大模型当老师,把学到的知识“搬”给一个小模型,让小模型在体积小、推理快的前提下尽量逼近老师的水平。
一句话讲清楚它解决什么问题:大模型负责把能力学到极致,小模型负责把能力用起来,蒸馏是两者之间那座桥。适合谁看?正在做端侧部署、模型上线、预算有限但不想放弃大模型效果的人,以及刚接触模型压缩、想系统搞懂蒸馏原理并能动手跑一遍完整流程的同学。这篇文章我尽量用大白话把原理讲透,再给一次能直接参考的完整实战。
1. 先搞明白:知识蒸馏到底在“搬”什么能力
1.1 大模型比小模型多出来的,不只有参数
很多人有个误解,觉得大模型强是因为参数多、算力大,所以蒸馏就是把参数“复制粘贴”过去。这个理解完全不对。参数没法直接复制,就算复制过去,小模型也装不下。真正可以迁移的,是模型在训练过程中形成的“决策方式”和“知识结构”。
举个最容易理解的例子。让大模型看一张图,它不仅能说出“这是一只猫”,还能说出“这只猫的耳朵比较尖,毛发偏橘色,脸型有点像狐狸”。前一句话是类别标签,后一句是模型在大量数据里学到的细微判断依据。小模型如果只看“这是一只猫”这种标签去学,它只能记住一个个孤立的事实,遇到没见过的猫就懵了。但如果让它跟着大模型学“这只猫因为耳朵尖、毛发橘、脸型像狐狸,所以大概率是猫”,它就学到了判断的边界和依据,泛化能力自然更强。
这个道理在自然语言处理里同样成立。比如情感分类,大模型面对一句“这家店的菜一般,但服务还不错”,它学到的不只是“正面”或“负面”这种单一标签,而是每个词对结论的贡献程度。“菜一般”偏向负,“服务不错”往正拉,两者对冲之后模型给出了“中性偏正”的概率分布。这种分布信息就是小模型真正该学的东西。
1.2 软标签、温度系数与前100字必须懂的过渡
知识蒸馏里最核心的一个概念叫“软标签”。传统训练给小模型看的是硬标签,比如“正面=1,负面=0”;蒸馏给小模型看的是软标签,比如“正面=0.55,负面=0.45”。软标签里藏着大模型对样本“犹豫”的程度,这种犹豫不是缺点,恰恰是知识最密集的地方。
为了让软标签的信息更充分暴露出来,Hinton那篇经典论文里引入了“温度系数”T。温度越高,模型输出的概率分布越平滑,越能暴露出那些原本被压到很小但依然有意义的概率值。打个比方,硬标签像一张只标了“终点”的地图,软标签像一张标了所有值得走的路线的地图,温度系数就是放大镜,帮你把地图上那些细碎但关键的路线也看得清清楚楚。
这里我先给一个最直观的认知,后面实战部分会带你算一遍完整的损失函数,你会亲眼看到温度系数是怎么影响训练的。
1.3 小模型和大模型的关系像“徒弟学师傅的思路”
用师徒关系来理解知识蒸馏特别贴切。师傅(大模型)做题时会把每一步的思路、可能的错误选项、正确解法的依据都讲给徒弟听,徒弟(小模型)不用死记硬背标准答案,而是去学师傅的分析路径。等徒弟上了考场,哪怕题目变了,他也能用师傅教的分析方式去推理,而不是凭着记忆里的标准答案硬套。
所以蒸馏训练时,小模型同时看两份材料:一份是真实标签,也就是标准答案;另一份是教师的软标签,也就是解题思路。两份材料一起学,效果通常比只看任何一份都好。这也是初学蒸馏的人最容易忽略的点——有人觉得软标签替代了硬标签,实际上两者互补,组合起来效果才稳。
2. 动手前的方案设计:模型怎么选、数据怎么备
2.1 选教师模型:不是越大越好,但要“本事过硬”
教师模型的选择直接决定蒸馏效果的天花板。如果教师自己水平就不行,教出来的学生肯定也强不到哪去。但教师也不是越大越好,大模型推理一次的成本如果在你的预算内高得离谱,整条方案就跑不动了。
从实操角度看,选教师有三个关键原则:
- 教师能力要显著优于你能接受的小模型下限,否则蒸馏没有意义。
- 教师和学生的任务必须完全一致,比如都是文本分类、都是序列标注,任务错位会导致知识迁移失败。
- 教师的输出格式要能方便地保存和复用,最好一次性离线把教师对所有训练样本的预测结果都存成文件,训练学生时直接读取,不用反复让教师推理。
我自己的经验是:文本分类这类任务,教师用同结构的更大模型(比如BERT-base 当教师、TinyBERT 当学生)就够;如果资源允许,用更大的生成式模型输出文本类的辅助信息也是锦上添花,但运算成本会陡增。第一次跑通流程,没必要盲目追求超大模型。
2.2 选学生模型:越小越好?还得看部署目标
学生模型的选择要回到你的部署目标上。如果最终目标是放进微信小程序或浏览器前端,那参数量要控制在几MB到几十MB级别;如果目标是跑在移动端App里,可以放宽到几十MB;如果在服务器上做高并发推理,几百MB也不是不能接受。
小模型的选择有一条隐藏原则:结构最好和教师保持一定的“血缘关系”。不是说必须一模一样,而是尽量选择同类的骨干网络。比如教师是BERT架构,学生选TinyBERT或者层数更少的BERT变体,知识迁移的摩擦会小很多。原因在于特征分布、注意力头数的继承性更好,学生更容易理解教师输出的表达方式。
如果你用的是完全不同的结构,比如教师是Transformer、学生是CNN,那也不是不行,但需要更多的调参和训练数据来弥补结构差异带来的分布偏移。新手不建议一上来就搞这种高难度玩法。
2.3 数据准备:没有额外标注数据也能做,但有三件事必须做
蒸馏最让人舒服的一点是:不需要额外的标注数据。你可以直接用原始训练集,把教师模型的预测当作标注来用。但有三件事必须提前做好,否则后面容易返工。
第一,数据质量要过一遍手。去重、清洗、长度截断这些常规操作不能少。教师模型虽然有较强的容错能力,但你让它学一堆乱数据,它照样会生产出乱七八糟的软标签。
第二,样本分布尽量均衡。如果分类任务里某个类别的样本特别少,教师在这个类上的判断能力也弱,学生跟着学就会“继承”这种偏科。条件允许的话做一点简单的数据增强,比如同义词替换、回译,把稀缺类别的样本量垫一下。
第三,一定要单独留出验证集和测试集。蒸馏训练过程中同样需要监控过拟合,没有验证集等于闭眼开车。测试集更不用说,是最后衡量学生模型水平的唯一标准。
3. 完整实战:一次文本分类模型的蒸馏全过程
3.1 实战目标与基线设定
为了让这次实战足够具体,我用一个经典的场景:中文情感分类,二分类(正向/负向),数据集规模在五万条左右。教师模型用完整版的中文BERT(参数量约1.1亿),学生模型用一个小型的6层Transformer(参数量约1500万)。这么设定比较接近真实的端侧部署需求,你也能直观感受到“模型体积缩到1/7左右,效果到底能保留多少”。
先说一下基线:直接用同样的数据训练小模型,不经过蒸馏,准确率大约在91.2%左右;教师模型的准确率是95.8%。我们的目标是让蒸馏后的小模型冲到94%以上,把小模型与大模型的差距从近5个点压缩到2个点以内。
3.2 步骤一:离线保存教师模型的软标签
这一步的核心就是把教师模型对所有训练样本的预测概率保存下来。因为学生训练整个epoch里要多次用到教师预测,如果每次都对教师做一次前向推理,成本太高。离线保存相当于把教师的知识固化成一个文件,学生直接读文件,不用再理教师模型。
这里给一段参考的PyTorch代码,思路很清晰:
import torch import numpy as np def generate_soft_labels(model, dataloader, temperature=4.0, output_path="soft_labels.npy"): model.eval() all_probs = [] all_labels = [] with torch.no_grad(): for batch in dataloader: input_ids = batch["input_ids"].cuda() attention_mask = batch["attention_mask"].cuda() logits = model(input_ids, attention_mask=attention_mask).logits # 用温度系数软化概率分布 probs = torch.softmax(logits / temperature, dim=-1) all_probs.append(probs.cpu().numpy()) all_labels.append(batch["label"].numpy()) all_probs = np.concatenate(all_probs, axis=0) all_labels = np.concatenate(all_labels, axis=0) np.savez(output_path, probs=all_probs, labels=all_labels) print(f"soft labels saved to {output_path}, shape: {all_probs.shape}")温度系数设置为4.0是我试下来比较稳的默认值,后面你会看到为什么不要太低也不要太高。这一步实际执行时,用教师的精度模式跑一遍全量训练集,两万条样本在单张消费级显卡上大概几分钟就能完成,非常快。
3.3 步骤二:构造蒸馏损失函数,重点理解温度的选择
蒸馏的损失函数由两部分组成:一部分是学生预测与真实硬标签的交叉熵,另一部分是学生预测与教师软标签的KL散度。总损失是两者的加权和。用公式表达就是:
[ L = \alpha \cdot L_{hard} + (1-\alpha) \cdot T^2 \cdot L_{soft} ]
其中 (L_{hard}) 是学生预测与真实标签的交叉熵,(L_{soft}) 是按同一温度软化后学生与教师的KL散度,(T) 是温度系数,(\alpha) 是平衡两个损失的超参数。
有一个细节很多人不理解,为什么KL散度部分要乘 (T^2)。原因是温度T把概率分布变平了,梯度的大小也会跟着变小。如果不乘 (T^2),高温下的软标签贡献会被稀释得几乎学不到东西。这个 (T^2) 就是用来把梯度“拉回来”的补偿项。
我自己实验里的默认配置是:温度T=4.0,(\alpha=0.7)。意思是七分看真实标准答案,三分学教师思路。如果你的数据量很小,可以适当调高 (\alpha) 到0.8甚至0.9,因为数据少的时候真实标签更稀缺珍贵,不能过度依赖教师的主观判断。数据量大了,再逐渐降低 (\alpha),让教师的知识多发挥作用。
3.4 步骤三:学生模型训练全流程与关键配置
学生模型结构上我做了如下选择:6层Transformer,隐藏维度384,8个注意力头,参数量约1500万。训练时batch size取64,学习率设5e-5,使用AdamW优化器,线性学习率衰减。训练轮数设为6个epoch,每轮结束在验证集上监控准确率,保存最优checkpoint。
训练主循环的参考代码如下:
def train_student(student_model, teacher_logits, train_loader, val_loader, config): student_model.train() optimizer = torch.optim.AdamW(student_model.parameters(), lr=config["lr"]) scheduler = torch.optim.lr_scheduler.LinearLR(optimizer, total_iters=config["epochs"]) temperature = config["temperature"] alpha = config["alpha"] best_acc = 0.0 for epoch in range(config["epochs"]): total_loss = 0.0 for batch, (probs_file_batch) in zip(train_loader, teacher_logits): input_ids = batch["input_ids"].cuda() attention_mask = batch["attention_mask"].cuda() labels = batch["label"].cuda() soft_labels = torch.tensor(probs_file_batch).cuda() logits = student_model(input_ids, attention_mask=attention_mask).logits loss_hard = torch.nn.functional.cross_entropy(logits, labels) loss_soft = torch.nn.functional.kl_div( torch.log_softmax(logits / temperature, dim=-1), soft_labels, reduction="batchmean" ) * (temperature ** 2) loss = alpha * loss_hard + (1 - alpha) * loss_soft optimizer.zero_grad() loss.backward() optimizer.step() total_loss += loss.item() val_acc = evaluate(student_model, val_loader) print(f"epoch {epoch+1} loss: {total_loss/len(train_loader):.4f} val_acc: {val_acc:.4f}") if val_acc > best_acc: best_acc = val_acc torch.save(student_model.state_dict(), "student_best.pt")这里有个关键点要特别强调:soft_labels必须和batch的样本顺序完全对应。如果你在生成软标签时打乱了数据集顺序,训练时也必须用相同的乱序逻辑去加载batch,否则学生学到的是错位的知识,效果会极其糟糕。这是一个特别容易踩的坑,我下面第4部分会专门展开讲。
3.5 步骤四:蒸馏结果对比与效果评估
整个训练过程在单张消费级显卡上跑完大约需要40分钟。我这边的实验结果如下表所示:
| 模型 | 参数量 | 准确率 | 推理时延(CPU, 单条) | 模型存储 |
|---|---|---|---|---|
| 教师 BERT-base | 1.1亿 | 95.8% | 约45ms | 约420MB |
| 小模型(无蒸馏) | 1500万 | 91.2% | 约8ms | 约60MB |
| 小模型(蒸馏) | 1500万 | 94.3% | 约8ms | 约60MB |
三组对比非常直观:小模型蒸馏之后比不蒸馏高了整整3.1个点,跟教师的差距从4.6个点缩小到1.5个点;推理时延却只有教师的六分之一不到,模型体积只有教师的七分之一。这就是蒸馏的意义——用不到教师10%的资源,保住教师95%以上的效果。
我又额外做了一个测试:把教师换成效果更强的更大模型,同一批数据,学生的准确率又往上浮了0.8个点。说明教师能力越强,学生上限越高。所以有条件的话,直接用你能用得起的最强模型来当教师。
3.6 进阶操作:蒸馏后量化与在端侧的落地路径
蒸馏不是终点,模型最终是要落地的。我接着把蒸馏出来的学生模型做了8bit量化,存储体积进一步从60MB压缩到16MB,推理时延从8ms降到5ms左右,准确率基本没有变化,保持在94.1%。这个体积和速度就直接具备进小程序、移动端App的条件了。
如果想把量化也纳入训练过程,可以考虑量化感知训练(QAT),在训练时就模拟量化的噪声,让模型提前适应低精度表达。实测下来QAT比训练后直接量化在极限压缩场景下能多保住1到2个点。代价是训练时间会增加,且超参数更敏感。第一次做端侧模型不追求极限压缩的话,训练后普通量化已经够用了。
4. 踩坑实录:这些坑我建议你别再踩一遍
4.1 温度系数不是越大越好,软标签别“软过头”
刚开始做蒸馏时我也迷信高温,觉得温度越高暴露出的小概率信息越多。结果在某个多分类任务上把温度调到8.0,学生模型直接学崩了,准确率比不蒸馏还低。原因是温度太高,所有类别的概率都趋于均匀,软标签里几乎没有有效信息,学生相当于在学一堆“几乎随机”的答案,反而干扰了它对真实标签的学习。
建议的做法是在3.0到6.0之间做一个小网格搜索,每个温度跑一个短训练,选验证集表现最好的值。大多数分类任务里,4.0是个比较稳的中等值,可以先从这个值起步再微调。
4.2 软标签和训练样本的顺序错位
这是我第一次实战时出的最离谱的问题。离线生成软标签后,我换了一种数据加载方式,batch里样本的排列顺序和生成软标签时对不上,但训练代码还是老老实实地按位置去取软标签,结果学生模型的准确率一路跌到85%。整整排查了两个小时才发现是顺序错位。
解决办法也很简单:生成软标签时,把样本的索引一并保存下来;训练时按索引去匹配软标签。或者最稳妥的方案是生成软标签和训练学生用同一个DataLoader实例,保证完全同序。
4.3 教师和学生输入格式不一致导致的知识断裂
有一次我用一个多模态教师去蒸馏纯文本学生,教师的输入里混入了图像特征,学生的输入只有文本。虽然任务标签一致,但知识迁移效果非常差。原因在于教师有一部分知识编码在图像特征里面,文本输入根本继承不到。
如果教师和学生结构差异过大,请务必在蒸馏前做一个通道对齐。最简单的做法是选择与教师同族的学生结构;另一种做法是额外加一个特征对齐损失,让学生的中间层特征去逼近教师对应位置的特征。这个属于进阶玩法,需要花时间调参,但确实管用。
4.4 蒸馏训练中的过拟合问题
小模型参数量少,通常不太容易过拟合,但在我做蒸馏时发现如果训练轮数排得太长、学习率又不够低,学生会在训练集上趋近教师的能力之后,继续死磕训练集中的噪声,验证集表现反而下滑。
建议学生在验证集上监控,连续三轮不升就开始做早停;学习率也要比正常训练小模型时略微降低一些,我一般会打个七折到八折。蒸馏的目标是继承泛化能力,不是死记硬背训练集。
4.5 常见问题速查表
| 现象 | 可能原因 | 解决方案 |
|---|---|---|
| 学生准确率低于基线 | 温度过高/过低 | 在3.0-6.0区间调参,重置为4.0验证 |
| 学生训练不收敛 | 软标签乱序 | 检查样本索引与batch顺序是否对应 |
| 验证集准确率持续下降 | 训练轮数过长 | 早停或降低学习率 |
| 学生学到了教师的缺点 | 教师本身水平有限 | 更换更强的教师或过滤教师置信度低的样本 |
| 小模型部署后效果骤降 | 量化损失过大 | 改用QAT或减少量化压缩倍数 |
4.6 我个人的一套蒸馏检查清单
每次做新的蒸馏任务,我都会先过一遍固定检查清单,省了很多返工时间。教师模型能力够不够强、软标签与训练样本顺序是否完全对应、温度系数和损失权重是否做过小范围搜索、验证集是否单独隔离且没有被教师见过的数据混入、学生模型推理性能是否达到部署目标、量化后的效果是否被再次评估。
这六项都是基础但致命的点。不要嫌繁琐,蒸馏这个技术看起来简单,真正稳定复现效果靠的就是这些细节。
5. 蒸馏之后,还可以做什么
蒸馏不是终点,现在做模型压缩早就不满足于只走一条路了。我通常会把蒸馏和剪枝、量化、模型结构搜索组合使用。先蒸馏出一个小而强的学生,再做结构化剪枝去掉冗余头,随后量化压缩到极致。三步走完,模型体积能比原始大模型压缩到二十分之一,速度提升二十倍以上,效果损失控制在两个点以内。
另一个趋势是把蒸馏能力用在多模态大模型上。之前提到的大模型知识抽取框架,本质上也是把多模态大模型里的跨模态知识蒸馏到单模态小模型里,让纯文本或纯图像的小模型也能沾到多模态的光。这个方向目前还很前沿,适合有一定基础的团队跟进。
如果你刚接触蒸馏,我的建议是先跑通一条最基础的单任务、单教师、单学生流程,把数据集和代码都调试顺了,再慢慢做组合优化,不要一上来就想搞花活。基础流程的每一步都理解了,后面所有进阶玩法都是举一反三。
最后分享一个小技巧:实战时可以顺手把你的软标签和硬标签之间的差异可视化出来看,差异大的样本往往是训练集中最难啃的硬骨头,也是学生模型最需要重点学的地方。这部分样本占总体通常不超过10%,但对最终效果的影响非常大,值得单独做损失加权。踩过几次坑之后再回看蒸馏这件事,我是真觉得它门槛不高、上限很高,值得每一个做模型落地的人都完整跑一遍流程。