最近在围观各类开源模型技术讨论时有一个词几乎每篇公告都会出现——蒸馏。模型越来越大部署成本越来越高很多团队开始把蒸馏当作“把大模型压缩成实用小模型”的重要手段。本文将围绕蒸馏技术展开从概念、原理到 PyTorch 代码实战完整拆解帮你亲手跑通一个最小的知识蒸馏项目并给出工程落地中的关键建议。无论你是算法工程师、后端开发还是刚入门深度学习的学习者都可以参考这套流程。1. 什么是模型蒸馏1.1 用一个例子理解蒸馏蒸馏这个名字听起来有点玄但思路其实很朴素。想象一个经验丰富的老师傅带新人。老师傅见过大量复杂案例能给出非常细腻的判断新人虽然学得快但经验不足。如果让新人只对着标准答案学他能学会大部分规则但很难掌握那些“只可意会”的判断细节。如果老师傅不仅告诉新人“正确答案是什么”还把自己的判断倾向、犹豫程度、各选项之间的权衡都讲清楚新人就能学得更快、更好。深度学习里的知识蒸馏就是这个过程。大模型担任“教师模型”小模型担任“学生模型”。教师模型在处理样本时会输出一个概率分布比如一张手写数字图片它可能输出“7 的概率是 0.71 的概率是 0.2其他数字各有少量概率”。这个分布里隐藏着模型对相似类别的理解这种信息被称为“软标签”或“暗知识”。学生模型不仅要学会预测正确类别还要尽量拟合教师模型给出的软标签。这样小模型就能把大模型的判断能力“打包”进自己有限的参数里。1.2 蒸馏为什么成为热门话题过去几年模型的参数规模快速增长从百万级增长到千亿级。参数大了能力确实强但成本也随之升高训练成本高、显存占用大、推理延迟高。而在实际业务里很多场景并不需要千亿模型。比如一个移动端离线场景、一个客服意图识别模块、一个边缘设备上的检测任务模型体积和响应速度往往比极端精度更重要。蒸馏恰好提供了一条路径把大模型的“内力”传给小模型让小模型在体积缩小几十倍甚至上百倍的情况下保持接近大模型的效果。最近这段时间“蒸馏”在开源社区里频繁出现也和开放模型的演进有关基础大模型越来越强但真正能被社区广泛部署的往往是经过蒸馏、量化等轻量化处理后的小模型。许多开发者把蒸馏看作开放模型走向实用化的关键路径。1.3 蒸馏、剪枝、量化有什么区别很多初学者会把模型轻量化的几种手段弄混。这里用一张表做区分。技术核心思路主要作用会不会改变模型结构知识蒸馏用小模型学习大模型的输出分布在缩小参数量的同时保留较强效果会学生模型结构通常独立设计剪枝删除不重要的权重或神经元减少计算量和存储量会模型变“瘦”量化把高精度参数转为低精度表示减少内存占用加快推理不改变结构但改变参数精度三个方向不冲突。实际工程中经常组合使用先蒸馏出一个较小的模型再对模型做量化和剪枝进一步压榨部署成本。2. 知识蒸馏的核心原理2.1 软标签隐藏的“暗知识”先看一个具体场景。训练一个区分猫和狗的分类器常规训练方式使用 one-hot 标签也就是“这张图是猫”就记为[1, 0]“这张图是狗”就记为[0, 1]。这种方式很干净但也丢掉了很多信息。比如一只长得特别像猫的狗模型在判断时可能输出“狗的概率 0.6猫的概率 0.4”。如果只看 one-hot 标签这个 0.4 的信息就丢失了。但教师模型输出的概率分布不会丢。它会把这种“模糊”保留下来。学生模型学习这个分布时就能知道原来在这个任务里猫和狗在某些特征上是接近的。这就是软标签的价值它比硬标签携带更多信息。2.2 温度系数 T 的作用为了让教师模型输出的概率分布更加平滑Hinton 等人提出了带温度 T 的 Softmax。普通 Softmax 公式是p_i exp(z_i) / sum(exp(z_j))加入温度系数后变成p_i exp(z_i / T) / sum(exp(z_j / T))T 越大概率分布越平滑T 越小分布越尖锐。当 T 等于 1 时就退化成普通 Softmax。为什么要把分布变平滑因为教师模型训练充分后输出往往非常自信比如“7 的概率 0.99其他概率几乎为 0”。这样的分布和 one-hot 标签差别不大学生学不到暗知识。调高 T 之后那些本来很小的概率会被放大学生就能看到教师模型在不同类别之间的细微判断。2.3 蒸馏损失函数与整体训练流程蒸馏训练时学生模型的损失通常由两部分组成。第一部分是硬标签损失也就是学生输出和真实标签之间的交叉熵保证学生模型能正确分类。第二部分是蒸馏损失也就是学生输出和教师输出之间的 KL 散度保证学生的判断倾向接近教师。整体损失可以写成Loss alpha * L_hard (1 - alpha) * T^2 * L_soft其中T^2是温度系数的平方用来修正梯度尺度。这是因为在蒸馏损失里对 logits 除以了 T梯度会变小乘上 T^2 可以让梯度量级恢复正常。整体流程可以概括为训练一个性能优秀的教师模型。固定教师模型参数用它对训练数据生成软标签。定义学生模型结构参数量远小于教师模型。同时计算硬标签损失和软标签蒸馏损失训练学生模型。评估学生模型在验证集上的效果。3. 蒸馏的几种主流形态3.1 离线蒸馏最成熟的方案离线蒸馏是应用最广泛的方式。流程是先把教师模型训练好并固定住然后用它指导学生模型训练。它的优点是流程简单、容易控制。教师模型可以是一个已经上线的大模型学生模型在训练时可以复用历史数据不需要额外设计复杂的交互逻辑。缺点是学生模型能学到什么程度完全取决于教师模型的能力上限。如果教师模型本身效果一般学生也很难超越教师。3.2 在线蒸馏与自蒸馏在线蒸馏不再区分训练阶段教师和学生一起训练。模型可以是一个大模型和小模型组成也可以两个结构相近的模型互相学习。这种方式适合教师模型事先不存在、希望学生模型在训练过程中快速成长的场景。自蒸馏则更进一步让模型自己指导自己。例如把同一模型不同深度的输出做蒸馏让浅层输出向深层输出对齐。这样不需要额外训练大模型也能提升模型效果。3.3 数据蒸馏把数据集压缩成“精华”数据蒸馏的思路和模型蒸馏相反它不蒸馏模型而是蒸馏数据。简单来说是从大规模数据集中选择或合成一小部分“精华样本”让模型只在这些样本上训练也能达到接近完整数据集训练的效果。这可以大幅缩短训练时间降低数据存储和处理成本。数据蒸馏目前还是研究热点实际工程中使用时需要谨慎评估信息损失。3.4 “蒸馏一本书”文档知识库蒸馏的工程思路最近社区里出现了一个比较形象的说法——“蒸馏一本书的 skill 知识库”。它本质上不是模型蒸馏而是一种工程路径把一本书、一份产品文档或一个领域的知识库蒸馏成一个小模型能够掌握的能力。常见做法是将文档切分成片段。使用更强的模型生成问答对或指令数据。用这些数据微调或蒸馏一个小模型。再把小模型接入检索增强生成或知识库系统。这种方式特别适合垂直领域。一个大模型无法覆盖所有专业细节但通过知识蒸馏可以让一个小模型牢牢掌握某个领域的核心逻辑并且保持部署成本可控。4. 环境准备与项目结构4.1 依赖版本本文的代码示例以 PyTorch 为例。推荐环境如下实际版本可以根据你本机情况调整。依赖项建议版本Python3.9 或更高版本PyTorch2.xtorchvision与 PyTorch 版本匹配CUDA可选没有 GPU 也能运行建议使用虚拟环境安装依赖。命令如下python -m venv venv source venv/bin/activate pip install torch torchvision如果你使用 GPU需要根据你的 CUDA 版本到 PyTorch 官网选择对应的安装命令。如果只是学习原理CPU 环境也完全够用。4.2 数据集本文使用 MNIST 手写数字数据集。它包含 6 万张训练图片和 1 万张测试图片每张图片是 28x28 的灰度图类别为 0-9 的数字。MNIST 数据量适中模型训练速度快很适合用来验证蒸馏流程。4.3 项目文件结构建议按下面的目录结构组织代码distill-demo/ ├── models.py # 定义教师模型和学生模型 ├── train_teacher.py # 训练教师模型并保存权重 ├── train_student.py # 分别用普通训练和蒸馏训练学生模型 └── data/ # MNIST 数据保存目录5. 代码实战用 PyTorch 完成一次知识蒸馏接下来我们完整实现一个蒸馏项目。教师模型用三层卷积网络学生模型用较浅的两层卷积网络最终对比学生在普通训练和蒸馏训练下的表现差异。5.1 定义教师模型和学生模型文件路径models.pyimport torch.nn as nn class TeacherNet(nn.Module): 教师模型参数量较大能力更强。 def __init__(self): super().__init__() self.conv nn.Sequential( nn.Conv2d(1, 32, kernel_size3, padding1), nn.ReLU(), nn.MaxPool2d(2), nn.Conv2d(32, 64, kernel_size3, padding1), nn.ReLU(), nn.MaxPool2d(2), ) self.fc nn.Sequential( nn.Flatten(), nn.Linear(64 * 7 * 7, 128), nn.ReLU(), nn.Dropout(0.5), nn.Linear(128, 10), ) def forward(self, x): return self.fc(self.conv(x)) class StudentNet(nn.Module): 学生模型参数更少结构更浅。 def __init__(self): super().__init__() self.conv nn.Sequential( nn.Conv2d(1, 16, kernel_size3, padding1), nn.ReLU(), nn.MaxPool2d(2), ) self.fc nn.Sequential( nn.Flatten(), nn.Linear(16 * 14 * 14, 64), nn.ReLU(), nn.Linear(64, 10), ) def forward(self, x): return self.fc(self.conv(x))教师模型的参数量明显大于学生模型。教师模型在两个卷积层上提取更丰富的特征学生模型只在第一层卷积后就直接连接全连接层参数量更少推理速度更快。5.2 训练教师模型文件路径train_teacher.pyimport torch from torch import nn from torch.utils.data import DataLoader from torchvision import datasets, transforms from models import TeacherNet def evaluate(model, loader, device): model.eval() correct 0 total 0 with torch.no_grad(): for x, y in loader: x, y x.to(device), y.to(device) pred model(x).argmax(dim1) correct (pred y).sum().item() total y.size(0) return correct / total def main(): device cuda if torch.cuda.is_available() else cpu print(using device:, device) transform transforms.Compose([ transforms.ToTensor(), transforms.Normalize((0.1307,), (0.3081,)), ]) train_ds datasets.MNIST(./data, trainTrue, downloadTrue, transformtransform) test_ds datasets.MNIST(./data, trainFalse, downloadTrue, transformtransform) train_loader DataLoader(train_ds, batch_size256, shuffleTrue) test_loader DataLoader(test_ds, batch_size256, shuffleFalse) teacher TeacherNet().to(device) optimizer torch.optim.Adam(teacher.parameters(), lr1e-3) criterion nn.CrossEntropyLoss() for epoch in range(8): teacher.train() total_loss 0 for x, y in train_loader: x, y x.to(device), y.to(device) optimizer.zero_grad() out teacher(x) loss criterion(out, y) loss.backward() optimizer.step() total_loss loss.item() * x.size(0) acc evaluate(teacher, test_loader, device) avg_loss total_loss / len(train_ds) print(fepoch{epoch 1}, loss{avg_loss:.4f}, test_acc{acc:.4f}) torch.save(teacher.state_dict(), teacher.pt) print(teacher model saved to teacher.pt) if __name__ __main__: main()这里使用交叉熵损失训练教师模型。训练 8 个 epoch 后模型在测试集上的准确率通常可以达到 99% 左右然后把权重保存到teacher.pt文件。训练过程中的total_loss / len(train_ds)计算的是整个训练集的平均损失MNIST 数据集共 60000 张图片这个计算方式没有问题。5.3 通过蒸馏训练学生模型文件路径train_student.pyimport torch from torch import nn from torch.utils.data import DataLoader from torchvision import datasets, transforms from models import TeacherNet, StudentNet def load_data(batch_size256): transform transforms.Compose([ transforms.ToTensor(), transforms.Normalize((0.1307,), (0.3081,)), ]) train_ds datasets.MNIST(./data, trainTrue, downloadTrue, transformtransform) test_ds datasets.MNIST(./data, trainFalse, downloadTrue, transformtransform) train_loader DataLoader(train_ds, batch_sizebatch_size, shuffleTrue) test_loader DataLoader(test_ds, batch_sizebatch_size, shuffleFalse) return train_loader, test_loader def evaluate(model, loader, device): model.eval() correct 0 total 0 with torch.no_grad(): for x, y in loader: x, y x.to(device), y.to(device) pred model(x).argmax(dim1) correct (pred y).sum().item() total y.size(0) return correct / total def train_with_distill(teacher, student, train_loader, test_loader, device, T4.0, alpha0.7, epochs8): optimizer torch.optim.Adam(student.parameters(), lr1e-3) hard_criterion nn.CrossEntropyLoss() soft_criterion nn.KLDivLoss(reductionbatchmean) teacher.eval() for epoch in range(epochs): student.train() total_loss 0 for x, y in train_loader: x, y x.to(device), y.to(device) with torch.no_grad(): teacher_logits teacher(x) student_logits student(x) hard_loss hard_criterion(student_logits, y) soft_loss soft_criterion( torch.log_softmax(student_logits / T, dim1), torch.softmax(teacher_logits / T, dim1), ) * T * T loss alpha * hard_loss (1 - alpha) * soft_loss optimizer.zero_grad() loss.backward() optimizer.step() total_loss loss.item() * x.size(0) acc evaluate(student, test_loader, device) avg_loss total_loss / len(train_loader.dataset) print(f[distill] epoch{epoch 1}, loss{avg_loss:.4f}, test_acc{acc:.4f}) torch.save(student.state_dict(), student_distill.pt) def train_without_distill(student, train_loader, test_loader, device, epochs8): optimizer torch.optim.Adam(student.parameters(), lr1e-3) criterion nn.CrossEntropyLoss() for epoch in range(epochs): student.train() total_loss 0 for x, y in train_loader: x, y x.to(device), y.to(device) optimizer.zero_grad() out student(x) loss criterion(out, y) loss.backward() optimizer.step() total_loss loss.item() * x.size(0) acc evaluate(student, test_loader, device) avg_loss total_loss / len(train_loader.dataset) print(f[normal] epoch{epoch 1}, loss{avg_loss:.4f}, test_acc{acc:.4f}) torch.save(student.state_dict(), student_normal.pt) def main(): device cuda if torch.cuda.is_available() else cpu print(using device:, device) train_loader, test_loader load_data() teacher TeacherNet().to(device) teacher.load_state_dict(torch.load(teacher.pt, map_locationdevice)) print(teacher loaded) student_normal StudentNet().to(device) print(training student without distill ...) train_without_distill(student_normal, train_loader, test_loader, device) student_distill StudentNet().to(device) print(training student with distill ...) train_with_distill(teacher, student_distill, train_loader, test_loader, device) if __name__ __main__: main()在蒸馏训练函数里需要重点理解几个点。teacher_logits是在torch.no_grad()下计算的因为教师模型已经训练完成不需要更新梯度。如果我们不对教师模型做no_grad训练时就会额外计算教师模型的梯度增加显存和耗时。soft_criterion使用的是 KL 散度。PyTorch 的KLDivLoss要求第一个参数是学生模型输出的对数概率第二个参数是教师模型输出的概率顺序不要写反。soft_loss * T * T是为了补偿温度缩放带来的梯度尺度变化。如果温度 T 设置得比较大不乘回T^2的话学生模型的学习速度会变慢。alpha表示硬标签损失的权重。这里设置为 0.7蒸馏损失的权重就是 0.3。如果你希望学生模型更贴近教师模型的判断可以调低alpha比如设成 0.5。5.4 运行结果与解读先运行教师模型训练脚本python train_teacher.py再运行学生模型训练脚本python train_student.py由于随机种子、设备和 PyTorch 版本不同准确率会有波动但整体趋势比较稳定教师模型在测试集上的准确率约在 98.5% 到 99.2% 之间。普通训练的学生模型准确率约在 97.5% 到 98.5% 之间。蒸馏训练的学生模型准确率通常会比普通训练高 0.5 到 1.5 个百分点在 98% 到 99% 之间。可以看到学生模型参数量更小但通过蒸馏可以获得更接近教师模型的效果。原因在于蒸馏训练不仅让学生模型学习正确答案还让它学到了教师模型对不同数字之间的“感知”。比如数字 4 和 9 在图形上有相似之处普通训练时这种相似性不会被显式表达蒸馏时却能通过软标签传递给小模型。6. 蒸馏在开放模型生态中的工程意义6.1 为什么开放模型社区都在聊蒸馏开源开放模型生态发展到现在一个明显的趋势是基础模型越来越强但直接部署大语言模型对算力、内存和带宽的要求都很高。普通开发者要做一个垂直应用不太可能直接跑一个几百 B 参数的模型。蒸馏让“更强的教师模型”反哺“更轻量的学生模型”成为可能。社区里很多团队将大模型生成的高质量数据用于蒸馏最终发布参数量小得多的开放模型。这些模型保留了较强的通用能力同时让开发者能在消费级显卡上完成推理甚至部署到端侧设备。因此蒸馏不是一个实验室概念而是开放模型生态中连接“前沿能力”和“真实部署”的关键桥梁。6.2 从“大而全”到“小而专”通用大模型确实很强大但它面对垂直领域时存在两个问题一是知识覆盖不够深二是推理成本高。工程上可以基于大模型蒸馏出面向特定领域的小模型比如法律问答、客服意图识别、医疗分诊辅助等。“蒸馏一本书”的思路就是这个方向。先把领域文档切分清洗再让大模型生成一批高质量的问答对最后用这批数据训练一个千百万参数级别的小模型。这样的模型虽然在综合能力上不能和大模型比但在特定领域内可能表现非常稳定而且部署成本低、响应快。实际项目中蒸馏出来的垂直小模型还可以与检索系统配合使用形成“小模型初筛 知识库召回 大模型兜底”的混合架构。6.3 蒸馏与开源协议、安全边界做模型蒸馏时不能只关注技术指标。如果在你的业务中需要把一个大模型蒸馏成另一个模型需要关注该模型的许可协议。不同模型对“是否允许使用其输出训练第三方模型”有不同规定训练前要确认使用条款避免合规风险。另外蒸馏过程会继承教师模型已有的偏见和错误甚至会放大某些数据分布不均匀带来的问题。如果训练数据包含用户隐私或敏感信息必须提前做脱敏处理。蒸馏结果上线前建议做一轮针对性评估尤其是面向真实用户的生成内容需要设置合理的过滤和审核机制。7. 常见问题与排查思路问题现象常见原因解决思路蒸馏后学生模型效果反而更差温度 T 设置不合适或教师模型本身太弱先确认教师模型精度再调整温度尝试 T3、4、8、10训练时显存不足教师模型未设置为 eval 或未使用 no_grad训练前对教师模型调用 eval并包裹 torch.no_grad蒸馏损失不下降KLDivLoss 的第一个参数没有取 log使用 torch.log_softmax(student_logits / T, dim1)学生模型输出概率过于平滑温度 T 过大降低温度让概率分布更接近真实判断学生模型只能学到硬标签信息alpha 设置过大蒸馏损失权重过小适当降低 alpha增大软标签影响加载权重时报 shape 不匹配教师模型和学生模型定义不一致检查模型类结构和保存权重的对应关系CPU 上训练太慢数据量大或模型计算量大可以减少 epoch、调小教师模型或使用 GPU其中最常见的问题是把KLDivLoss的输入顺序写反。PyTorch 的KLDivLoss(pred, target)要求pred是模型预测的对数概率target是目标概率分布传反之后梯度方向会出错。另外一个容易被忽略的问题是温度T对蒸馏损失量级的影响。如果温度很大但忘记乘T^2学生模型的蒸馏信号会被削弱效果提升不明显。8. 最佳实践与工程建议8.1 先确认教师模型足够好蒸馏的前提是有一个可靠的教师模型。教师模型如果本身欠拟合或过拟合学生模型学到的“经验”就是错误的。工程上建议先充分训练教师模型确保验证集指标稳定。检查教师模型的错误案例确认是数据问题还是模型能力不足。如果条件允许使用多个教师模型做集成再蒸馏到学生模型通常效果更稳定。8.2 温度、损失权重、数据增强怎么调蒸馏场景下需要调整的参数主要有三个温度T、硬标签损失权重alpha、训练数据规模。温度T的经验值通常在 3 到 10 之间。任务越复杂类别越多所需的温度通常越高。你可以先固定T4观察学生模型的蒸馏损失表现再逐步调大或调小。alpha控制了学生模型对真实标签和教师软标签的依赖程度。建议从 0.7 开始对比普通训练和蒸馏训练的差距再决定是否降低alpha。蒸馏训练对数据质量的要求同样重要。即使教师模型很强如果训练数据分布与实际场景偏差很大学生模型依然无法泛化。保持数据增强策略和真实部署场景一致是蒸馏工程中比较容易忽略但收益极高的地方。8.3 蒸馏后的评估、量化与上线模型上线前需要做完整的评估。不能只看准确率还要关注不同类别上的表现是否有明显偏斜。输入扰动下模型是否稳定。推理延迟和吞吐量是否满足线上要求。显存占用是否符合部署环境限制。如果蒸馏后的模型仍然偏大可以继续做量化。比如把 PyTorch 模型转为 INT8 精度推理速度和内存占用都会有明显改善。蒸馏和量化属于两个独立优化维度组合使用能获得更极致的部署收益。8.4 生产环境注意事项在生产环境使用蒸馏模型时建议建立监控和回滚机制。对模型输入输出做日志采样关注线上数据与训练数据的分布差异。设置置信度阈值低置信度请求可以转发给更大模型兜底。使用 A/B 实验对比蒸馏模型和旧模型的效果再决定是否全量上线。保留上一个大版本模型权重以便快速回滚。安全方面要做到最小权限原则。蒸馏模型如果部署在服务端接口需要做认证和限流如果发布到端侧要考虑模型文件被提取后的知识产权风险避免投喂未授权的敏感信息。9. 总结与下一步实践建议本文从“蒸馏是什么”讲起介绍了软标签、温度系数、损失函数并给出了一套完整的 PyTorch 蒸馏代码。通过教师模型训练、学生模型普通训练和蒸馏训练三组实验的对比你应该已经看到蒸馏让小模型逼近大模型效果的基本过程。接下来可以做三件事第一把文章里的代码跑通然后修改T和alpha两个参数观察学生模型精度的变化加深对蒸馏原理的理解。第二将 MNIST 换成自己的业务数据。如果你的训练数据比较少可以用一个外部大模型生成软标签再蒸馏一个小模型这是一种很实用的冷启动方案。第三学习进阶方向。比如在线蒸馏、自蒸馏、数据蒸馏以及蒸馏与量化的组合优化。蒸馏本质上不是“模型变小的魔法”而是一种“让已有知识更高效传递”的训练范式。当你能控制教师模型、学生模型、温度参数和训练数据之间的关系时你就已经掌握了模型压缩中的一个核心武器。建议直接上手跑一遍示例代码再回到自己的项目里调整很快就能感受到蒸馏带来的实际收益。