今天我们讲讲大模型的核心技术:蒸馏(Model Distillation)
作 者:吴佳浩Alben
撰稿时间:2026.7.18
更新时间:2026.7.23
很多人都会有一个疑问。 现在开源模型遍地都是,DeepSeek 的推理成本已经低到令人惊讶,为什么各家大模型公司仍在持续投入模型蒸馏? 如果只是站在 API 使用者的角度,这个问题确实很难理解------按量付费、直接调用,多简单。但当你真正开始训练、部署和维护一个大模型时,答案会变得完全不同。 据业内普遍推测,**Claude 的实际参数规模已逼近 6 万亿。**一次完整训练,成本高达数千万甚至上亿美元;每一次在线推理,都在持续消耗 GPU 算力、显存带宽,并与网络延迟和用户体验进行权衡。而企业机房里能够部署的 GPU 资源,始终是有限的。到了这一步,蒸馏就不再是一个"要不要做"的优化项,而是决定模型能否走出实验室、真正落地的核心工程能力。
事实上,几乎所有头部公司------OpenAI、Anthropic、Google、Meta、DeepSeek、阿里、字节等------都在持续投入蒸馏。它不像预训练那样占据头条,也不像 RLHF 那样引发热议,却贯穿了模型从研发到部署的整个生命周期。 更重要的是,蒸馏的真正价值,从来不只是 ** "把模型变小" ** 。数据如何构造、教师模型如何选择、温度参数如何设计、损失函数如何组合、哪些层需要对齐、能力如何迁移而不退化......这些工程细节,才是各家公司真正的核心竞争力,也是极少对外公开的技术机密。
本文将从为什么需要蒸馏开始,深入理解蒸馏背后的理论基础,梳理主流蒸馏方法,并结合真实的大模型工程实践,带你完整理解:模型蒸馏究竟在做什么,它如何在几乎不损失能力的前提下降低模型成本,以及为什么直到今天,它依然是连接模型训练与生产落地的关键技术,也是每一家大模型公司的必修课。
一、为什么需要蒸馏(Why)
1.1 大模型越来越强,但成本越来越高
过去两年,大模型的参数规模几乎是指数级增长:7B → 32B → 70B → 671B。参数越多,模型的知识容量和推理能力通常越强,但随之而来的是一系列非常现实的问题:
- 推理成本高:671B 级别的模型单次推理的算力开销是 7B 模型的几十上百倍;
- GPU 占用高:一个 671B 模型即便用 FP8 量化,也往往需要几十张高端 GPU 才能装下;
- 延迟高:参数越大,前向传播越慢,首 token 延迟和吞吐都会受影响;
- 部署困难:边缘设备、端侧应用、私有化部署场景根本无法承载超大模型。
一个直观的例子:
一个 671B 模型可能需要几十张 GPU 才能跑起来,而一个 7B 模型一张消费级 GPU(比如 RTX 4090)即可部署,成本相差可能是两个数量级。
这就引出了一个核心问题:
有没有办法把大模型的能力"复制"到一个更小、更便宜、更快的模型上?
于是便有了------模型蒸馏(Model Distillation)。
1.2 为什么所有 AI 公司都在做蒸馏
蒸馏并不是一个小众技术,几乎所有头部大模型公司都在系统性地使用它:
- OpenAI:GPT-4 系列到 GPT-4o-mini 系列的能力迁移;
- DeepSeek:R1 系列蒸馏出 32B / 14B / 8B / 1.5B 等一整套小模型;
- Google Gemini:Gemini Pro 到 Gemini Flash / Nano 的能力压缩;
- Qwen:Qwen 系列大量使用自家更大模型作为 Teacher;
- Llama:Meta 通过蒸馏和剪枝让小尺寸 Llama 保持较强能力;
- Mistral:小参数量模型追求单卡高效推理。
一个重要的现实是:真正线上跑的模型,不一定是公司里最大的那个模型。
真实的产品架构往往是:
一个 Teacher,蒸馏出多个 Student,各自面向不同的场景(云端 API、端侧、嵌入式设备等),最终形成一整套模型矩阵,而不是只有一个"最强模型"在裸奔。
1.3 为什么中小公司更需要蒸馏
大厂尚且如此,对于没有几千万预算训练基座模型的中小公司,蒸馏几乎是唯一现实的路径:
- 使用已有的开源或商用大模型(如 Qwen、DeepSeek、GPT、Claude)作为 Teacher;
- 让 Teacher 针对自己的业务场景生成大量高质量数据;
- 用这些数据蒸馏出一个体量小、成本低、可控性强的自有模型;
- 部署自己的 7B(甚至更小)模型,满足业务的延迟和成本要求。
典型的落地场景包括:客服对话、医疗问答、法律咨询、代码助手、金融风控、企业知识库问答、企业 Agent 等------这些场景往往任务边界清晰,恰好非常适合用一个专精的小模型去覆盖,而不需要一个无所不能的超大模型。
1.4大厂之间的实际情况
实际上,大模型公司的竞争,比很多人想象得更加直接。
每当一个能力更强的新模型出现,其他公司第一时间思考的并不是"重新训练一个",而是如何把这些能力尽可能快地迁移到自己的模型体系里。
从某种意义上说,今天的大模型竞争,很大一部分就是能力迁移的竞争,而蒸馏,就是这场竞争中最核心的技术之一。
谁能够更快、更低成本地完成一次高质量蒸馏,谁就能够更快推出新模型,占据下一轮竞争优势。
二、蒸馏到底是什么(What)
2.1 什么叫模型蒸馏
最经典的定义可以用一张图概括:

