“大模型蒸馏”这四个字,最近在圈子里出现的频率实在太高了。朋友圈、技术群、开源社区,隔三差五就有人晒出同款标题的分享:1000条数据,蒸馏出一个领域专家模型。说实话,第一次看到这种帖子我也心动过——不需要几十万条标注数据,不需要几十张A100,就能让一个小模型在特定领域里逼近甚至超越大模型的表现,这个说法对任何团队都有巨大的吸引力。但当我真正带着团队把一个蒸馏项目从想法推到线上之后,才明白这句话背后藏着一大堆前提条件:数据从哪来、教师模型的输出怎么处理、学生模型选多大、损失函数怎么配、验证怎么做,每个环节都有讲究。做对了,1000条数据确实能省下几十万的标注预算;做错了,你只会得到一个看起来在跑、一用就露馅的花架子。
这篇文章我就拿最近做的一个法律问答领域蒸馏项目当例子,把整条链路完整复盘一遍,包括数据构造、蒸馏实现、效果对比和踩坑心得。想上蒸馏但不知道怎么下手的,照着这份流程走能少折腾至少两周。
1. 先搞清楚:大模型蒸馏到底在“蒸”什么
1.1 从师带徒说起:知识蒸馏的本质
知识蒸馏这个概念的经典表述,是让一个参数量巨大的教师模型(Teacher)把自己“做判断的方式”传授给一个小参数的学生模型(Student)。注意,这里说的是“做判断的方式”,不是“标准答案”。区别在哪?我给你打个比方。
老厨师带新徒弟,如果只给他一本菜谱,让他背配料表,他做出来的菜顶多是“能吃”;但老厨师要是站在旁边,让徒弟看他怎么颠勺、怎么掌握火候、怎么在收汁的最后一分钟判断浓稠度,徒弟学到的才是“做菜的手感”。大模型蒸馏也是这个道理——教师模型在生成每一个token的时候,内部会计算出一整份概率分布:它觉得“A”有70%的可能,“B”有20%,“C”有10%。这个分布就是教师的“手艺”。微调只告诉你结果选A,蒸馏则把A/B/C之间的概率关系也一并传给学生,这多出来的信息就是“手感”。
所以蒸馏的本质,是把一个黑盒大模型的软输出(soft label)当作监督信号,让学生模型去逼近这份概率分布,而不是单纯逼近那个最终选中的token。
1.2 为什么1000条数据可能“够用”
标题里的“1000条”是最容易被误读的地方。不少人以为这是说蒸馏只需要1000条训练数据,然后随便抓1000条QA就往上灌,最后效果一塌糊涂,回头骂“标题党”。真相是,1000条数据本身确实有可能够,但前提是这1000条里装的不是简单的“问题-答案”对,而是带着完整概率分布的教师输出。
我算过一笔账:一条法律问答,假设教师模型生成300个token,每个token附带一个覆盖5万词表的概率分布。即便我们只在分布里保留概率最高的前20个token,那一条数据携带的监督信息也比单纯一个硬标签高出一到两个数量级。换句话说,蒸馏场景下1000条高质量软标签数据,信息量大致可以等效成几千条甚至上万条硬标注数据,前提是教师模型足够强、输出质量足够稳定。
这也是为什么蒸馏特别适合“领域专家模型”这个目标:大模型已经把通用知识学得差不多了,学生模型不需要重新学知识点,它只需要学大模型在这些领域问题上“怎么组织回答、怎么处理不确定、怎么避开胡说八道”的行为模式。行为模式这种东西,用少量样本就能学个八九不离十。
1.3 蒸馏和微调到底差在哪
把蒸馏和微调放在一起对比,能更清楚地理解“省数据”是怎么发生的。我用一张表来说明:
| 对比维度 | 传统微调 | 知识蒸馏 |
|---|---|---|
| 学习对象 | 标准答案(硬标签) | 教师模型的输出概率分布(软标签) |
| 数据需求 | 通常需要上万条才稳定 | 高质量数据几百到几千条可启动 |
| 输出特性 | 容易“背题”,换了问法就翻车 | 学的是答题风格和边界感,泛化更好 |
| 对错误标注的容忍度 | 低,一条脏数据就能带偏 | 较高,教师模型自身有纠错能力 |
| 典型成本 | 标注人力高、清洗成本高 | 重点是算力和数据设计 |
微调解决的是“知道答案”,蒸馏解决的是“像一位专家那样作答”。领域专家模型的核心竞争力不在背诵个别法条,而在于面对真实用户那些口语化、模糊化、甚至带坑的提问时,依然能给出结构清楚、分寸得当的回答。这恰恰是蒸馏的强项。
2. 关键不是数据量,是这1000条数据的含金量
2.1 1000条数据怎么铺满一个领域
很多人第一步就栽在数据分布上。1000条看着不少,但如果全是“某法条是什么”这种单选题式问答,蒸馏出来的模型一放到真实场景里立刻就现原形——用户又不是考题机器,没人会按你训练集的样子提问。
我的做法是先画一张场景矩阵。拿法律问答来说,我把整个业务切成五个场景:法条定位与释义、案例分析、流程咨询、文书生成、风险与拒答。然后按照真实流量占比分配数据条数。比如法条定位占30%,那我就给它300条;案例分析但难度大,给它250条;流程咨询解决大部分用户需求,给200条;文书生成写起来费劲,给150条;风险拒答必须有,但样本不用多,100条足够。
每个场景下面再拆“问题类型”:法条定位里有直接问法条的、有给案情让找法条的、有比较多个法条差异的。这样每一类问题都能保证有足够的代表性样本,模型不会因为某个类型只见过两三次而完全学不会。
2.2 高质量问答对的三个特征
数据质量怎么判断?我在项目里定了三条硬标准,缺一条就返工。
第一,问题要贴近真实用户。别拿教科书里那种规范表述当问题,真实用户会问“我朋友欠我三万块不还怎么办”,不会问“民间借贷纠纷中债权人如何实现债权”。我用了一个笨但有效的办法:去知乎、贴吧、法律咨询平台,把真实提问原封不动拿回来洗一遍,而不是自己编。
第二,答案要符合领域规范且保持风格一致。同一部法律,教师在回答里一会儿说“根据XX法第几条”,一会儿说“法规规定”,学生模型学到的输出风格就会飘。我要求所有答案在开头统一结构,引用法条统一格式,结论统一放在末尾。风格一致性越强,1000条数据能发挥的效果越好。
第三,必须包含“拒答样本”。领域专家不是什么都答,遇到明显要律师介入的个案咨询,专业做法是提示风险并建议线下咨询,而不是硬编一个答案。蒸馏模型如果不专门学这部分,它会在所有问题上都“强行输出”,这是领域模型最招人烦的毛病。
2.3 训练集和验证集怎么划分才不算自欺欺人
1000条数据,我建议留120到150条做验证集,而且划分的时候不要随机抽,要分层抽。什么叫分层?就是每个场景、每种问题类型都按比例留出验证样本,保证验证集能代表整个领域分布。随机抽的验证集容易出现某类问题一条都没有,测出来分数再好都是虚的。
更关键的是防泄漏。验证集里的问题不能跟训练集里的问题在语义上过于相似。我这里举一个真实翻车案例:训练集里有一条“合同纠纷诉讼时效是几年”,验证集里放了一条“合同纠纷起诉的诉讼时效是多久”,两个问题本质上是一个问题,模型在训练时已经见过几乎一样的表述,验证分数虚高到没有参考价值。后来我加了embedding相似度去重的步骤,把训练集和验证集之间相似度超过0.85的样本全部剔掉重划,验证分数才恢复到可信水平。
3. 蒸馏实操:从教师模型到学生模型的完整链路
3.1 教师模型怎么选
教师模型是整个蒸馏项目的上限。学生模型永远不可能稳定超过教师,所以教师的能力必须至少是你目标水平的1.2倍以上。我在这个项目里用了两种方案:本地部署的qwen2.5-72b-instruct,以及商业API的顶级闭源模型。两条腿走路的原因很实际——本地模型方便批量跑不花钱,商业API质量更高但费钱,最后我用商业API跑了一遍,本地模型跑了一遍,两套软标签都保留,训练时随机选一份用,相当于给数据做了一点增强。
如果你算力有限,也不强求非要上闭源API。一个经验是:找个开源社区里公认推理能力强的大模型当教师,效果通常比用同系列的小模型自己蒸馏自己要好得多。教师和学生之间的“能力差”如果太小,蒸馏出来的信息量会很有限。
3.2 软标签生成:温度和logits的配合
软标签不是教师模型跑完一遍正常对话输出就完事的。标准的做法是在模型推理时调整温度参数T,用带温度的softmax重新计算概率分布。公式是 softmax(z / T),z是logits。温度T越高,分布越平滑,小概率token的相对差异被放大,这样学生模型能看到更多“教师原本会怎么犹豫”的信息;T越低,分布越尖锐,越接近硬标签。
我的初始经验值是T=4.0,具体做法是把教师模型在验证集上一共跑5遍,每遍用不同的随机种子,温度设置在3.0到6.0之间浮动,然后把5份概率分布取平均。这样生成的软标签比单次推理稳定得多。下面是核心的生成脚本:
import json import torch from transformers import AutoModelForCausalLM, AutoTokenizer model_name = "qwen2.5-72b-instruct" tokenizer = AutoTokenizer.from_pretrained(model_name) model = AutoModelForCausalLM.from_pretrained( model_name, torch_dtype=torch.bfloat16, device_map="auto" ) model.eval() data = [json.loads(line) for line in open("train_questions.jsonl", encoding="utf-8")] results = [] for i, item in enumerate(data): question = item["question"] reference_answer = item["reference_answer"] messages = [ {"role": "system", "content": "你是资深法律顾问,回答请基于现行法律并标明法条依据。"}, {"role": "user", "content": question} ] input_ids = tokenizer.apply_chat_template( messages, add_generation_prompt=True, return_tensors="pt" ).to(model.device) # 记录所有候选token的logits logits_list = [] max_new_tokens = 512 for _ in range(5): with torch.no_grad(): outputs = model.generate( input_ids, max_new_tokens=max_new_tokens, temperature=4.0 + (torch.rand(1).item() - 0.5) * 2, do_sample=True, output_scores=True, return_dict_in_generate=True, renormalize_logits=True ) # outputs.scores 是每个生成步的logits张量 logits_step = torch.stack(outputs.scores, dim=0) # (max_new_tokens, 1, vocab_size) logits_list.append(logits_step.cpu().float()) # 平均5份logits,然后除以温度得到软标签 logits_avg = torch.mean(torch.stack(logits_list), dim=0) soft_probs = torch.softmax(logits_avg / 4.0, dim=-1) generated_tokens = outputs.sequences[0, input_ids.shape[-1]:].tolist() results.append({ "id": i, "question": question, "reference_answer": reference_answer, "generated_token_ids": generated_tokens, "soft_probs": soft_probs.numpy().tolist() # 注意:实际存储建议用np.save更省空间 }) with open("soft_labels.json", "w", encoding="utf-8") as f: json.dump(results, f, ensure_ascii=False)注意上面代码里我循环生成了5次,但实际上在内存里保留完整的soft_probs非常占空间——7万词表的概率分布乘以512个token,1000条数据就能撑爆单机内存。实际工程里建议只保留每一生成步概率最高的Top-20 token的概率值,其他全部置零,训练时用稀疏表示来算KL散度,内存占用能降低80%以上。
3.3 学生模型选多大
学生模型的选择直接关系到部署成本和效果上限。我在项目里分别试了1.5B、3B和7B三个规格,结论是:如果目标是领域专家模型,起步建议7B,特别吃紧再降到3B,1.5B只适合做验证概念原型。
蒸馏有一个特点:学生模型越小,对教师分布的学习精度越低,但反过来说,小模型因为容量有限,反而会更集中地学习那些高频行为模式,不会东学一点西学一点。1.5B模型在1000条数据下也能学到“像模像样的回答结构”,但稍微深入一点的法律推理就露馅。7B模型则能承接教师更复杂的推理路径,输出稳定性明显上了一个台阶。如果你部署环境只有16G显存,那3B是个比较平衡的选择。
3.4 混合损失函数怎么配
蒸馏的损失函数不是只有KL散度一项。我先给公式,再解释为什么。
总的损失 L = α * T² * KL(Student_logits / T, Teacher_logits / T) + (1 - α) * CE(Student_logits, hard_label)
第一项是让学生的概率分布去贴近教师,第二项是让学生的最终预测贴合标准答案,两者混合。乘上T²是因为温度T放大了分布梯度,要除以T²才让梯度尺度回到和硬标签训练一致的量级,否则温度一高loss就直接飞了。α是两者的权重,我的经验值是0.7,也就是主要靠蒸馏信号,同时用硬标签兜底防止学生完全跟着教师偶尔的推理漂移走。
训练时用LoRA降低了显存压力。LoRA参数配的是r=16、lora_alpha=32,学习率2e-4,训练3到5个epoch,batch size设4,配合梯度累积到等效batch size 16。蒸馏训练本来就容易过拟合,epoch数宁少勿多。我做过一次对照组实验,第3个epoch验证loss还在降,到第5个epoch评测分数反而掉了,明显是开始死记硬背那1000条数据了。
训练的核心代码片段如下:
import torch import torch.nn.functional as F from transformers import AutoModelForCausalLM, AutoTokenizer, get_cosine_schedule_with_warmup from peft import LoraConfig, get_peft_model, TaskType student_name = "qwen2.5-7b-instruct" tokenizer = AutoTokenizer.from_pretrained(student_name) model = AutoModelForCausalLM.from_pretrained( student_name, torch_dtype=torch.bfloat16, device_map="auto" ) lora_config = LoraConfig( task_type=TaskType.CAUSAL_LM, r=16, lora_alpha=32, lora_dropout=0.1, target_modules=["q_proj", "k_proj", "v_proj", "o_proj", "gate_proj", "up_proj", "down_proj"] ) model = get_peft_model(model, lora_config) model.train() T = 4.0 alpha = 0.7 optimizer = torch.optim.AdamW(model.parameters(), lr=2e-4) scheduler = get_cosine_schedule_with_warmup( optimizer, num_warmup_steps=100, num_training_steps=len(train_loader) * 4 ) for epoch in range(4): for batch in train_loader: input_ids = batch["input_ids"].to(model.device) attention_mask = batch["attention_mask"].to(model.device) teacher_logits = batch["teacher_topk_logits"] # 稀疏结构 teacher_top_indices = batch["teacher_topk_indices"] hard_labels = batch["labels"].to(model.device) student_logits = model(input_ids, attention_mask=attention_mask).logits # 用TopK方式计算KL散度 student_topk = torch.gather( F.log_softmax(student_logits / T, dim=-1), dim=-1, index=teacher_top_indices ) teacher_probs_topk = F.softmax(teacher_logits / T, dim=-1) kd_loss = F.kl_div( student_topk, teacher_probs_topk, reduction="batchmean" ) * (T * T) ce_loss = F.cross_entropy( student_logits.view(-1, student_logits.size(-1)), hard_labels.view(-1), ignore_index=-100 ) loss = alpha * kd_loss + (1 - alpha) * ce_loss loss.backward() optimizer.step() scheduler.step() optimizer.zero_grad()3.5 训练参数最省心的配置项
上面那个训练循环里有几个参数是我来回调过好几轮的,这里统一给结论。
LoRA的rank,16和32我都试过,在1000条数据场景下16就够了,32反而更容易过拟合。学习率2e-4适合7B模型,7B以上可以降到1e-4。warmup步数100没什么讲究,重点是让损失先稳定爬升再进入训练状态。精度方面,NVIDIA显卡就用bf16,老显卡不支持的就用fp16,训练时把gradient_checkpointing打开,7B模型可以在24G显存下跑起来。
另外一定要记得,蒸馏训练跟普通微调不一样,验证集不能只用硬标签准确率评价,要在验证集上同时看KL散度和生成质量。我见过一个情况:验证集KL散度降得很好,但生成出来的回答句子不完整。原因是学生模型在学概率分布时,把注意力过度放在了高频token上,低频的句尾标点和承接词没学到。解决办法是给KL散度加一个mask,只计算在教师输出中概率高于0.01的token位置的损失,丢掉那些纯噪声的低概率位置。
4. 效果怎么验证才不算“自嗨”
4.1 三个维度缺一不可
蒸馏项目的评测不能只看“答对率”。我设计了三个维度,每个维度都有独立的评测集。
第一个维度是领域能力。我单独攒了一个150题的评测集,覆盖前面说的五个场景,全部是真实用户问题,答案由资深律师人工审核打分。这个评测集在训练前和训练后都跑一遍,会得到一个能力提升的基线分。
第二个维度是通用能力回退。领域模型最担心的就是学会了法律、忘掉了常识。我抽了C-Eval里的常识、逻辑、数学三个子集,共100题,蒸馏前后各跑一次。比如模型如果因为强化法律表达而在普通数学题上也强行“根据相关规定作答”,那就说明通用能力受损。
第三个维度是稳定性。同一个问题问10遍,统计答案的词汇重合度。蒸馏出来的学生模型如果温度设成0,输出必须保持几乎一致;如果输出抖动很大,说明概率分布还没训稳。
我在项目里最常用的一句话是:不要用选择题的思维去评价生成模型的进步。领域专家模型的价值在于长文本回答的条理性、引用准确性和对不确定问题的处理,这些指标都要靠人工或半人工的评估框架。
4.2 和纯微调做一个A/B对比
为了让老板心服口服,我跑了一组对比实验:同1000条数据,一份做纯微调(只用硬标签),一份做蒸馏(软标签+硬标签),学生模型相同,训练轮数和LoRA配置相同。最终评测结果如下(简化后的示意数据):
| 评估维度 | 纯微调 | 蒸馏 | 提升 |
|---|---|---|---|
| 领域问答准确率 | 71.2% | 78.6% | +7.4% |
| 法条引用正确率 | 62.4% | 74.3% | +11.9% |
| 常识子集回退率 | -3.8% | -1.2% | +2.6% |
| 10次回答词汇重合度 | 68% | 86% | +18% |
法条引用正确率提升了近12个百分点,是蒸馏最明显的收益,原因就是软标签让模型学到了教师“在不确定时先给分析路径再下结论”的习惯,而不是硬编码地背法条编号。纯微调模型遇到没见过的问法就容易全错,蒸馏模型则保留了一条“退路”。
如果你没有人力做大规模人工评测,一个简化的替代方案是把教师的回答、微调模型的回答、蒸馏模型的回答三个放一起,丢给一个更强的模型按规则打分。虽然不是百分之百准确,但能快速筛出明显差距。
5. 实战中踩过的坑和排查方法
5.1 数据泄漏带来的“假高分”
这是我第一个踩的坑,上面在验证集划分里已经提过。补充一个当时排查的经过:训练到第2个epoch时,验证准确率突然从70%跳到88%,我很兴奋,赶紧让团队庆祝了一下。结果第二天用真实用户问题一测,退回72%。分析后发现,验证集和训练集有接近一成的题目是语义重复的,模型相当于开卷考试。后来加了embedding相似度去重才把分数打回原形。记住一句话:评测分数高得越突兀,越要先怀疑数据泄漏。
5.2 训练loss降了但生成质量越来越差
第二个坑出现在训练后期。loss一路降到很低,但生成的回答开始出现重复片段,比如“综上所述综上所述综上所述”。排查后确定是过拟合:模型把训练数据中那些高概率token路径背下来了,开始进入循环生成。解决方案有三个,按优先级排序:降低epoch数、增大LoRA dropout到0.15、把蒸馏权重α从0.7调到0.8。这三个动作做完,重复问题基本消失。
5.3 温度T调太大的后果
我最初以为温度越高,软标签信息量越大,于是试过T=8.0。结果训练出来的学生模型什么问题都回答得模棱两可,连“合同是否有效”这种本来可以给出明确结论的问题,也输出一堆“可能、或许、视情况而定”。原因在于温度过高导致概率分布过于平均,教师原本包含的“置信度区分”信息被抹平了。学生学不到“这件事教师很有把握、那件事教师很犹豫”的边界感。T的合适区间在2.0到5.0,具体需要在小验证集上扫一遍。
5.4 常见问题速查表
| 症状 | 最可能的原因 | 排查与解决 |
|---|---|---|
| 验证loss很低但生成乱码 | 软标签只存了Top-1概率 | 检查软标签稀疏格式,保留Top-20以上 |
| 输出重复片段 | 过拟合 | 减epoch、增大dropout、提高α |
| 参考答案完全错误 | 教师模型本身能力不足 | 换更强的教师,或两三个模型输出投票 |
| 法条编号张冠李戴 | 训练数据里法条引用不统一 | 统一教师系统提示词要求格式,清洗数据 |
| 训练时显存OOM | 梯度累积不生效 | 检查gradient_checkpointing是否开启 |
| 问答过于啰嗦 | T调太高 | 降到2.0~4.0重新蒸馏 |
6. 聊聊我个人的真实体会
项目收尾之后回头看,我越来越觉得“1000条数据蒸馏出一个领域专家模型”这句话被很多人理解得太功利了。它确实成立,但成立的前提是你愿意把功夫花在数据构造和验证设计上,而不是花在“跑模型”本身。跑模型的时间可能只占整个项目周期的三成,剩下七成都在折腾数据怎么提问、怎么清洗、怎么防止泄漏、怎么验证。如果你只是随手搜1000条QA硬灌进去,出来的东西大概率只能当玩具;反过来,如果你愿意静下心来打磨数据、调好软标签、设计一套不骗自己的评估体系,这笔投资比请人标一万条数据划算得多。
顺着这个方向,后续我打算尝试在技能蒸馏(skill distillation)上继续延展——同样是蒸馏思路,但目标从“学会回答”变成“学会调用工具和规划步骤”,这对数据的利用效率更高。等项目跑通了,我再单独写一篇实操复盘。