logo

Model Distillation 实战:把大模型能力压缩下来

Published on

Model Distillation 实战:把大模型能力压缩下来

知识蒸馏听起来很像一句口号:“让大模型教小模型”。真正动手时,你会发现困难并不在最后那条训练命令,而在于三个问题:教师到底教什么,哪些答案值得学生学习,怎样证明学生没有把教师的错误一起继承下来。

我做蒸馏实验时,最初也犯过一个很自然的错误:批量调用大模型,让它生成几万条答案,然后直接拿去训练。训练集很快就有了,Loss 也下降得很漂亮,可学生模型上线后表现并不稳定。后来回头检查,数据里有重复问题、格式不一致、教师幻觉,还有一些答案虽然写得通顺,却没有真正回答任务。

蒸馏不是把“大模型输出”复制到文件里,而是把教师的能力变成经过筛选、可测量、可回滚的训练资产。下面以“用户问题分类”为例,走一遍比较稳妥的流程。

一、先定义学生模型的任务边界

假设我们要把客服问题分成 8 类:退款、物流、账户、发票、优惠、售后、投诉和其他。输入是一句话,输出必须是固定 JSON:

{
  "label": "物流",
  "confidence": 0.91,
  "reason": "问题询问订单配送状态"
}

任务越明确,蒸馏越容易成功。不要一开始就让学生模型“学会客服”,因为这个目标包含分类、检索、对话、情绪判断和业务决策,无法知道失败究竟发生在哪一层。

先定义任务规格:

const taskSpec = {
  labels: ['退款', '物流', '账户', '发票', '优惠', '售后', '投诉', '其他'],
  outputSchema: 'ClassificationResultV1',
  maxInputTokens: 256,
  maxOutputTokens: 80,
}

同时准备一小批人工标注的基准集。哪怕只有 500 条,也要覆盖真实表达、错别字、口语、混合意图和应该归为“其他”的问题。这批数据不让教师参与生成,后面专门用来检验蒸馏是否真的有效。

二、让教师输出稳定,而不是只追求聪明

教师模型的 Prompt 要把任务、标签定义、输出格式和拒答条件写清楚:

你是一个客服意图分类器。
只能从以下标签中选择一个:退款、物流、账户、发票、优惠、售后、投诉、其他。
如果问题包含多个意图,选择用户最主要的诉求。
只输出 JSON,不要输出 Markdown,不要增加额外字段。

批量生成时还要保存教师版本、Prompt 版本、采样参数和请求状态。否则一个月后你看到一批训练数据,却不知道它是哪个模型生成的,也无法解释为什么重新生成的结果不一致。

type TeacherSample = {
  sampleId: string
  question: string
  rawOutput: string
  teacherModel: string
  promptVersion: string
  temperature: number
  status: 'generated' | 'failed' | 'reviewed'
}

分类任务通常可以把 temperature 设低一些,让标签更加稳定;开放式生成则需要根据任务决定。无论参数怎么选,都不要把一次生成结果当成真理。

三、数据清洗比训练更重要

1. 先做结构校验

教师输出必须先经过 JSON Schema 校验:

function validateSample(sample: unknown): sample is ClassificationResult {
  return (
    isObject(sample) &&
    typeof sample.label === 'string' &&
    labels.includes(sample.label) &&
    typeof sample.reason === 'string'
  )
}

解析失败的样本不要偷偷修补后继续训练。应该记录失败原因,观察是 Prompt 不稳定、模型输出异常,还是问题本身超出了任务范围。

2. 去重和去污染

训练数据里大量相同问题,会让模型记住模板而不是学会能力。可以先做文本规范化,再使用哈希或相似度去重。还要检查训练集和测试集是否有近似样本,避免评测结果虚高。

const key = normalize(question).toLowerCase()
if (seen.has(key)) return
seen.add(key)

严格来说,语义近似去重比字符串去重更难,但至少要先处理完全重复和模板化重复。工程质量经常不是来自一个复杂算法,而是来自不忽略这些基础清理。

3. 教师一致性检查

可以让同一个教师用不同 Prompt 或不同采样参数生成两次结果。如果标签经常冲突,这个问题就不适合直接蒸馏,或者标签定义需要重写。对于高风险分类,还应该用第二个模型或人工抽样复核。

同一问题 → 教师 A:退款
         → 教师 B:售后
         → 进入人工复核,而不是随便选一个

蒸馏数据宁可少一点,也不要用大量互相矛盾的样本把学生训练得犹豫不决。

四、两种蒸馏目标

1. 硬标签蒸馏

最简单的方式是把教师最后选出的标签当作普通监督数据:

