ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

知识蒸馏实战教程:用PyTorch将大模型压缩为小模型

知识蒸馏实战教程:用PyTorch将大模型压缩为小模型 “什么时候蒸馏我自己”——看到这个标题很多同学可能会会心一笑。这句话表面上是一句程序员的自我调侃但拆开来看它恰好指向了深度学习里一个非常实用的技术方向知识蒸馏Knowledge Distillation。通俗点说知识蒸馏就是“训练一个大模型当老师再让一个小模型跟着老师学”最终让小模型在参数量大幅减少的情况下尽量逼近大模型的精度表现。本文会用一篇完整的实战教程把知识蒸馏的原理、代码、训练过程和踩坑点讲清楚让你看完之后不仅能理解“蒸馏”到底在蒸什么还能自己动手跑通一个完整的蒸馏训练流程。这篇文章适合以下几类读者刚入门深度学习想了解模型压缩与加速的同学。已经能跑通基础 PyTorch 训练但还不清楚如何做“师生模型”训练的人。在端侧、边缘设备或 Web 端部署模型时发现模型太大、推理太慢想通过蒸馏压缩模型的工程师。读完本文你将掌握知识蒸馏的核心思想、损失函数的构造方法、温度系数 T 的作用以及一套基于 PyTorch 的完整可运行示例。下面我们正式开始。1. 背景与核心概念1.1 为什么需要知识蒸馏先来看一个常见的现实问题在图像分类、文本分类等任务中通常模型越大、层数越深效果越好。一个 ResNet-50、BERT-base 甚至更大的模型在 GPU 服务器上训练和推理都没有问题。但一旦要部署到手机 App、浏览器、嵌入式设备上就会遇到几个棘手的问题模型参数量太大占用存储空间。推理延迟高用户点击一次要等好几秒。内存占用过高低端设备直接崩溃。有人会说那直接换一个小模型不就行了确实MobileNet、SqueezeNet 这类轻量网络在结构上做了很多优化但如果直接用小型化数据去训练一个轻量模型效果通常比大模型差一截。原因是小模型容量有限很难从零开始完全学到复杂数据中的规律。知识蒸馏提供了一种折中方案先用大模型教师网络在数据上学到丰富的知识然后让小模型学生网络去模仿大模型的输出。因为教师网络的输出包含了“类间相似度”这种软信息学生网络能学到的内容比只依赖真实标签One-Hot 硬标签要多得多。1.2 什么是知识蒸馏知识蒸馏这个概念最早由 Hinton 等人在 2015 年的论文Distilling the Knowledge in a Neural Network中正式提出。它的核心思路可以概括为三步训练一个性能较好的大模型称为教师网络Teacher。设计一个参数量更少的小模型称为学生网络Student。在学生网络的训练过程中不仅让它学习真实标签还要让它模仿教师网络的输出概率分布。这里的关键在于“概率分布”。普通的分类任务使用交叉熵损失目标是把真实类别的概率推向 1其他类别推向 0。但教师网络在预测时除了正确类别之外其他类别的概率并不是正好为 0而是有高有低。比如一张图片是一只猫教师网络可能输出类别概率猫0.75狗0.15老虎0.06其他0.04这个概率分布里“狗”和“老虎”的概率比“其他”高说明教师网络认为猫和狗、猫和老虎在视觉特征上有一定的相似性。这种相似性就是教师网络从海量数据中学到的“软知识”。学生网络如果只学习硬标签就只会得到一个“猫”的判断结果但如果学习教师网络的软输出它就能额外知道“猫有点像狗但不太像老虎”。这种软知识正是蒸馏能够提升小模型效果的根本原因。1.3 知识蒸馏的常见应用场景知识蒸馏并不是只能用于图像分类。在工程实践中它的应用场景非常广泛模型压缩将大规模 Transformer 或 ResNet 蒸馏为小模型用于移动端部署。跨架构迁移把 Transformer 的知识蒸馏到 LSTM 或 CNN 结构中。多任务与大模型简化用一个大模型同时训练出多个专用小模型。半监督与自监督增强利用教师模型对无标注数据生成伪标签结合蒸馏训练学生模型。推荐系统与排序模型用复杂排序模型蒸馏出轻量召回模型。理解了背景之后我们接下来从环境准备开始一步步完成一个知识蒸馏的实战项目。2. 环境准备与版本说明2.1 基础环境本文的示例代码基于 Python 和 PyTorch。建议使用以下环境Python3.8 或更高版本。PyTorch1.10 及以上2.x 版本也可以。torchvision与 PyTorch 版本对应的版本。操作系统Windows / Linux / macOS 均可。如果你使用的是 GPU 环境CUDA 11.x 或更高版本训练速度会快很多。如果没有 GPU用 CPU 训练也能完成本文的示例只是耗时会长一些。创建虚拟环境并安装依赖# 创建虚拟环境可选 python -m venv distill_env # 激活虚拟环境 # Windows: distill_env\Scripts\activate # Linux/macOS: source distill_env/bin/activate # 安装依赖 pip install torch torchvision matplotlib版本说明不同版本的 PyTorch 在 API 上略有差异但本文用到的nn.CrossEntropyLoss、nn.KLDivLoss、F.log_softmax都是长期稳定的接口在各版本中均可使用。如果你使用的 PyTorch 版本非常新遇到个别 API 变动以官方文档为准。2.2 示例项目结构为了方便阅读和运行我们采用下面的项目结构knowledge_distillation/ ├── train.py # 完整训练脚本 ├── models.py # 教师网络与学生网络定义 ├── distill.py # 蒸馏训练逻辑 └── README.md # 项目说明为了让代码更清晰这里把模型定义、蒸馏训练逻辑和主脚本分开。实际项目中可以按照自己的习惯组织目录结构但保持模块化是一个好习惯。3. 知识蒸馏核心原理拆解3.1 软标签与温度系数 T前面提到教师网络输出的概率分布就是“软标签”。但这里有一个问题如果教师网络的输出过于自信比如正确类别的概率是 0.99其他类别几乎为 0那么软标签和硬标签的区别就不明显了学生网络也学不到太多额外信息。为了解决这个问题Hinton 在论文中引入了“温度系数 T”。在计算概率分布时不再直接使用网络输出的 logits 做 Softmax而是先除以 T[ q_i \frac{\exp(z_i / T)}{\sum_j \exp(z_j / T)} ]其中 (z_i) 是网络输出的原始 logitsT 是温度系数。当 T1 时这就是普通的 Softmax 输出当 T1 时概率分布变得更“平滑”类别之间的差异被缩小模型就能暴露出更多“软知识”。下图是一个简单的示意T1输出概率分布比较尖锐。T4输出概率分布更平均类间关系更明显。T 越大分布越接近均匀分布但超过一定范围后有效信息反而被淹没。在蒸馏训练中教师网络和学生网络在计算蒸馏损失时通常使用相同的温度 T且 T 往往大于 1例如 T4 或 T5。3.2 蒸馏损失函数的构成知识蒸馏的总损失由两部分组成硬标签损失Student Loss学生网络的输出与真实标签之间的交叉熵损失。蒸馏损失Distill Loss学生网络的高温输出与教师网络的高温输出之间的 KL 散度损失。数学表达可以写成[ L \alpha \cdot L_{hard} (1 - \alpha) \cdot L_{soft} ]其中(L_{hard}) 是学生网络输出与真实标签的交叉熵。(L_{soft}) 是学生网络与教师网络的软标签之间的 KL 散度。(\alpha) 是硬标签损失的权重通常设置在 0.1 到 0.7 之间。为什么需要两个损失如果只看蒸馏损失学生网络只学到“模仿老师”如果只看硬标签损失那就退化成普通训练。两者结合学生既能从老师的经验中受益又不会完全被老师的错误判断带偏。3.3 为什么用 KL 散度而不是交叉熵KL 散度Kullback-Leibler Divergence用于衡量两个概率分布之间的差异。在蒸馏中教师网络的输出概率作为“目标分布”学生网络的输出概率作为“预测分布”KL 散度越小说明学生越接近老师。计算方式如下# 教师网络和学生网络的输出都先经过 log_softmax loss_kl nn.KLDivLoss(reductionbatchmean)( F.log_softmax(student_output / T, dim1), F.softmax(teacher_output / T, dim1) ) * (T * T)注意这里有一个细节F.log_softmax在前F.softmax在后。因为KLDivLoss的输入要求第一个参数是对数概率第二个参数是普通概率。最后的(T * T)是梯度缩放修正从论文中沿用下来的处理方式目的是让损失值在不同温度下保持合理的尺度。3.4 蒸馏训练的整体流程一次完整的蒸馏训练过程可以概括为以下步骤加载数据集分成训练集和测试集。定义教师网络和学生网络。固定教师网络参数先训练教师网络或直接加载一个预训练好的教师模型。在每一轮训练中同时向学生网络输入真实标签和教师网络的软输出。计算硬标签损失和蒸馏损失加权求和后反向传播更新学生网络参数。在测试集上评估学生网络效果。接下来就用 PyTorch 把上面这套流程完整实现出来。4. 完整实战案例PyTorch 实现知识蒸馏4.1 创建项目结构首先创建项目目录和文件mkdir knowledge_distillation cd knowledge_distillation touch train.py models.py distill.py4.2 定义教师网络和学生网络为了演示方便我们使用 MNIST 手写数字数据集。MNIST 有 10 个类别图像尺寸是 28x28结构简单训练速度快非常适合用来验证蒸馏流程。这里教师网络使用一个参数量较多的 CNN学生网络使用一个参数量较少的 CNN。完整定义放在models.py中# 文件路径models.py import torch import torch.nn as nn import torch.nn.functional as F class TeacherNet(nn.Module): 教师网络参数量较多、表达能力更强。 结构两层卷积 三层全连接 def __init__(self): super(TeacherNet, self).__init__() self.conv1 nn.Conv2d(1, 32, kernel_size3, padding1) self.conv2 nn.Conv2d(32, 64, kernel_size3, padding1) self.pool nn.MaxPool2d(2, 2) self.fc1 nn.Linear(64 * 7 * 7, 512) self.fc2 nn.Linear(512, 256) self.fc3 nn.Linear(256, 10) def forward(self, x): x self.pool(F.relu(self.conv1(x))) x self.pool(F.relu(self.conv2(x))) x x.view(x.size(0), -1) x F.relu(self.fc1(x)) x F.relu(self.fc2(x)) x self.fc3(x) return x class StudentNet(nn.Module): 学生网络参数量少、结构轻量。 结构一层卷积 两层全连接 def __init__(self): super(StudentNet, self).__init__() self.conv1 nn.Conv2d(1, 16, kernel_size3, padding1) self.pool nn.MaxPool2d(2, 2) self.fc1 nn.Linear(16 * 14 * 14, 64) self.fc2 nn.Linear(64, 10) def forward(self, x): x self.pool(F.relu(self.conv1(x))) x x.view(x.size(0), -1) x F.relu(self.fc1(x)) x self.fc2(x) return x说明教师网络采用两层卷积加三层全连接参数约 50 万级别。学生网络只保留一层卷积加两层全连接参数约 3 万级别。两个网络的输入输出维度一致都是 MNIST 的 1x28x28 输入10 类输出。4.3 编写蒸馏训练逻辑接下来在distill.py中编写蒸馏训练的核心逻辑。这里包括训练教师网络、蒸馏训练学生网络两个环节。因为教师网络定义较复杂我们先训练它蒸馏时冻结它的参数。# 文件路径distill.py import torch import torch.nn as nn import torch.optim as optim import torch.nn.functional as F from torch.utils.data import DataLoader from torchvision import datasets, transforms def train_teacher(model, device, train_loader, optimizer, criterion, epoch): 普通训练教师网络 model.train() running_loss 0.0 for batch_idx, (data, target) in enumerate(train_loader): data, target data.to(device), target.to(device) optimizer.zero_grad() output model(data) loss criterion(output, target) loss.backward() optimizer.step() running_loss loss.item() if batch_idx % 100 0: print(fEpoch {epoch} [{batch_idx * len(data)}/{len(train_loader.dataset)}] fLoss: {loss.item():.6f}) def evaluate(model, device, test_loader): 评估模型准确率 model.eval() correct 0 total 0 with torch.no_grad(): for data, target in test_loader: data, target data.to(device), target.to(device) output model(data) pred output.argmax(dim1, keepdimTrue) correct pred.eq(target.view_as(pred)).sum().item() total target.size(0) acc 100.0 * correct / total print(fTest accuracy: {acc:.2f}%) return acc def distill_train(student, teacher, device, train_loader, optimizer, T, alpha): 蒸馏训练学生网络。 参数说明 - student: 学生网络 - teacher: 教师网络训练完成后冻结 - T: 温度系数 - alpha: 硬标签损失的权重 student.train() teacher.eval() # 教师网络固定不参与梯度更新 hard_criterion nn.CrossEntropyLoss() soft_criterion nn.KLDivLoss(reductionbatchmean) running_loss 0.0 for batch_idx, (data, target) in enumerate(train_loader): data, target data.to(device), target.to(device) optimizer.zero_grad() # 学生网络输出 student_output student(data) # 教师网络输出不计算梯度 with torch.no_grad(): teacher_output teacher(data) # 硬标签损失 loss_hard hard_criterion(student_output, target) # 蒸馏损失使用高温 softmax loss_soft soft_criterion( F.log_softmax(student_output / T, dim1), F.softmax(teacher_output / T, dim1) ) * (T * T) # 总损失 loss alpha * loss_hard (1 - alpha) * loss_soft loss.backward() optimizer.step() running_loss loss.item() if batch_idx % 100 0: print(fDistill [Batch {batch_idx}] Total Loss: {loss.item():.6f}, fHard Loss: {loss_hard.item():.6f}, Soft Loss: {loss_soft.item():.6f})这段代码是蒸馏训练的核心有几个地方值得仔细说明teacher.eval()确保教师网络中的 BatchNorm/Dropout 不产生影响。推理模式下必须设置。with torch.no_grad()包裹教师网络的前向传播节省显存和计算量。KL 散度损失中教师网络输出经过softmax学生网络输出经过log_softmax顺序不能颠倒。最终损失是硬标签损失和蒸馏损失的加权和通过alpha控制两者的权重。4.4 编写主训练脚本主脚本train.py负责加载数据、初始化模型、依次训练教师网络和学生网络。这里使用 MNIST 数据集如果本地没有数据PyTorch 会自动下载。# 文件路径train.py import torch import torch.nn as nn import torch.optim as optim from torch.utils.data import DataLoader from torchvision import datasets, transforms from models import TeacherNet, StudentNet from distill import train_teacher, evaluate, distill_train def main(): # 设备设置 device torch.device(cuda if torch.cuda.is_available() else cpu) print(fUsing device: {device}) # 数据预处理 transform transforms.Compose([ transforms.ToTensor(), transforms.Normalize((0.1307,), (0.3081,)) ]) # 加载 MNIST train_dataset datasets.MNIST( root./data, trainTrue, downloadTrue, transformtransform ) test_dataset datasets.MNIST( root./data, trainFalse, downloadTrue, transformtransform ) batch_size 128 train_loader DataLoader(train_dataset, batch_sizebatch_size, shuffleTrue) test_loader DataLoader(test_dataset, batch_sizebatch_size, shuffleFalse) # 初始化模型 teacher TeacherNet().to(device) student StudentNet().to(device) # 第一步训练教师网络 print( Training Teacher Network ) teacher_optimizer optim.Adam(teacher.parameters(), lr0.001) teacher_criterion nn.CrossEntropyLoss() for epoch in range(5): train_teacher(teacher, device, train_loader, teacher_optimizer, teacher_criterion, epoch 1) train_acc evaluate(teacher, device, train_loader) test_acc evaluate(teacher, device, test_loader) print(fTeacher Epoch {epoch 1}: train_acc{train_acc:.2f}%, test_acc{test_acc:.2f}%) # 第二步蒸馏训练学生网络 print( Distilling Student Network ) student_optimizer optim.Adam(student.parameters(), lr0.001) T 4 # 温度系数 alpha 0.3 # 硬标签损失权重 # 冻结教师网络 for param in teacher.parameters(): param.requires_grad False for epoch in range(5): distill_train(student, teacher, device, train_loader, student_optimizer, T, alpha) test_acc evaluate(student, device, test_loader) print(fStudent Distill Epoch {epoch 1}: test_acc{test_acc:.2f}%) # 保存模型 torch.save(student.state_dict(), student_model.pth) torch.save(teacher.state_dict(), teacher_model.pth) print(Models saved.) if __name__ __main__: main()4.5 运行与验证在项目目录下运行python train.py预期输出大致如下Using device: cuda Training Teacher Network Epoch 1 [0/60000] Loss: 0.358264 ... Test accuracy: 98.21% Teacher Epoch 1: train_acc97.52%, test_acc98.21% ... Distilling Student Network Distill [Batch 0] Total Loss: 1.982103, Hard Loss: 1.823456, Soft Loss: 0.553211 ... Student Distill Epoch 5: test_acc97.26%为了验证蒸馏的有效性可以再单独训练一个不使用蒸馏的 StudentNet并对比测试准确率。通常来说经过蒸馏的学生网络会比从零训练的学生网络有 0.5 到 2 个百分点的提升具体取决于数据集、温度和损失权重的设置。5. 常见问题与排查思路5.1 教师网络输出过拟合如果教师网络在训练集上效果很好但在测试集上泛化很一般那么学生学到的是教师“背下来的答案”而不是“理解的规律”。这种情况常见于教师网络过度训练或者训练轮次过多。建议教师网络训练到验证集准确率不再提升时就停止并尽量使用早停或正则化手段。5.2 温度 T 设置不当温度 T 太小软标签接近硬标签蒸馏退化为普通训练。温度 T 太大软标签接近均匀分布学生网络学不到类别间差异。T 取值范围效果1无蒸馏效果2 到 5常用区间推荐从 4 开始尝试10 以上分布过于平滑容易丢失有效信息推荐做法将 T 作为超参数在验证集上分别尝试 2、3、4、5找到最合适的值。5.3 KL 散度损失出现 NaN如果 KL 散度损失出现 NaN通常是因为教师网络或学生网络的输出经过 Softmax 后出现极端值或者 logits 太小导致数值不稳定。排查顺序检查数据是否归一化。检查网络输出是否有限值可以在前向传播后打印output.min()和output.max()。尝试在 softmax 前对 logits 做 clip例如限制在 [-5, 5]。降低学习率。5.4 蒸馏后学生网络没有提升这是最常见的问题。可能的原因有教师网络不够强软知识价值有限。学生网络结构太简单容量严重不足无论怎么学都学不会。温度 T 设置过低软标签和硬标签几乎一样。alpha 权重过高蒸馏损失的影响微乎其微。训练轮次不足学生网络还没收敛就提前结束。排查时可以先做一个 base line直接训练学生网络并记录准确率然后对比蒸馏训练后的准确率。如果两者几乎一样优先调整 T 和 alpha。5.5 显存不足蒸馏训练需要同时加载教师网络和学生网络显存开销比单独训练学生网络要高。可以尝试减少 batch size。使用 CPU 训练MNIST 这类小数据集完全可行。分阶段推理预先用教师网络生成所有样本的软标签保存到磁盘蒸馏训练时直接读取不需要在内存中保留教师网络。这也是工程上推荐的优化方式。具体思路是训练教师网络后遍历训练集把教师输出的软标签存为.npy或.pt文件训练学生时只加载这些文件。6. 最佳实践与工程建议知识蒸馏在学术与工程中的实现方式远比上面的示例丰富这里整理一些实用的工程建议。6.1 先做模型容量评估蒸馏不是万能的。如果教师网络和学生网络的容量差距过大学生网络可能根本无法吸收教师的知识。建议先用下面的思路评估单独训练学生网络看它能达到的准确率上限。用教师网络结构替换学生网络确定目标效果的“天花板”。如果两者差距非常大考虑增加学生网络容量或者改用更优的轻量结构而不是一味地加大蒸馏损失权重。6.2 软标签预处理与离线蒸馏在大型数据集上训练时如果每个 epoch 都要前向传播教师网络开销非常大。工程上更常见的方式是训练好教师网络后对全部训练数据执行一次前向传播。将教师网络的 softmax(T) 输出保存到磁盘作为预计算软标签。训练学生网络时不再加载教师网络直接从磁盘读取软标签。这样做的好处有两点一是显著减少显存和训练耗时二是软标签可以重复使用更换温度 T 或损失权重时无需重新跑教师网络。6.3 动态权重与课程蒸馏比较进阶的做法是动态调整 alpha。训练初期学生网络还很弱可以适当提高蒸馏损失权重让它多观察教师的行为训练中后期学生逐渐成熟再慢慢提高硬标签损失权重让它与真实任务对齐。简单实现伪代码如下# 动态 alpha 示例 alpha max(0.7 - epoch * 0.1, 0.1)6.4 集成蒸馏与自蒸馏除了单教师蒸馏还有几种常见的扩展方式多教师蒸馏多个教师网络同时提供软标签取平均或加权融合学生网络可以获得更全面的知识。自蒸馏让同一个网络的大版本来蒸馏小版本例如在训练过程中使用 EMA指数滑动平均模型作为教师。特征蒸馏不只模仿输出概率还让学生网络模仿教师网络中间层的特征图常用于检测、分割等任务。这些方案都是建立在基础蒸馏框架之上的优化。先跑通本文的基础代码再逐步引入复杂度是比较稳妥的学习路径。6.5 验证集与超参管理蒸馏训练有两个网络、多个超参数很容易过拟合验证集。建议建立一套规范固定教师网络只调学生网络的超参数。使用独立的验证集调参测试集只在最终评估时使用一次。记录每次实验的温度 T、权重 alpha、学习率、教师准确率、学生准确率方便对比。7. 总结与学习路线本文从“什么时候蒸馏我自己”这句调侃出发系统梳理了知识蒸馏的核心概念、原理与完整实战。相信你现在已经能回答这几个关键问题蒸馏到底蒸的是什么教师网络软输出中包含的类间相似性信息。温度 T 的作用控制概率分布的平滑程度暴露更多软知识。总损失如何构成硬标签的交叉熵损失 教师软标签的 KL 散度损失。如何用 PyTorch 实现训练教师网络后冻结再训练学生网络。接下来如果你想继续深入这个方向可以从三个维度扩展模型结构角度尝试将教师学生模型替换为不同架构观察蒸馏效果的差异。任务角度将示例从 MNIST 换成 CIFAR-10或者尝试在 NLP 文本分类任务中使用蒸馏。部署角度将训练好的学生模型转换为 ONNX 或 TorchScript在端侧推理框架中部署。知识蒸馏并不是高不可攀的前沿技术它更像是一套“站在巨人肩膀上学习”的训练范式。希望这篇文章能帮你跑通第一条蒸馏训练流水线并在实际项目中真正用起来。如果本文对你有帮助欢迎收藏备用。你也可以在自己的数据集上调整温度 T 和权重 alpha看看学生网络能在多大程度上接近教师网络的表现。动手试一下跑出你自己的蒸馏模型下一次再问“什么时候蒸馏我自己”的时候你就能自信地回答现在就可以。
返回列表