AI模型蒸馏折腾手记
这次的具体需求有三个硬指标:
- 模型大小要控制在 500MB 以内(目标设备是 4GB 内存)
- 推理延迟不能超过 100ms(需要实时响应)
- 准确率不能比原方案下降超过 5%
为什么要做这件事
事情是这样的。我们有个文本分类任务,原本直接用 GPT-4o 做推理,准确率确实不错,但每条消息都要调一次 API,延迟加上成本,看着就心疼。老板说:“能不能把它搞到本地跑?别老烧钱。”
本地跑有两种路子:要么找现成的开源模型直接微调,要么自己搞个蒸馏版本。
第一种路子试过了,BERT-base 在中文数据上表现还凑合,但离我们的要求差了一截;RoBERTa-large 准确率够,但显存吃紧,推理慢得让人怀疑人生。这时候,蒸馏就成了看起来唯一可行的方向。
什么是模型蒸馏
用人话说,模型蒸馏就是让一个"老师模型"把自己的知识教给一个"学生模型"。老师模型通常很大、很聪明,但跑起来很贵;学生模型小一点、笨一点,但跑起来快。
关键在于别只抄最终答案,还得学老师对各类别的打分方式。老师对一个问题的判断往往有多层信息,学生如果只记对错,会丢掉大量细节。
传统训练只看标签对不对,蒸馏训练要看"软标签"——也就是老师对每个可能答案的置信度分布。比如老师认为这个句子是正面情感的概率是 0.8、中性 0.15、负面 0.05,学生就要尽量学这个分布,而不是只记住"正面"这个标签。

