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、很多开源小推理模型都是这套流程

流程:

  1. 准备一批 prompt(业务数据 / 公开数据集);
  2. 调用教师 API,生成完整回答(包含 CoT 思考过程);
  3. 清洗、过滤、去重,做成 prompt→response SFT 数据集;
  4. 拿这份数据集对学生底座做 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)

Logo

欢迎加入DeepSeek 技术社区。在这里,你可以找到志同道合的朋友,共同探索AI技术的奥秘。

更多推荐