知识蒸馏:把大模型的能力“搬“进小模型,原理与一次完整实战

知识蒸馏:把大模型的能力"搬"进小模型,原理与一次完整实战

"硬标签只告诉你答案是猫;软标签还会告诉你,它有多像狐狸、几乎不像汽车。多出来的这部分信息,就是蒸馏真正要教的东西。"

🔥 摘要:量化改的是数值精度,剪枝改的是结构连接,而蒸馏改的是知识的来源——让一个小模型去学大模型输出的概率分布,而不是只学"标准答案"。本文从"软标签到底在教什么"讲起,用一组可复现的数字说清温度 $T$ 的作用;拆开蒸馏的三种范式(离线/在线/自蒸馏)与两条监督通路(Logits 蒸馏 / 特征蒸馏);最后给一份 PyTorch 完整可跑的实战代码:从头训一个 Teacher,再把它"蒸"进一个参数量小一个量级的 Student,并与"从零训练同款小模型"做精度对比。看完你会知道蒸馏什么时候值、什么时候纯属折腾。

🎯 阅读收益:① 真正理解"暗知识"和温度参数在数学上做了什么;② 掌握蒸馏损失的标准写法与 $T^2$ 这个容易被漏掉的系数;③ 分清三种范式与两条监督通路的适用场景;④ 拿到一份可直接运行、自带对比实验的完整 PyTorch 代码;⑤ 避开 7 个真实训练里会把效果做崩的坑。

⚠️ 说明:文中温度为演示用的示意 logits,数值由 softmax(z/T) 直接算出,可自行复现;实战代码为教学简化版,使用 CIFAR-10 与自制小网络,目的是让流程可在一张消费级显卡(甚至 CPU)上跑完,工业级蒸馏需按你的模型与数据调整结构与超参。


一、为什么量化、剪枝之后,还需要蒸馏

模型压缩有三条路线,它们改的东西完全不同:

在这里插入图片描述

路线 改什么 参数量变化 结构变化 需要训练吗
量化 数值精度(FP16→INT8/INT4) 不变 不变 通常不需要(PTQ)
剪枝 结构连接(删权重/通道/层) 减少 改变 剪后一般需要微调
蒸馏 知识来源(跟谁学) 任意设计 可完全重设计 需要完整训练

三者的关键差别在于自由度:

  • 量化和剪枝都是在原有模型上做减法,天花板被原模型锁死:7B 量化后还是 7B 的骨架,能力上限不会超过原模型。
  • 蒸馏是重新造一个学生:学生的结构你可以完全自定义(更浅、更窄、换算子、换注意力实现),只要它最终能模仿老师的输出分布。

所以蒸馏常与量化组合使用:先蒸馏出一个小骨架,再量化,这在端侧部署里是最常见的路径。

还有一个容易被忽略的价值:蒸馏可以跨结构迁移。你可以把一个 Transformer 老师的能力,蒸进一个更适合移动端的小 CNN 或混合结构里——这是量化和剪枝都做不到的事。

二、软标签到底在教什么

2.1 硬标签丢掉了什么

假设一张图片的真实标签是「猫」,模型输出的 logits 是:

z = [5.0, 3.0, 1.5, 0.5, -1.0]   # 对应 [猫, 狗, 狐狸, 狼, 汽车]

硬标签(one-hot)是:[1, 0, 0, 0, 0]

它只说了一件事:这是猫。至于"它有点像狗,更像狐狸,完全不像汽车"——全部丢掉了。

而模型真正学到的知识恰恰藏在这些非正确类别的相对大小里,这被称为暗知识(Dark Knowledge)。

2.2 温度 $T$ 做了什么

蒸馏的核心操作,是在 softmax 里引入温度 $T$:

$$
p_i = \frac{\exp(z_i / T)}{\sum_j \exp(z_j / T)}
$$

  • $T = 1$:就是标准 softmax;
  • $T > 1$:分布被"抹平",非正确类别的概率被放大;
  • $T \to \infty$:趋向均匀分布。

用上面那组 logits 实算一下,效果非常直观:

在这里插入图片描述

类别 T=1 T=2 T=4 T=8
猫 0.848 0.589 0.389 0.288
狗 0.115 0.217 0.236 0.225
狐狸 0.026 0.102 0.162 0.186
狼 0.009 0.062 0.126 0.164
汽车 0.002 0.029 0.087 0.136

看最右边两列:$T=1$ 时,"狼"和"汽车"的概率都小到几乎可以忽略(0.009 和 0.002),模型学不到任何关于它们的信息;$T=8$ 时,二者的差异被放大到 0.164 vs 0.136,"虽然都不对,但狼比汽车靠谱得多"这个信息就传过去了。

一句话:温度不是让模型"更不确定",而是把老师脑子里类别之间的相对关系暴露出来给学生看。

2.3 蒸馏损失的标准写法

$$
\mathcal{L} = \alpha \cdot T^2 \cdot \mathrm{KL}\left(p^{T}{teacher} ,|, p^{T}\right) ;+; (1-\alpha) \cdot \mathrm{CE}(y,; p_{student})
$$

三点必须注意:

  1. 老师和学生要用同一个 $T$。只在老师侧除 $T$、学生侧不除,是新手最常见的错误,损失会直接跑偏。
  2. $T^2$ 不能漏。因为 $\mathrm{KL}$ 项里除了 $T$,梯度会被缩小约 $1/T^2$ 倍,乘回 $T^2$ 是为了让软标签损失和硬标签损失在量级上可比,这样 $\alpha$ 才有意义。漏掉它,$\alpha$ 怎么调都不对。
  3. 硬标签项用 $T=1$。学生的最终输出必须是"真实温度"下的分布,否则 eval 时会发现概率全被抹平、准确率暴跌。

公式里的 $\alpha$ 控制"跟老师学"和"跟标准答案学"的权重,常用取值在 0.3 ~ 0.7;温度 $T$ 常用 2 ~ 8,任务越复杂、类别越多,$T$ 可以适当调大。

三、三种范式,两条通路

3.1 三种范式

范式 做法 优点 代价
离线蒸馏 先训好老师并冻结,再训学生 最简单、最常用、可复用老师 需要事先有一个好老师
在线蒸馏 老师和学生同时训练 不需要预训练老师,可互相促进 训练更复杂,显存占用更高
自蒸馏 模型自己教自己(深层教浅层 / 历史教当前) 零额外模型成本 提升幅度通常有限

入门和业务落地,选离线蒸馏就够了。

3.2 两条监督通路

在这里插入图片描述

mermaid diagram

  • Logits 蒸馏:只对齐最终输出分布。实现简单、与学生结构完全解耦,是首选。
  • 特征蒸馏:额外对齐中间层特征(或注意力图)。监督信号更密,学生上限更高,但要处理层数与维度对齐(常见的做法是给学生中间层加一个 Linear 投影到老师的维度)。

选型建议:先做 Logits 蒸馏,跑通并拿到基线;效果不够再加特征蒸馏。一上来就上特征对齐,很容易在维度对齐上耗掉大量时间而看不到收益。

四、完整实战:把大模型"蒸"进小模型

4.1 实验设计

为了让你在一张消费级显卡上就能跑完,这里不用预训练大模型,而是自制一对师生:

角色 结构 通道宽度 参数量级
Teacher 3 层 CNN 64 / 128 / 256 约 130 万
Student 同样的 3 层 CNN 16 / 32 / 64 约 8 万(约 1/16)

对比三组结果:

  1. Teacher 从零训练;
  2. Student 从零训练(对照组);
  3. Student 用蒸馏训练(实验组)。

关键看 3 比 2 高多少——这才是蒸馏真正的收益,而不是"学生能不能追上老师"。

4.2 环境准备

pip install torch torchvision
# 有 NVIDIA 显卡时建议装对应 CUDA 版本的 torch;没有也能用 CPU 跑完(会慢一些)
python -c "import torch; print(torch.cuda.is_available())"

4.3 完整代码

# -*- coding: utf-8 -*-
"""
kd_cifar10.py —— 知识蒸馏完整可跑示例
三组对比:Teacher 从零 / Student 从零 / Student 蒸馏
运行:python kd_cifar10.py
"""
import torch
import torch.nn as nn
import torch.nn.functional as F
from torch.utils.data import DataLoader
from torchvision import datasets, transforms

DEVICE = "cuda" if torch.cuda.is_available() else "cpu"
BATCH, EPOCHS, LR = 128, 5, 1e-3


# ---------- 模型:同一个骨架,靠通道宽度控制容量 ----------
def conv_block(cin, cout):
    return nn.Sequential(
        nn.Conv2d(cin, cout, 3, padding=1),
        nn.BatchNorm2d(cout),
        nn.ReLU(inplace=True),
        nn.MaxPool2d(2),
    )