这张图想说明:硬标签训练等于只记"选哪个",蒸馏训练更像是在模仿老师整张概率分布——每个类别各占多少,而不只是谁最大。
需求与限制
这次的具体需求有三个硬指标:
- 模型大小要控制在 500MB 以内(目标设备是 4GB 内存)
- 推理延迟不能超过 100ms(需要实时响应)
- 准确率不能比原方案下降超过 5%
环境限制也不少:
- 目标设备是嵌入式 Linux,只有 4GB RAM 和一个低端 GPU
- 数据集约 10 万条中文文本,分布比较均匀
- 训练资源有限,只有两块 RTX 3090 可用
实现过程
第一步:准备老师模型
老师模型我们选了 GPT-4o,通过 API 获取软标签。为了节省成本,我写了个批量脚本,每批 100 条一起请求,然后解析返回的 logprobs。
import openai
import json
from tqdm import tqdm
client = openai.Client(api_key="your-api-key")
def get_teacher_labels(texts, batch_size=100):
"""批量获取老师的软标签"""
results = []
for i in tqdm(range(0, len(texts), batch_size)):
batch = texts[i:i+batch_size]
response = client.chat.completions.create(
model="gpt-4o",
messages=[{"role": "user", "content": f"分类这段文本:{text}"} for text in batch],
max_tokens=10,
temperature=0.0,
logprobs=True,
top_logprobs=5
)
# 解析 logprobs 提取软标签
batch_results = []
for choice in response.choices:
logprobs = choice.logprobs.content[0].top_logprobs
soft_labels = {lp.token: np.exp(lp.logprob) for lp in logprobs}
batch_results.append(soft_labels)
results.extend(batch_results)
return results
这一步花了大概两小时,10 万条数据,费用有点肉疼。但软标签的质量确实比硬标签好不少,尤其是那些边界案例。
第二步:设计学生模型
学生模型选了 DistilBERT-base,参数量约 66M,推理速度快,适合中文场景。结构上做了点调整:
- 去掉了部分注意力头(12 个砍到 6 个)
- 隐藏层维度从 768 降到 512
- 层数从 6 层砍到 4 层
from transformers import DistilBertForSequenceClassification, DistilBertConfig
config = DistilBertConfig(
num_labels=3, # 正面、中性、负面
n_heads=6,
dim=512,
n_layers=4,
vocab_size=21128 # 中文词汇表
)
student_model = DistilBertForSequenceClassification(config)
这样设计后,模型大小大概 180MB,远低于 500MB 的上限。
第三步:蒸馏训练
蒸馏的核心是损失函数。我们要同时考虑:
- 软标签损失(KL 散度):学生要模仿老师的概率分布
- 硬标签损失(交叉熵):学生还要记住真实标签
- 温度系数 T:控制软标签的平滑程度
import torch
import torch.nn.functional as F
def distillation_loss(
student_logits,
teacher_logits,
labels,
temperature=3.0,
alpha=0.5
):
"""蒸馏损失函数"""
# 软标签损失(KL 散度)
soft_loss = F.kl_div(
F.log_softmax(student_logits / temperature, dim=-1),
F.softmax(teacher_logits / temperature, dim=-1),
reduction='batchmean'
) * (temperature ** 2)
# 硬标签损失(交叉熵)
hard_loss = F.cross_entropy(student_logits, labels)
# 加权组合
return alpha * soft_loss + (1 - alpha) * hard_loss
训练参数也没怎么调,就用了常用的套路:
- 学习率 5e-5,线性衰减
- Batch size 32,梯度累积 4 步
- 训练 10 个 epoch,早停 patience 3
踩坑记录
这一路踩的坑不少,挑几个有代表性的说说。
坑 1:软标签太"软"
一开始 T 设成 5.0,发现学生学得很"平均",不管什么问题都喜欢给个 0.4、0.35、0.25 之类的分布,像个不会做选择的老好人。
后来才明白,温度太高会过度平滑老师的判断,学生反而学不到边界感。改到 T=3.0 才好点,但具体数值还得看任务。二分类任务可能 2.0 就够了,多分类复杂点可以上到 4.0。
坑 2:硬标签权重调不对
alpha 最初设成 0.9,就是说 90% 的损失来自软标签。结果学生跟着老师跑偏了,遇到老师没见过的数据就瞎猜。
改成 alpha=0.5 后才正常,软标签和硬标签各占一半。这个值确实得调,不能想当然就选个极端值。原则是:软标签够好时可以高一点,数据量不够时就得给硬标签多点权重。
坑 3:学生模型太小
第一次设计的模型太小,层砍到 3 层,维度降到 384,结果根本学不动。训练过程中损失一直在震荡,准确率比随机猜测强不了多少。
后来才意识到:学生再小,也得有基本的学习能力。DistilBERT-base 原来的结构是有道理的,砍得太多就伤了命。现在这个 4 层、512 维的版本算是折中,能学又不太大。
坑 4:推理还是慢
模型训练好了,放到目标设备上跑,发现推理延迟 80ms 左右,离 100ms 的红线很近。进一步 profiling 发现, tokenizer 耗了一半时间。
后来换了个更快的 tokenizer 实现,又加了缓存机制,把常用词的编码结果缓存起来,总算降到了 60ms 以内。但这提醒我:模型大小不是唯一的瓶颈,预处理也得考虑。
结果与反思
折腾了几周,最后的结果还算能接受:
| 指标 | 原方案 (GPT-4o) | 蒸馏后 | 变化 |
|---|---|---|---|
| 准确率 | 92.3% | 88.7% | -3.6% |
| 模型大小 | API (不计) | 178MB | - |
| 推理延迟 | 350ms | 58ms | -83% |
| 推理成本 | $0.001/条 | $0 (本地) | -100% |
准确率确实降了一些,但在可接受范围内;延迟和成本改善很明显,尤其是成本这块,跑一个月就能把之前烧的钱省回来。
这次实践下来,我更清楚大模型强在哪:理解上下文、处理模糊输入、做细粒度判断——这些比"参数多"难压缩得多。
蒸馏能做的是:把老师已经明确掌握的知识高效地传递给学生。但老师的"直觉"、“常识”、“推理能力”,这些说不清的东西,学生很难照单全收。
另一个教训是:别指望一次蒸馏就搞定。中间的 T、alpha、学生结构,都得调。尤其是学生结构,得根据具体任务和数据特点来设计,不能盲目照抄论文里的配置。
最后说几句
模型蒸馏这事,说起来是个技术活,但背后其实是个权衡问题:你要什么,你愿意放弃什么,你能接受多大的损失。
我们的场景刚好适合这种 trade-off:准确率要求没那么高,但成本和延迟必须紧着来。如果你的任务是医疗诊断、金融风控这类敏感场景,可能就得换条路子。
技术这行就是这样,没有银弹。蒸馏是工具,不是神药。用对地方它能帮你省钱省时间,用错了就是白忙活。
这轮蒸馏先收在这儿。指标达标了,中间踩过的坑更值得留着——下次换老师模型或换任务,还能对照着少走弯路。
版权声明: 本文首发于 指尖魔法屋-AI模型蒸馏折腾手记(https://blog.thinkmoon.cn/post/276-ai-model-distillation-teacher-student-practice/) 转载或引用必须申明原指尖魔法屋来源及源地址!
评论
使用 GitHub 账号登录后即可留言,支持 Markdown。