问题 → 退款
问题 → 物流
问题 → 投诉

它适合分类、抽取和格式化任务,训练实现简单,也容易解释。缺点是教师对“退款”和“售后”有多大把握,学生看不到这些信息。

2. 软目标蒸馏

如果能够拿到教师对各个类别的概率分布,就可以让学生学习更丰富的判断边界。常用做法是使用温度参数 T 平滑分布:

p_teacher = softmax(logits_teacher / T)
p_student = softmax(logits_student / T)
蒸馏损失 = KL(p_teacher || p_student) × T²

温度提高后,原本很小的类别概率会变得更明显,学生能学到“这个问题虽然主要像物流,也有一点售后倾向”。但软目标只有在教师概率可信时才有价值。很多 API 只返回一个最终答案,没有真正的完整概率分布,这时不要假装自己拿到了软标签,可以使用多个候选输出近似,但必须标注这是工程折中。

实际训练经常混合真实标签和教师标签:

总损失 = 0.6 × 真实标签交叉熵
       + 0.4 × 教师输出蒸馏损失

比例不是固定答案,要通过验证集比较。真实标签质量高时,应该让它占更大权重;教师覆盖长尾表达时,可以适当提高蒸馏数据的作用。

五、一个最小训练骨架

下面是接近 PyTorch 的伪代码,重点是表达训练关系,不绑定某个具体模型:

for batch in train_loader:
    input_ids = batch["input_ids"].to(device)
    labels = batch["labels"].to(device)
    teacher_probs = batch["teacher_probs"].to(device)

    student_logits = student(input_ids).logits
    hard_loss = cross_entropy(student_logits, labels)

    student_log_probs = log_softmax(student_logits / temperature, dim=-1)
    soft_loss = kl_div(
        student_log_probs,
        teacher_probs,
        reduction="batchmean",
    ) * temperature ** 2

    loss = alpha * hard_loss + (1 - alpha) * soft_loss
    optimizer.zero_grad()
    loss.backward()
    optimizer.step()

训练时要记录训练集、验证集和基准集的指标,不要只盯着 Loss。Loss 下降而验证集不涨,通常说明模型在记忆训练样本;验证集上涨但关键少数类下降,则可能需要重新调整采样和损失权重。

六、评测蒸馏是否值得

至少比较四个版本:规则基线、原始学生模型、蒸馏后的学生模型、教师模型。指标包括总体 F1、每个类别的召回率、拒答准确率、P95 延迟、内存占用和单请求成本。

尤其要看长尾类别。一个模型总体准确率从 91% 提升到 94%,听起来不错,但如果“投诉”类别召回从 88% 降到 65%,对客服系统可能是倒退。

还要做错误分析:

  • 教师和学生都错:可能是标签定义或问题本身有歧义。
  • 教师对、学生错:学生容量或训练数据不足。
  • 教师错、学生也学错:蒸馏数据污染。
  • 学生答对但格式错:需要更强的输出约束或格式训练。

这四类错误的修复方向完全不同。只看一个总分,无法告诉你下一步应该补数据、改标签还是换模型。

七、把不确定性交给大模型

蒸馏模型不需要假装无所不知。上线时可以设置置信度阈值:高置信度的简单问题由学生处理,低置信度或高风险问题升级给大模型或人工。

学生模型
  ├─ 高置信度 + 低风险 → 直接返回
  └─ 低置信度 / 冲突 / 高风险 → 升级处理

阈值必须在独立验证集上校准,不能凭感觉写成 0.8。同时保留路由原因,后续才能知道升级比例是不是过高,或者某一类问题的置信度普遍虚高。

总结:蒸馏是在复制能力,也是在复制责任

这次实践让我最深的感受是,蒸馏最有价值的产物不是一个更小的权重文件,而是一套经过理解的数据资产。你必须知道教师教了什么、哪些内容被排除了、学生在哪些地方失败,以及当数据被撤回时如何找到它。

大模型适合做能力上限和复杂长尾,小模型适合做高频、稳定、低延迟的任务。蒸馏把两者连接起来,但不会替你消除幻觉、偏差和数据合规问题。教师不可靠,学生只会更便宜地犯错;评测不完整,部署以后只会更快地暴露问题。

所以别把蒸馏理解成“让模型变小”的魔法。它更像一位老师整理自己的讲义:先定义课程,再筛选例题,批改错误,安排考试,最后让学生在不会的地方举手求助。这样训练出来的小模型,才是真的能在生产环境里承担任务,而不是只在 Demo 里显得聪明。

🤪 您也可以编辑此页: