17.知识蒸馏
1.知识蒸馏(Distillation)
用大模型(教师 Teacher),去训练一个更小的模型(学生 Student)
把大模型的思考方式、输出风格迁移给小模型,得到一个更快、更省显存的轻量模型
类比:学霸(大教师)不直接把脑子复制给学徒,而是大量做题,把自己的思考过程、倾向性、答案写成笔记;
学徒拿着这套笔记重新学习,学徒本身规模很小,但学会学霸的做题习惯。
2. 核心概念
1)教师模型 Teacher:大参数量、效果好、推理慢、吃显存,蒸馏阶段一般冻结,只做推理输出样本,不参与参数更新。
2)学生模型 Student:参数量小,新模型,要训练更新参数,目标尽量复刻老师的行为。
3)软标签 Soft Label / 暗知识 Dark Knowledge
- 普通训练只用硬标签:只有标准答案。
- 蒸馏用老师输出的完整概率分布:不只看最终答案,还看老师认为各个备选答案的可能性,里面藏着模型的推理倾向,叫暗知识。
4)蒸馏损失:学生一边对齐真实标准答案,一边对齐老师输出的概率分布,两者加权一起训练学生模型。
3. 标签
标签就是给每个候选 token 的目标分数
3.1 硬标签
目标值只有 0 或者 1。
例子(正确 token 是猫)
猫:1 狗:0 鸟:0
目标:输出尽量变成 [1,0,0]。 这里的标签就是一组目标分数,非 0 即 1。
3.2 软标签
目标是一堆小数概率,不是只有 0 和 1
猫:0.88 狗:0.11 鸟:0.01
目标:输出尽量贴近 [0.88,0.11,0.01]。 这组小数,就是软标签。
4. LLM 行业里两种常见蒸馏
4.1 Logit 蒸馏(原生知识蒸馏)
学生直接学老师输出 token 的完整概率分布,对训练算力要求高,效果上限高。
LLM 每一步要选下一个 token:词汇表有几万候选词,模型给每一个候选词打一个原始分数,这个分数就叫 logit。
4.2 输出蒸馏(行业最常用,也叫 SFT 蒸馏)
拿一堆 prompt,调用教师 API 拿到回答,把prompt+teacher回答做成 SFT 数据集,拿这份数据集微调学生模型。
现在很多开源小模型就是这么做的:调用 DeepSeek、GPT‑4 大量生成样本,拿样本训自己的 7B/14B 小底座。 这是学输出文本,不学原始 logit 概率,工程简单,但上限低于原生 logit 蒸馏。
SFT Supervised Fine‑Tuning 监督微调,用一批「用户问题 + 理想回答」的标注数据,继续微调预训练大模型,让模型学会按对话格式输出。
5.蒸馏 vs 量化
|
对比点 |
蒸馏 Distillation |
量化 Quantization |
|
做什么 |
训练出新的更小模型 ,知识迁移 |
不改模型结构 ,把参数的比特位降低(FP16→INT4/INT8) |
|
是否重新训练 |
必须训练学生模型 |
PTQ 后量化不需要训练;QAT 量化感知训练才要训 |
|
改变参数量 |
可以从 70B→7B,参数量直接变小 |
参数量不变,只是每个数字存得更粗糙 |
|
目的 |
把大模型能力教给小模型 |
同一个模型,压缩体积、加速推理 |
|
关系 |
经常组合:先蒸馏得到小模型,再对小模型做 INT4 量化,进一步压缩部署 |
举个端侧 AI 部署链路例子: 70B 大模型 →【蒸馏】→7B 学生模型 →【INT4 量化】→本地端侧可跑的小权重文件。
6.输出蒸馏介绍
输出蒸馏也叫硬样本蒸馏、API 蒸馏,DeepSeek‑R1‑Distill、很多开源小推理模型都是这套流程
流程:
- 准备一批 prompt(业务数据 / 公开数据集);
- 调用教师 API,生成完整回答(包含 CoT 思考过程);
- 清洗、过滤、去重,做成 prompt→response SFT 数据集;
- 拿这份数据集对学生底座做 SFT 微调。
输出蒸馏代码
# ====================== 硬样本蒸馏链路相关代码 ======================
import json
import os
from openai import OpenAI
# ---------- 配置:教师API(DeepSeek / 豆包都兼容openai格式) ----------
client = OpenAI(
api_key=os.environ["DEEPSEEK_API_KEY"],
base_url="https://api.deepseek.com/v1"
)
def load_my_business_prompt():
"""加载业务prompt列表,可以从txt/json读取,这里模拟"""
return [
"写一段python快速排序",
"解释什么是logit",
"简单讲什么是知识蒸馏",
"写rust实现二分查找"
]
prompts = load_my_business_prompt()
distill_data = []
# ---------- 步骤1:调用教师API生成蒸馏硬样本 ----------
for q in prompts:
resp = client.chat.completions.create(
messages=[{"role":"user", "content": q}],
model="deepseek-chat",
temperature=0.7
)
ans = resp.choices[0].message.content
distill_data.append({
"messages": [
{"role":"user", "content": q},
{"role":"assistant", "content": ans}
]
})
# ---------- 步骤2:简单清洗过滤,保存数据集json ----------
def filter_sample(sample):
"""简单过滤逻辑:回答为空、过短直接丢弃,业务可再加更多校验"""
assistant_text = sample["messages"][1]["content"]
if not assistant_text or len(assistant_text.strip()) < 10:
return False
return True
distill_data = [s for s in distill_data if filter_sample(s)]
with open("hard_distill_dataset.json", "w", encoding="utf‑8") as f:
json.dump(distill_data, f, ensure_ascii=False, indent=2)
print(f"生成蒸馏数据集完成,共 {len(distill_data)} 条,保存 hard_distill_dataset.json")
# ---------- 步骤3:调用 LLaMA‑Factory 执行SFT训练(硬样本蒸馏) ----------
"""
注意:llama‑factory是命令行工具,不在python内直接import跑;
两种方式:
A) 脚本调用subprocess拉起llamafactory-cli命令(下面示例)
B) 终端直接执行:llamafactory-cli train sft_distill.yaml
"""
import subprocess
cmd = [
"llamafactory-cli",
"train",
"sft_distill.yaml"
]
print("开始执行LLaMA‑Factory SFT蒸馏训练...")
subprocess.run(cmd, check=True)
print("训练完成,蒸馏后的学生模型输出在配置的 output_dir 目录")
输出学生模型位置配置在sft_distill.yaml
### sft_distill.yaml
model_name_or_path: Qwen/Qwen2‑1.5B‑Instruct
dataset: hard_distill_dataset
dataset_dir: ./
template: qwen
finetuning_type: lora
lora_target: all
stage: sft
do_train: true
overwrite_cache: true
cutoff_len: 1024
max_samples: null
per_device_train_batch_size: 2
gradient_accumulation_steps: 2
learning_rate: 5e‑5
num_train_epochs: 3
logging_steps: 10
save_steps: -1
save_total_limit: 2
output_dir: ./student_hard_distill_lora
bf16: true
fp16: false
optim: paged_adamw_8bit
val_size: 0.0
report_to: none
训练结束后,导出完整权重代码
# 训练完,合并LoRA,导出完整学生模型
cmd_merge = [
"llamafactory-cli",
"export",
"--model_name_or_path", "Qwen/Qwen2‑1.5B‑Instruct",
"--adapter_name_or_path", "./student_hard_distill_lora",
"--export_dir", "./student_distilled_full",
"--export_size", "2"
]
subprocess.run(cmd_merge, check=True)
更多推荐



所有评论(0)