class Net(nn.Module):
    """CIFAR-10 是 32x32,经过 3 次 pool 后变成 4x4"""

    def __init__(self, widths=(64, 128, 256), num_classes=10):
        super().__init__()
        w1, w2, w3 = widths
        self.features = nn.Sequential(conv_block(3, w1), conv_block(w1, w2), conv_block(w2, w3))
        self.head = nn.Linear(w3 * 4 * 4, num_classes)

    def forward(self, x):
        return self.head(torch.flatten(self.features(x), 1))


# ---------- 数据 ----------
def loaders():
    tf = transforms.Compose([
        transforms.ToTensor(),
        transforms.Normalize((0.4914, 0.4822, 0.4465), (0.2470, 0.2435, 0.2616)),
    ])
    tr = datasets.CIFAR10(root="./data", train=True, download=True, transform=tf)
    te = datasets.CIFAR10(root="./data", train=False, download=True, transform=tf)
    return (DataLoader(tr, BATCH, shuffle=True, num_workers=2),
            DataLoader(te, 256, shuffle=False, num_workers=2))


# ---------- 评估 ----------
@torch.no_grad()
def evaluate(model, loader):
    model.eval()
    correct = total = 0
    for x, y in loader:
        x, y = x.to(DEVICE), y.to(DEVICE)
        correct += (model(x).argmax(1) == y).sum().item()
        total += y.size(0)
    return correct / total


# ---------- 普通训练(只学硬标签) ----------
def train_plain(model, train_loader, test_loader, epochs=EPOCHS):
    model.to(DEVICE)
    opt = torch.optim.Adam(model.parameters(), lr=LR)
    for ep in range(epochs):
        model.train()
        for x, y in train_loader:
            x, y = x.to(DEVICE), y.to(DEVICE)
            loss = F.cross_entropy(model(x), y)
            opt.zero_grad()
            loss.backward()
            opt.step()
        print("  epoch %d | CE %.4f | test acc %.4f"
              % (ep + 1, loss.item(), evaluate(model, test_loader)))
    return model


# ---------- 蒸馏训练(硬标签 + 软标签) ----------
def train_distill(teacher, student, train_loader, test_loader,
                  epochs=EPOCHS, T=4.0, alpha=0.7):
    teacher.to(DEVICE).eval()
    for p in teacher.parameters():
        p.requires_grad_(False)          # 冻结老师:省显存也省算力
    student.to(DEVICE)
    opt = torch.optim.Adam(student.parameters(), lr=LR)

    for ep in range(epochs):
        student.train()
        for x, y in train_loader:
            x, y = x.to(DEVICE), y.to(DEVICE)

            with torch.no_grad():
                t_logits = teacher(x)                       # 老师不参与反传
            s_logits = student(x)

            # ① 软标签损失:师生两侧同除 T,再乘回 T^2 还原梯度量级
            kd = F.kl_div(F.log_softmax(s_logits / T, dim=1),
                          F.softmax(t_logits / T, dim=1),
                          reduction="batchmean") * (T * T)
            # ② 硬标签损失:学生必须在 T=1 的真实分布上对齐
            ce = F.cross_entropy(s_logits, y)

            loss = alpha * kd + (1 - alpha) * ce
            opt.zero_grad()
            loss.backward()
            opt.step()

        print("  epoch %d | KD %.4f | CE %.4f | test acc %.4f"
              % (ep + 1, kd.item(), ce.item(), evaluate(student, test_loader)))
    return student


if __name__ == "__main__":
    torch.manual_seed(42)
    train_loader, test_loader = loaders()

    teacher = Net(widths=(64, 128, 256))
    scratch = Net(widths=(16, 32, 64))
    kd_stu = Net(widths=(16, 32, 64))

    tp = sum(p.numel() for p in teacher.parameters())
    sp = sum(p.numel() for p in kd_stu.parameters())
    print("Teacher: {:,} 参数 | Student: {:,} 参数 | 压缩约 {:.1f}x\n".format(tp, sp, tp / sp))

    print("[1/3] 训练 Teacher(从零)")
    train_plain(teacher, train_loader, test_loader)
    acc_teacher = evaluate(teacher, test_loader)

    print("\n[2/3] 训练 Student(从零,对照组)")
    train_plain(scratch, train_loader, test_loader)
    acc_scratch = evaluate(scratch, test_loader)

    print("\n[3/3] 训练 Student(蒸馏)")
    train_distill(teacher, kd_stu, train_loader, test_loader)
    acc_kd = evaluate(kd_stu, test_loader)

    print("\n" + "=" * 46)
    print("最终结果(CIFAR-10 test accuracy)")
    print("=" * 46)
    print("  Teacher              : %.4f" % acc_teacher)
    print("  Student 从零训练     : %.4f" % acc_scratch)
    print("  Student 蒸馏训练     : %.4f" % acc_kd)
    print("  蒸馏带来的增益       : %+.2f 个百分点" % ((acc_kd - acc_scratch) * 100))