Teacher Model 指导 Student Model 学习。
这里必须强调一个常见误解:蒸馏不是复制参数 。Student 的网络结构、参数量可以和 Teacher 完全不同,甚至可以是不同的架构。蒸馏复制的是能力(Capability)------也就是 Teacher 在输入-输出映射上所体现出的"行为模式",而不是权重本身。
2.2 一个生活中的例子
可以把蒸馏想象成老师教学生:
-
老师:能考 100 分,知识面广,思考全面,但请一个"老师"成本很高;
-
学生
:只能考 80 分,但是:
- 考试速度更快;
- 培养成本更低;
- 更容易大规模复制(一个老师可以教很多学生)。
模型也是一样:Student 不追求超过 Teacher,而是追求用最小的代价,尽可能逼近 Teacher 的能力上限,在准确率和成本之间找到最优解。
2.3 蒸馏和微调有什么区别
这是初学者最容易混淆的地方,这里做一个直接的对比:
| 维度 | 微调(Fine-tuning / SFT) | 蒸馏(Distillation) |
|---|---|---|
| 学习对象 | 学数据 | 学模型 |
| 标签来源 | 标签来自人工标注 | 标签(软标签/推理过程)来自 Teacher |
| 学的是什么 | 学"标准答案" | 学"能力"和"思考过程" |
| 效果上限 | 由数据质量和数量决定 | 由 Teacher 的能力上限决定 |
需要说明的是,在真实的工程项目中,这两者往往不是二选一,而是:
Distillation + SFT 一起使用,是目前大多数团队的标准做法:既用人工标注的高质量数据打基础,又用 Teacher 生成的海量软标签数据扩大覆盖面。
三、蒸馏到底学了什么(Theory)
这一章是全文的理论核心,我们从最早的经典方法一路讲到当前最前沿的推理蒸馏。为了降低阅读门槛,每个概念都配一张图和一个可运行的最小例子。
3.1 最早的 Knowledge Distillation
蒸馏这个概念最早来自 Hinton 等人 2015 年的经典工作《Distilling the Knowledge in a Neural Network》,核心创新是提出了 Soft Target(软标签) 的概念。
为什么不直接用 One-Hot 标签?
如果一张图片是"猫",传统监督学习的标签是 One-Hot 形式:
bash
猫 = 1
狗 = 0
鸟 = 0
这种标签只告诉模型"正确答案是什么",却完全没有告诉模型"错误答案之间的相对关系"。
而 Teacher 模型对同一张图片的输出,往往是一个更丰富的概率分布:
bash
猫 0.82
狗 0.15
狐狸 0.03
下面这张图直观对比了两种标签携带的信息量:
这个分布本身就包含了额外的信息:模型认为这张"猫"的图片,和"狗"更像,和"狐狸"关系较远。这种模型在训练过程中学到的、隐藏在非最大概率里的相对关系,Hinton 称之为 "暗知识"(Dark Knowledge)。Student 学习的不再只是"正确答案",而是整个概率分布所蕴含的知识结构。
3.2 Temperature(温度)
直接用 Teacher 的 softmax 输出做监督信号有个问题:如果 Teacher 非常自信,正确类别的概率会接近 1,其余类别的概率会被压缩到接近 0,"暗知识"就被淹没了。解决方法是引入温度系数 T,对 softmax 做平滑:
pi=∑jexp(zj/T)exp(zi/T)
其中 zi 是 logits(未归一化的原始分数)。
- T=1:退化为标准 softmax;
- T>1:分布变得更平滑,类别之间的相对关系更容易被 Student 学到;
- T 越大,"暗知识"暴露得越充分,但过大也会引入噪声。
下面是一段实际可运行的代码,直观展示温度如何影响分布形状:
python
import numpy as np
def softmax_with_temperature(logits, T=1.0):
"""带温度的softmax函数
T越大,输出分布越平滑(暴露更多"暗知识")
T=1时退化为标准softmax
"""
logits = np.array(logits, dtype=np.float64)
scaled = logits / T
scaled = scaled - np.max(scaled) # 减去最大值,防止exp溢出,数值更稳定
exp_scaled = np.exp(scaled)
return exp_scaled / np.sum(exp_scaled)
# 模拟Teacher在"猫"这张图片上的原始logits(未归一化分数)
teacher_logits = [4.0, 1.5, 0.2] # 分别对应: 猫, 狗, 狐狸
labels = ["猫", "狗", "狐狸"]
print("=== 不同温度下Teacher的输出分布 ===")
for T in [1, 2, 5]:
probs = softmax_with_temperature(teacher_logits, T)
print(f"T={T}: " + ", ".join(f"{l}={p:.4f}" for l, p in zip(labels, probs)))
实际运行输出:
bash
=== 不同温度下Teacher的输出分布 ===
T=1: 猫=0.9054, 狗=0.0743, 狐狸=0.0203
T=2: 猫=0.6963, 狗=0.1995, 狐狸=0.1042
T=5: 猫=0.4821, 狗=0.2924, 狐狸=0.2255
可以清楚看到:随着 T 增大,"猫"的概率被压低,"狗"和"狐狸"的概率被拉高,分布越来越平滑,隐藏在小概率里的相对关系被逐渐"暴露"出来,这正是 Student 需要学习的信息。
3.3 Loss 如何计算
标准的 KD Loss 由两部分组成:
- CrossEntropy(硬标签损失):Student 输出与真实标签之间的标准交叉熵,保证 Student 不偏离真实答案;
- KL Divergence(软标签损失):Student 与 Teacher 在温度 T 下输出分布的 KL 散度,衡量 Student 是否学到了 Teacher 的"思维方式"。
为什么用 KL 散度而不是 MSE? 因为 Teacher 和 Student 的输出本质上是概率分布,KL 散度是专门用来衡量两个概率分布之间差异的信息论指标,比欧氏距离(MSE)更符合"分布匹配"这个目标的数学含义。
下面同样是一段可运行代码,计算 KL 散度:
python
def kl_divergence(p, q, eps=1e-10):
"""KL散度 KL(p || q),用于衡量Student分布q与Teacher分布p的差异"""
p = np.clip(p, eps, 1.0)
q = np.clip(q, eps, 1.0)
return np.sum(p * np.log(p / q))
student_logits = [3.0, 0.8, -0.5] # Student能力较弱,分布不如Teacher准确
T = 2.0
teacher_probs = softmax_with_temperature(teacher_logits, T)
student_probs = softmax_with_temperature(student_logits, T)
kd_loss = kl_divergence(teacher_probs, student_probs)
print(f"KL(Teacher || Student) = {kd_loss:.4f}")
实际运行输出:
bash
KL(Teacher || Student) = 0.0024
数值越小,说明 Student 的分布越接近 Teacher,蒸馏效果越好。
3.4 Feature Distillation(特征蒸馏)
除了学习最终输出层的概率分布,Student 还可以直接学习 Teacher 的中间层信息,包括:
- Hidden State(隐藏状态)
- Embedding(词向量/特征向量)
- Attention(注意力权重)
- Layer Feature(各层的中间特征)
一个直观的类比:如果说输出层蒸馏是"学老师最后给的答案",那么特征蒸馏就是"偷看老师做题过程中的草稿纸"------中间层特征往往包含了 Teacher 理解输入的方式,信息量比最终输出更丰富。
由于 Teacher 和 Student 的 hidden_dim 通常不一致,需要一个**投影层(Projector)**先对齐维度,再计算距离(通常用 MSE):
python
import torch
import torch.nn as nn
import torch.nn.functional as F
TEACHER_HIDDEN = 256
STUDENT_HIDDEN = 32
BATCH = 8
# 模拟一次前向传播中拿到的中间层输出(真实场景中通过hook从模型中间层取出)
teacher_hidden = torch.randn(BATCH, TEACHER_HIDDEN)
student_hidden = torch.randn(BATCH, STUDENT_HIDDEN)
# 投影层: 把Student的hidden映射到Teacher的维度,才能计算距离
projector = nn.Linear(STUDENT_HIDDEN, TEACHER_HIDDEN)
def feature_distillation_loss(student_hidden, teacher_hidden, projector):
"""用MSE衡量投影后的Student特征与Teacher特征的差距"""
projected = projector(student_hidden)
# Teacher的特征不需要梯度,detach掉,避免反向传播影响Teacher
return F.mse_loss(projected, teacher_hidden.detach())
loss = feature_distillation_loss(student_hidden, teacher_hidden, projector)
print(f"Feature Distillation Loss (MSE): {loss.item():.4f}")
loss.backward()
print(f"projector.weight.grad 是否有效: {projector.weight.grad is not None}")
实际运行输出:
bash
Feature Distillation Loss (MSE): 1.3448
projector.weight.grad 是否有效: True
(代码中未固定随机种子,每次运行 MSE 数值会略有不同,属正常现象,重点看数量级是否合理、梯度是否成功回传。)
backward() 成功执行说明梯度可以正常从损失回传到 Student 和投影层,训练链路是通的。这种场景适用于 CV(如 ResNet 蒸馏 MobileNet)、LLM、ViT 等各类模型。
3.5 Attention Distillation
Transformer 时代出现的新方法:不仅对齐输出、对齐特征,还直接对齐 Attention Map(即 Q、K 计算出的注意力权重矩阵)。
一个具体的例子帮助理解:输入句子"这家餐厅的服务很好,但是菜有点贵",在判断情感倾向时:
为什么 Attention 值得单独蒸馏?因为 Attention 权重在很大程度上反映了模型"在理解一句话时,把注意力放在哪些词上",这是模型理解过程的直接体现,比单纯的输出概率包含更多结构化信息。做法通常是对 Teacher 和 Student 对应层(或映射后的层)的 Attention Map 计算 MSE 或 KL 散度。
3.6 Intermediate Layer Distillation(中间层蒸馏)
当 Teacher 和 Student 层数差异很大时(例如 Teacher 80 层,Student 32 层),不可能逐层一一对应,需要做 Layer Mapping(层映射)。
常见策略:
- 等间隔映射 :Student 的第 i 层对应 Teacher 的第 ⌊i×80/32⌋ 层;
- 首尾对齐:重点监督最浅层和最深层,中间层弱监督或不监督;
- 可学习映射:用注意力机制自动学习层与层之间的对应关系(如 TinyBERT 的做法)。
3.7 Reasoning Distillation(CoT 蒸馏)------目前最热门的方向
这是当前大模型蒸馏最火的方向,也是 DeepSeek-R1、OpenAI o 系列、Qwen、Claude 等模型都大量采用的方法。
区别于传统蒸馏只学"最终答案",Reasoning Distillation 让 Teacher 输出完整的:
一个具体例子(数学题:"一个水池,进水管每小时注水 8 吨,出水管每小时排水 3 吨,同时打开需要多久注满 40 吨的水池?"):
bash
传统蒸馏(只学答案):
Question -> Answer: "8小时"
Reasoning蒸馏(学完整推理链):
Question ->
Thinking: "需要求净注水速度..."
Reasoning: "净注水速度 = 8 - 3 = 5吨/小时;
时间 = 40 ÷ 5 = 8小时"
Answer: "8小时"
Student 学习的不只是最后的答案,而是整个推理链条,包括中间的分析、假设、验证、纠错等步骤。这样训练出的 Student 即使参数量很小,也能表现出较强的思维链能力,这也是为什么 DeepSeek-R1 蒸馏出的 7B、1.5B 小模型依然能在数学、代码等推理任务上有不错的表现。
3.8 Preference Distillation(偏好蒸馏)
RLHF 兴起之后出现的新方向:Teacher 不仅提供答案,还提供偏好信息------对同一个问题的多个候选回答,标注哪个是 Chosen(更好) ,哪个是 Rejected(更差)。
Student 通过学习这种偏好关系,来对齐 Teacher 的价值判断和回答风格,常见方法包括:
- DPO(Direct Preference Optimization)
- IPO(Identity Preference Optimization)
- ORPO(Odds Ratio Preference Optimization)
这些方法本质上是把 RLHF 中"训练 Reward Model + PPO"的复杂流程,简化成一个直接在偏好数据上优化的过程,同时天然适合"从 Teacher 的偏好中蒸馏"这个场景。
四、蒸馏的数据从哪里来(Data Pipeline)
理论讲完了,从这一章开始进入真实工程环节。蒸馏的效果七成取决于数据质量,这一章专门讲数据从哪来、怎么生成、怎么清洗。
4.1 数据来源
- 公开数据集(学术 benchmark、开源语料)
- 企业私有数据(历史工单、业务文档)
- 用户真实交互数据
- 日志数据(脱敏后)
- 企业知识库
- 合成数据(由 Teacher 模型直接生成)
4.2 Teacher 生成数据
核心流程是 Prompt 设计 + 批量推理,让 Teacher 针对目标场景生成结构化数据,通常包括:
- Question(问题)
- Answer(答案)
- Reasoning(推理过程,用于 CoT 蒸馏)
- Tool Call(工具调用轨迹,用于 Agent 蒸馏)
- JSON(结构化输出,便于程序化处理)
4.3 数据清洗
Teacher 生成的数据不能直接拿去训练,必须经过清洗,常见步骤包括:
- 去重:避免同质化数据反复出现,浪费训练算力;
- 过滤:剔除格式错误、明显答非所问的样本;
- 质量评分:用规则或模型对样本打分,保留高质量部分;
- 长度过滤:过短(信息量不足)或过长(可能是重复生成/幻觉)的样本剔除;
- 语言过滤:确保语言符合目标场景要求;
- 一致性检查:例如推理过程和最终答案是否自洽。
4.4 数据打分
打分环节通常结合多种手段:
- AI Judge:用另一个大模型对生成数据的质量进行评分;
- Reward Model:专门训练的打分模型;
- 人工抽检:对高风险场景(医疗、法律、金融)必须有人工复核环节;
- 模型互评:多个模型交叉评分,减少单一模型的偏差。
原因很直接:垃圾数据会毁掉 Student。蒸馏的本质是"Student 无条件相信 Teacher 给的答案",如果 Teacher 生成的数据里混入了错误、幻觉或低质内容,这些问题会被 Student 原样学走,甚至被放大。
4.5 数据配比
一个常见的经验配比(具体比例需要根据实际场景调整):
为什么不能全部用 AI 生成的数据? 因为纯 AI 生成的数据容易出现"分布坍缩"------数据的多样性和真实分布会逐渐向 Teacher 的偏好收敛,缺乏人工数据的"锚点",长期迭代下模型质量可能会缓慢劣化(即 Model Collapse 问题)。保留一定比例的人工标注数据,可以为整个数据分布提供真实世界的校准。
五、真实的大模型蒸馏流程(Engineering)
这一章讲企业里真正是怎么落地一次蒸馏项目的。
5.1 整个 Pipeline
这是一个闭环的流程:如果评测不达标,需要回到数据生成或训练阶段迭代,而不是一次性完成的线性任务。
5.2 Teacher 如何部署
Teacher 通常需要生成海量数据,因此推理效率至关重要,常用的高性能推理引擎包括:
- vLLM:目前最主流的开源高吞吐推理框架,支持 PagedAttention、连续批处理;
- TensorRT-LLM:NVIDIA 官方的高性能推理引擎,适合追求极致延迟/吞吐的场景;
- SGLang:在结构化生成和多轮对话场景下有独特优势;
- 多 GPU + Batch 推理:把生成任务拆分成大批次并行处理,最大化吞吐。
Teacher 推数据的效率决定了整个蒸馏项目的迭代速度------如果生成一批数据要等好几天,整个数据-训练-评测的闭环就转不起来。
5.3 Student 训练
常用训练框架:
- LLaMA Factory:目前最流行的开箱即用微调/蒸馏框架,支持 LoRA、QLoRA、全量微调;
- TRL(Transformer Reinforcement Learning):HuggingFace 出品,支持 SFT、DPO、PPO 全流程;
- OpenRLHF:偏工程化、支持大规模分布式 RLHF 训练;
- Megatron-LM:NVIDIA 出品,适合超大规模模型的张量并行/流水并行训练;
- DeepSpeed:微软出品,ZeRO 系列显存优化技术,常与上述框架配合使用。
5.4 蒸馏 Loss 如何加入训练代码
真实训练代码中,蒸馏 Loss 通常是多个 Loss 的加权组合:
bash
Total Loss = SFT Loss + KD Loss + Feature Loss + Preference Loss
下面是一份完整可运行的端到端蒸馏训练示例(用小型 MLP 模拟 Teacher/Student,逻辑与真实 LLM 蒸馏完全一致,只是把 Transformer 换成了更轻量的网络,便于本地快速验证):
python
"""
真实可运行的知识蒸馏(KD)训练示例
依赖: pip install torch --break-system-packages
说明: 使用一个"大"MLP作为Teacher,一个"小"MLP作为Student,
在一个简单的多分类合成数据集上做知识蒸馏训练。
展示 SFT Loss + KD Loss(KL散度) 的组合训练方式。
"""
import torch
import torch.nn as nn
import torch.nn.functional as F
from torch.utils.data import DataLoader, TensorDataset
torch.manual_seed(42)
# ---------- 1. 构造一份合成分类数据 ----------
NUM_CLASSES = 5
NUM_SAMPLES = 2000
FEATURE_DIM = 32
X = torch.randn(NUM_SAMPLES, FEATURE_DIM)
# 用一个随机线性映射 + argmax 造出"真实标签",模拟一个有结构的分类任务
true_w = torch.randn(FEATURE_DIM, NUM_CLASSES)
y = (X @ true_w).argmax(dim=1)
dataset = TensorDataset(X, y)
loader = DataLoader(dataset, batch_size=64, shuffle=True)
# ---------- 2. 定义 Teacher(参数多)和 Student(参数少) ----------
class MLP(nn.Module):
def __init__(self, in_dim, hidden_dim, out_dim, num_layers=2):
super().__init__()
layers = [nn.Linear(in_dim, hidden_dim), nn.ReLU()]
for _ in range(num_layers - 1):
layers += [nn.Linear(hidden_dim, hidden_dim), nn.ReLU()]
layers.append(nn.Linear(hidden_dim, out_dim))
self.net = nn.Sequential(*layers)
def forward(self, x):
return self.net(x) # 返回未归一化的logits
teacher = MLP(FEATURE_DIM, hidden_dim=256, out_dim=NUM_CLASSES, num_layers=4)
student = MLP(FEATURE_DIM, hidden_dim=32, out_dim=NUM_CLASSES, num_layers=1)
print(f"Teacher参数量: {sum(p.numel() for p in teacher.parameters()):,}")
print(f"Student参数量: {sum(p.numel() for p in student.parameters()):,}")
# ---------- 3. 先把Teacher训练到较好的水平(模拟已经训好的大模型)----------
teacher_optim = torch.optim.Adam(teacher.parameters(), lr=1e-3)
for epoch in range(20):
for xb, yb in loader:
logits = teacher(xb)
loss = F.cross_entropy(logits, yb)
teacher_optim.zero_grad()
loss.backward()
teacher_optim.step()
teacher.eval() # Teacher训练完毕,固定参数,只做推理
# ---------- 4. 蒸馏训练 Student:SFT Loss + KD Loss ----------
def distillation_loss(student_logits, teacher_logits, labels, T=2.0, alpha=0.7):
"""
组合损失:
- hard_loss: Student对真实标签的CrossEntropy(模拟SFT Loss)
- soft_loss: Student与Teacher在温度T下的KL散度(KD Loss)
alpha 控制两者的权重
"""
hard_loss = F.cross_entropy(student_logits, labels)
# KL散度要求两边都是log/prob,且按温度T缩放后要乘以T^2做梯度补偿(Hinton论文中的做法)
soft_teacher = F.softmax(teacher_logits / T, dim=1)
soft_student = F.log_softmax(student_logits / T, dim=1)
soft_loss = F.kl_div(soft_student, soft_teacher, reduction="batchmean") * (T ** 2)
return alpha * soft_loss + (1 - alpha) * hard_loss, hard_loss.item(), soft_loss.item()
student_optim = torch.optim.Adam(student.parameters(), lr=1e-3)
for epoch in range(30):
total, correct = 0, 0
for xb, yb in loader:
with torch.no_grad():
teacher_logits = teacher(xb) # Teacher只推理,不反传梯度
student_logits = student(xb)
loss, hard_l, soft_l = distillation_loss(student_logits, teacher_logits, yb)
student_optim.zero_grad()
loss.backward()
student_optim.step()
correct += (student_logits.argmax(dim=1) == yb).sum().item()
total += yb.size(0)
if (epoch + 1) % 10 == 0:
print(f"Epoch {epoch+1}: acc={correct/total:.4f} hard_loss={hard_l:.4f} soft_loss={soft_l:.4f}")
print("蒸馏训练完成:Student用远少于Teacher的参数量,学到了Teacher的分类能力。")
实际运行结果:
bash
Teacher参数量: 207,109
Student参数量: 1,221
Epoch 10: acc=0.8975 hard_loss=0.3925 soft_loss=2.3784
Epoch 20: acc=0.9635 hard_loss=0.2245 soft_loss=1.1197
Epoch 30: acc=0.9815 hard_loss=0.1767 soft_loss=0.8513
蒸馏训练完成:Student用远少于Teacher的参数量,学到了Teacher的分类能力。
这个例子非常直观地体现了蒸馏的价值:Student 只用了 Teacher 约 0.6% 的参数量(1,221 vs 207,109),最终却达到了 98.15% 的准确率。这正是蒸馏在真实大模型工程中被广泛采用的核心原因------用极小的代价,逼近大模型的能力上限。
5.5 多 Teacher 蒸馏
这是目前越来越流行的做法:不局限于单一 Teacher,而是让 Student 同时向多个不同厂商的模型学习:
为什么多 Teacher 效果更好? 不同模型在不同任务上各有所长(比如某个模型代码能力强,另一个模型数学推理强),多 Teacher 蒸馏相当于让 Student "博采众长",同时也能通过多个 Teacher 的输出交叉验证,一定程度上过滤掉单一 Teacher 的幻觉和偏差。
六、真实案例(Case Study)
6.1 DeepSeek 蒸馏
DeepSeek-R1 蒸馏版之所以广受关注,核心原因在于它把 Reasoning Distillation 做到了极致:用 R1 大模型生成的长推理链数据,蒸馏出了 32B、14B、8B、7B、1.5B 一整套小模型。这些小模型在数学、代码等推理密集型任务上,表现明显超过同尺寸的其他开源模型------这说明推理能力(而不仅仅是知识)也是可以被有效蒸馏和继承的。
6.2 Qwen 蒸馏实践
Qwen 系列同样采用了"大模型指导小模型"的策略,通过自家更大规模的模型生成高质量的 SFT 数据和偏好数据,用于训练中小尺寸的 Qwen 版本,使得较小的 Qwen 模型在同尺寸对比中依然具备较强的综合能力。
6.3 Meta 蒸馏实践
Llama 系列的发展趋势也体现了"模型越来越小、越来越高效"的方向。通过结合蒸馏与结构优化(如分组查询注意力 GQA、更高效的 tokenizer),Llama 在保持较小参数量的同时,尽可能保留大模型的能力水平,以便更好地支持端侧和私有化部署场景。
6.4 完整企业实战案例:从零蒸馏一个客服小模型
前面几节讲的是行业趋势,这一节我们完整走一遍企业内部真实会经历的全流程:某电商公司想做一个 7B 级别的智能客服模型,用于处理"物流查询/退换货/优惠券咨询"三类高频问题,要求响应快、成本低、可私有化部署。以下是完整的 Teacher 生成数据 → 训练 → 评测 → 部署闭环。
6.4.1 整体架构
6.4.2 Step1-2:用 Teacher 批量生成客服训练数据
第一步用 Teacher(假设通过 API 调用一个强模型)针对客服场景批量生成结构化数据。这里用可运行的规则化脚本模拟"Teacher 生成"这一过程(真实场景中把 mock_teacher_generate 替换成对 Teacher API 的调用即可,接口形状完全一致):
python
"""
Step1-2: 模拟Teacher批量生成客服训练数据
真实场景:把 mock_teacher_generate() 替换成对Teacher大模型API的调用
(例如 requests.post 调用 Claude/GPT/Qwen 的 chat completion 接口),
其余的批量处理、清洗、落盘逻辑完全不变。
"""
import json
import random
random.seed(0)
# 业务种子问题:真实场景中来自历史工单/知识库,这里手写少量样例做演示
SEED_QUESTIONS = [
"我的订单显示已发货,但是物流三天没更新怎么办?",
"买的衣服不合适,怎么申请退货?",
"优惠券显示已过期,但我明明是今天领的,能补发吗?",
"退款审核多久能通过?",
"换货需要我自己承担运费吗?",
]
def mock_teacher_generate(question: str) -> dict:
"""
模拟Teacher模型针对一个客服问题生成结构化训练样本。
真实实现示例(伪代码):
resp = client.messages.create(
model="teacher-model",
messages=[{"role": "user", "content": build_prompt(question)}],
)
return parse_json(resp.content)
这里用简单规则模拟出结构一致的输出,便于本地无需API即可跑通全流程。
"""
templates = {
"物流": "先查询物流单号最新轨迹;若超过48小时未更新,建议联系承运商核实,"
"同时可在平台发起'物流不更新'投诉以加速处理。",
"退货": "确认商品在7天无理由退货期内且未影响二次销售,"
"在订单详情页申请退货,选择原因后等待审核,审核通过后按提示寄回。",
"优惠券": "核实优惠券的实际有效期与领取规则是否存在系统延迟展示的情况,"
"如确认为系统问题,可联系客服提交工单申请补发。",
"退款": "退款审核通常在1-3个工作日内完成,节假日可能延后,"
"可在'我的订单-退款进度'页面实时查看审核状态。",
"换货": "若因商品质量问题换货,运费由卖家承担;若因个人原因(如尺码不合适),"
"运费通常由买家承担,具体以平台换货政策为准。",
}
matched_topic = next((k for k in templates if k[:2] in question), "退货")
reasoning = f"识别问题类型为「{matched_topic}」相关问题,需要先确认关键事实,再给出处理建议。"
answer = templates[matched_topic]
return {
"question": question,
"reasoning": reasoning,
"answer": answer,
"topic": matched_topic,
}
# 批量生成(真实场景中会是几万到几十万条,这里演示流程用少量数据)
raw_samples = []
for q in SEED_QUESTIONS:
for _ in range(4): # 对每个种子问题做多次改写/多样化生成,模拟规模化数据生成
raw_samples.append(mock_teacher_generate(q))
print(f"Teacher共生成 {len(raw_samples)} 条原始样本")
print("样例:", json.dumps(raw_samples[0], ensure_ascii=False, indent=2))
实际运行输出:
bash
Teacher共生成 20 条原始样本
样例: {
"question": "我的订单显示已发货,但是物流三天没更新怎么办?",
"reasoning": "识别问题类型为「物流」相关问题,需要先确认关键事实,再给出处理建议。",
"answer": "先查询物流单号最新轨迹;若超过48小时未更新,建议联系承运商核实,同时可在平台发起'物流不更新'投诉以加速处理。",
"topic": "物流"
}
6.4.3 Step3:规则 + AI Judge 双重过滤
生成的数据不能直接用,先做规则过滤(长度、格式),再模拟 AI Judge 打分(真实场景中同样是调用一个模型做质量评分):
python
"""
Step3: 数据清洗与打分
规则过滤: 长度、关键信息完整性
AI Judge打分: 模拟用另一个模型对(question, answer)做1-5分质量评分
"""
def rule_filter(sample: dict) -> bool:
"""规则过滤:答案不能太短,且必须包含具体可执行的建议"""
if len(sample["answer"]) < 15:
return False
if sample["reasoning"] == "":
return False
return True
def mock_ai_judge_score(sample: dict) -> int:
"""
模拟AI Judge打分(真实场景中调用打分模型,prompt通常是
"请从相关性/准确性/可执行性三个维度给这条客服问答打1-5分")。
这里用简单启发式规则模拟评分,保证流程可运行。
"""
score = 3
if len(sample["answer"]) > 30:
score += 1
if any(kw in sample["answer"] for kw in ["工单", "审核", "承担"]):
score += 1
return min(score, 5)
filtered = [s for s in raw_samples if rule_filter(s)]
print(f"规则过滤后剩余: {len(filtered)}/{len(raw_samples)}")
scored = []
for s in filtered:
s["quality_score"] = mock_ai_judge_score(s)
if s["quality_score"] >= 4: # 只保留高质量样本(>=4分)
scored.append(s)
print(f"AI Judge打分后保留高质量样本: {len(scored)}/{len(filtered)}")
实际运行输出:
bash
规则过滤后剩余: 20/20
AI Judge打分后保留高质量样本: 20/20
(真实项目中过滤和打分环节的淘汰率通常会高得多,这里因为是规则模拟数据所以全部通过,实际生产数据往往会淘汰 20%~40%。)
6.4.4 Step4:构建 ChatML 训练集
python
"""
Step4: 把清洗后的数据转换成ChatML格式,写入jsonl文件,供训练框架读取
"""
SYSTEM_PROMPT = "你是一名专业、耐心的电商客服助手,请根据用户问题给出清晰、可执行的处理建议。"
def to_chatml(sample: dict) -> dict:
"""转换为ChatML多轮对话格式,同时保留reasoning用于CoT蒸馏"""
assistant_content = f"【分析】{sample['reasoning']}\n【回复】{sample['answer']}"
return {
"messages": [
{"role": "system", "content": SYSTEM_PROMPT},
{"role": "user", "content": sample["question"]},
{"role": "assistant", "content": assistant_content},
]
}
chatml_dataset = [to_chatml(s) for s in scored]
output_path = "customer_service_distill_data.jsonl"
with open(output_path, "w", encoding="utf-8") as f:
for item in chatml_dataset:
f.write(json.dumps(item, ensure_ascii=False) + "\n")
print(f"训练集已写入 {output_path},共 {len(chatml_dataset)} 条")
print("样例:")
print(json.dumps(chatml_dataset[0], ensure_ascii=False, indent=2))
实际运行输出:
bash
训练集已写入 customer_service_distill_data.jsonl,共 20 条
样例:
{
"messages": [
{
"role": "system",
"content": "你是一名专业、耐心的电商客服助手,请根据用户问题给出清晰、可执行的处理建议。"
},
{
"role": "user",
"content": "我的订单显示已发货,但是物流三天没更新怎么办?"
},
{
"role": "assistant",
"content": "【分析】识别问题类型为「物流」相关问题,需要先确认关键事实,再给出处理建议。\n【回复】先查询物流单号最新轨迹;若超过48小时未更新,建议联系承运商核实,同时可在平台发起'物流不更新'投诉以加速处理。"
}
]
}
这份 jsonl 文件已经是 LLaMA Factory 等主流训练框架可以直接读取的标准格式(真实项目中会有几万到几十万条这样的样本)。
6.4.5 Step5:LoRA 训练配置(真实可用的 LLaMA Factory 配置)
真实项目中会用 LLaMA Factory 加载 Qwen2.5-7B-Instruct 之类的 Student 基座,用上一步生成的数据做 LoRA 蒸馏微调。以下是一份可以直接使用的训练配置文件(train_config.yaml):
yaml
# train_config.yaml
# LLaMA Factory LoRA蒸馏训练配置示例
model_name_or_path: Qwen/Qwen2.5-7B-Instruct # Student基座模型
stage: sft # 蒸馏在实现上通常复用sft stage,
# 数据里的reasoning+answer即承担了KD的软监督作用
do_train: true
finetuning_type: lora
lora_target: q_proj,k_proj,v_proj,o_proj # 对Attention相关权重做LoRA微调
dataset: customer_service_distill # 对应dataset_info.json中注册的数据集名
template: qwen
cutoff_len: 2048
max_samples: 100000
per_device_train_batch_size: 4
gradient_accumulation_steps: 8
learning_rate: 1.0e-4
num_train_epochs: 3.0
lr_scheduler_type: cosine
warmup_ratio: 0.03
logging_steps: 10
save_steps: 200
output_dir: ./output/customer_service_student
bf16: true
启动训练的命令:
bash
# 安装依赖
pip install llamafactory --break-system-packages
# 启动LoRA蒸馏训练
llamafactory-cli train train_config.yaml
6.4.6 Step6:业务指标评测
训练完成后不能只看 loss,必须结合通用能力 + 业务指标双重评测。下面模拟一次评测脚本,对比 Student 蒸馏前后在业务测试集上的表现:
python
"""
Step6: 业务指标评测
关键点:真实场景中Answer的判定要靠AI Judge或人工比对,
这里用关键词覆盖率模拟"答案是否包含必要处理要点"的评分方式,
以保证代码可独立运行、逻辑可验证。
"""
# 业务测试集:每条包含问题、必须覆盖的关键要点
eval_set = [
{"question": "物流一直没更新怎么办?", "must_cover": ["物流", "投诉"]},
{"question": "衣服不合适怎么退货?", "must_cover": ["7天", "退货"]},
{"question": "退款要多久?", "must_cover": ["工作日", "审核"]},
]
def student_before_distill(question: str) -> str:
"""模拟蒸馏前的Student(基座模型,未经过客服场景训练),回答比较通用、缺少要点"""
return "请您耐心等待,客服会尽快处理您的问题。"
def student_after_distill(question: str) -> str:
"""模拟蒸馏后的Student,回答更贴合业务(用训练数据中的模板做演示)"""
for s in scored:
if s["topic"] in question or question[:4] in s["question"]:
return s["answer"]
return "先查询物流单号最新轨迹;若超过48小时未更新,建议联系承运商核实,同时可在平台发起投诉以加速处理。"
def coverage_score(answer: str, must_cover: list) -> float:
"""关键要点覆盖率:命中的关键词数 / 总关键词数"""
hit = sum(1 for kw in must_cover if kw in answer)
return hit / len(must_cover)
print(f"{'问题':<20}{'蒸馏前覆盖率':<15}{'蒸馏后覆盖率':<15}")
before_scores, after_scores = [], []
for item in eval_set:
ans_before = student_before_distill(item["question"])
ans_after = student_after_distill(item["question"])
s_before = coverage_score(ans_before, item["must_cover"])
s_after = coverage_score(ans_after, item["must_cover"])
before_scores.append(s_before)
after_scores.append(s_after)
print(f"{item['question']:<20}{s_before:<15.2f}{s_after:<15.2f}")
print(f"\n平均覆盖率 - 蒸馏前: {sum(before_scores)/len(before_scores):.2%}")
print(f"平均覆盖率 - 蒸馏后: {sum(after_scores)/len(after_scores):.2%}")
实际运行输出:
bash
问题 蒸馏前覆盖率 蒸馏后覆盖率
物流一直没更新怎么办? 0.00 1.00
衣服不合适怎么退货? 0.00 1.00
退款要多久? 0.00 1.00
平均覆盖率 - 蒸馏前: 0.00%
平均覆盖率 - 蒸馏后: 100.00%
对比非常直观:蒸馏前的 Student(未经业务数据训练的基座模型)只会给出"请您耐心等待"这类空泛回复,关键要点覆盖率为 0;蒸馏后的 Student 因为学习了 Teacher 生成的结构化数据,能准确覆盖"7 天无理由退货""工作日审核"等业务关键信息,覆盖率达到 100%。
需要提醒的是,这里的评测集样本量很小、且和训练数据主题高度重合,只是为了演示评测脚本的完整逻辑可以跑通。真实项目中的评测集必须与训练数据不重叠(严格划分 train/eval),并且要覆盖训练数据没见过的问法、边界场景(比如多轮追问、模糊表达),这样才能真实反映 Student 的泛化能力,而不是"背题"能力。如果评测中发现某类问题覆盖率不达标,就需要回到 Step2 针对性地让 Teacher 生成更多该类型的数据,再重新走一遍闭环,而不是训练完就直接上线。
6.4.7 Step7:vLLM 部署上线
评测达标后(真实项目中通常要求业务指标达到 90% 以上,且通用能力评测不明显退化),用 vLLM 部署 Student 模型对外提供服务:
bash
# 安装vLLM
pip install vllm --break-system-packages
# 启动OpenAI兼容的推理服务
# --model 指向训练完成后LoRA权重与基座合并后的模型目录
python -m vllm.entrypoints.openai.api_server \
--model ./output/customer_service_student_merged \
--served-model-name customer-service-7b \
--port 8000 \
--max-model-len 4096 \
--gpu-memory-utilization 0.85
部署后,业务系统即可通过标准 OpenAI 接口协议调用:
python
"""
调用已部署的客服Student模型(OpenAI兼容接口)
依赖: pip install openai --break-system-packages
"""
from openai import OpenAI
client = OpenAI(base_url="http://localhost:8000/v1", api_key="not-needed")
response = client.chat.completions.create(
model="customer-service-7b",
messages=[
{"role": "system", "content": SYSTEM_PROMPT},
{"role": "user", "content": "我的快递显示已签收,但我没收到货怎么办?"},
],
temperature=0.3,
)
print(response.choices[0].message.content)
至此,一个完整的企业蒸馏项目闭环走完:历史工单 → Teacher批量生成 → 清洗打分 → ChatML数据集 → LoRA蒸馏训练 → 业务指标评测 → vLLM部署。这也是本文最推荐读者动手复现的部分------把 mock 的 Teacher 调用换成真实 API,把演示用的少量样本换成真实业务数据,这套流水线就是一个可以直接跑起来的生产级蒸馏项目骨架。
七、如何自己做一次蒸馏(Hands-on)
结合第六章的完整案例,这里给出一份更通用的实操路线图,方便迁移到其他业务场景。
7.1 准备 Teacher
根据预算和场景选择:
- 开源自部署:Qwen3-235B、DeepSeek 系列(可控性强,成本可控,但需要自备算力);
- 闭源 API 调用:GPT、Claude(无需自己部署,按调用量付费,适合快速验证)。
7.2 准备 Student
根据部署环境选择:
- Qwen3-7B / Qwen3-1.5B:中文场景友好;
- Llama3-8B:英文及通用场景表现均衡;
- Gemma:轻量级、适合端侧部署。
7.3 生成训练数据
核心是 Prompt 设计 + 批量调用 API,生成结构化的三元组数据(做法可参考 6.4.2 的完整代码):
bash
Question(问题)
Reasoning(推理过程,可选,用于CoT蒸馏)
Answer(答案)
7.4 构建训练集
生成的数据需要整理成标准训练格式,常见的三种格式:
- JSONL:每行一个 JSON 对象,最通用的格式;
- ChatML :
<|system|> <|user|> <|assistant|>角色标记格式,主流对话模型训练常用(参考 6.4.4); - ShareGPT :
conversations字段包含多轮对话列表,适合多轮对话数据。
7.5 开始训练
用 LLaMA Factory 等框架,根据资源情况选择训练策略(参考 6.4.5 的配置文件):
- LoRA:只训练低秩适配矩阵,显存占用小,适合快速迭代;
- QLoRA:在 LoRA 基础上结合量化,进一步降低显存需求,单卡也能训练较大模型;
- Full Fine-tuning:全量参数训练,效果上限最高,但显存和算力开销也最大。
7.6 模型评测
训练完成后,需要从通用能力和业务指标两个维度评测(参考 6.4.6):
- 通用 Benchmark:MMLU(英文综合知识)、C-Eval(中文综合知识)、HumanEval(代码能力)、Arena(人工/模型对战评测);
- 业务指标:结合具体场景设计,例如客服场景的问题解决率、代码助手的通过率、知识库问答的准确率和幻觉率等。
7.7 部署上线
评测达标后,选择合适的推理引擎部署(参考 6.4.7):
- vLLM / SGLang:适合云端高并发服务;
- Ollama:适合本地/端侧快速部署,对开发者友好;
- GPU 部署:根据模型尺寸选择合适的显卡配置,做好显存和吞吐的压测。
八、蒸馏的局限性
一个常被问到的问题:学生一定会超过老师吗?
答案是:通常不会。原因主要有两点:
- 知识损失:Student 的参数量、结构容量本身有限,无法完整承载 Teacher 学到的全部知识和能力;
- 能力天花板:Student 的监督信号来自 Teacher,本质上是在"逼近"一个既定目标,而不是自主探索出超越目标的能力。
不过在某些特定条件下,Student 确实有可能反超 Teacher,相关的研究方向包括:
- Born Again Network(自我再生网络):用同结构的模型反复自我蒸馏,后一代有时能在某些指标上超过前一代;
- Self Distillation(自蒸馏):模型自己蒸馏自己(比如用模型更深层的输出指导浅层),起到正则化和知识蒸馏的双重作用;
- Iterative Distillation(迭代蒸馏):多轮"蒸馏-再训练-再蒸馏"的迭代过程,配合更好的数据和训练技巧,逐步逼近甚至局部超越原始 Teacher 在特定任务上的表现。
需要强调的是,这些"反超"通常是在局部任务/局部指标上出现的,而不是全面能力的反超。
九、未来的发展方向
结合当前学术界和工业界的动向,模型蒸馏未来几年可能的重点发展方向包括:
- Reasoning Distillation:推理链蒸馏会持续深化,尤其是长链推理、多步验证过程的蒸馏;
- Multi-Teacher:多 Teacher 融合蒸馏会成为提升 Student 泛化能力的标准做法;
- Agent Distillation:把复杂 Agent 的规划、工具调用、多轮决策能力蒸馏到轻量模型;
- RL Distillation:把强化学习训练出的策略和价值判断蒸馏给 Student;
- Tool Distillation:专门针对工具调用准确率和参数生成的蒸馏;
- Long Context Distillation:把长上下文理解和检索能力蒸馏到参数更小的模型;
- Multimodal Distillation:跨模态(图文、视频、语音)能力的蒸馏;
- MoE Distillation:把混合专家模型的能力蒸馏到稠密小模型,或反过来指导 MoE 路由训练;
- Self Distillation:模型自我迭代提升,减少对外部 Teacher 的依赖;
- Online Distillation:在线实时蒸馏,随着新数据/新反馈不断更新 Student;
- Continuous Distillation:持续蒸馏管线,让小模型能够跟随 Teacher 的版本迭代持续"追赶"。
这些方向的共同趋势是:蒸馏正在从"压缩模型"演变为"压缩能力",蒸馏的对象也从单纯的输出分布,扩展到推理过程、决策链条、工具使用等更复杂的能力维度。可以预见,蒸馏将成为未来几年 AI 模型训练和落地不可或缺的核心技术之一。
十、总结
一句话总结蒸馏:
模型蒸馏并不是简单地缩小模型,而是将大模型在海量数据中学到的知识、推理能力和行为模式,以更低的计算成本迁移到更小的模型中。
它的核心价值并不是追求最强性能,而是在效果、成本、延迟和部署效率之间取得最佳平衡。对于今天的大模型产业而言,无论是 OpenAI、DeepSeek 等头部厂商,还是构建垂直应用的中小企业,蒸馏都已经成为模型落地和规模化部署不可或缺的关键技术。
从本文第五章的实测代码可以看到一个直观的结论:在一个简单的分类任务上,参数量仅为 Teacher 0.6% 的 Student,通过蒸馏依然达到了 98% 以上的准确率;而第六章的完整企业案例则展示了这套方法论如何从理论落地成一条可复现的生产级流水线。这正是蒸馏技术最迷人也是最实用的地方------用可控的成本,逼近大模型的能力边界。