4.4 怎么读这个结果

正常情况你会看到:

  • Student 从零训练 落后 Teacher 若干个点(小模型容量不足);
  • Student 蒸馏训练 明显高于 Student 从零训练——这两个数的差,就是蒸馏的净收益;
  • 蒸馏后的 Student 通常仍略低于 Teacher(蒸馏不是免费的午餐),但参数量只有它的几十分之一。

想做消融的话,按这个顺序调,每次只动一个变量:

想验证什么 怎么改
温度的影响 T 分别取 1 / 2 / 4 / 8 / 16
软硬标签配比 alpha 分别取 0.1 / 0.3 / 0.5 / 0.7 / 0.9
容量差距的影响 把 Student 宽度改成 (8,16,32) 或 (32,64,128)
训练量的影响 EPOCHS 改成 10 / 20,看增益是放大还是收敛

特别建议做一次 T=1 的对照:你会发现去掉温度之后,蒸馏的收益会大幅缩水——这比任何文字解释都更能说明"暗知识"是什么。

五、踩坑清单

  1. 学生侧忘记除 $T$。只在老师侧除温度,KL 直接跑偏,损失下不去。师生必须用同一个 $T$。
  2. 漏掉 $T^2$。软标签损失被稀释到几乎不起作用,alpha 怎么调都像没加蒸馏。
  3. 推理时还带着蒸馏温度。evaluate 用的是 model(x) 的原始 logits($T=1$),如果你在推理时也除 $T$,准确率会明显偏低。
  4. 老师没冻结、没切 eval()。BatchNorm 的统计量会跟着当前 batch 漂移,Dropout 会随机丢神经元,老师的"软标签"每天都在变;同时白白浪费一份前向+反向的显存。
  5. alpha 设成 1.0。完全不学真实标签,老师犯的错会被学生原样继承。保留一部分 CE 是重要的"纠偏项"。
  6. 师生容量差距过大。这是蒸馏里最反直觉的坑:老师太强时,学生反而学不好(分布过于尖锐,学生拟合不动)。差距大时,可以引入一个中等规模的助教模型做中间过渡。
  7. 用了错误的数据。蒸馏最好用老师训练时同分布的数据,或者干脆用原始训练集;用与老师知识无关的数据去蒸,等于让老师在自己没见过的领域瞎编。
  8. 只跑 1 个 epoch 就下结论。蒸馏的收益需要一定的训练量才会显现,短训对比出来的结论往往不稳定。

六、下一步:能练手,也能接单

  1. 跑通本文实验并做完整消融,把「温度 × alpha × 容量差距」三维结果整理成一张表,这是社区里很受欢迎的一类硬核实测文。
  2. 换成真实模型:把 Teacher 换成 torchvision 里的 resnet18/mobilenet_v3,Student 换成更窄的自定义网络,重跑一遍,流程完全一致。
  3. 做一次"蒸馏 + 量化"组合:蒸馏出小模型后再做 INT8 量化,记录精度与推理速度的三方对比。
  4. 封装成蒸馏小工具:支持配置师生结构、温度、alpha,自动产出对比报告——这是一份很实在的接单作品。

结语

蒸馏的本质,是把"答案"换成"老师对答案的看法"来教学生。

记住三句话就够了:

  • 硬标签只说"是什么",软标签才说"像什么、差多远"——后者叫暗知识;
  • 温度 $T$ 负责把暗知识放大到可学习的量级,别忘了乘回 $T^2$;
  • 蒸馏不是量化的替代品,而是在量化之前,先把骨架换小的那一步。

最后留一句最实用的判断:如果你只是想在现有硬件上跑得更快,先量化;如果你需要的是一个结构完全不同、能塞进端侧的模型,才需要蒸馏。

你在蒸馏时用的是什么师生组合?温度取了多少?欢迎在评论区贴出你的消融结果,一起看看哪套配置最划算。

posted @ 2026-08-29 22:17  橘和柠  阅读(14)  评论(0)    收藏  举报