最近在做模型压缩时反复对比了多种知识蒸馏方案。一开始我把注意力全部放在“教师模型输出分数”上也就是用 soft label 去指导学生模型但效果始终差一口气。后来把教师模型中间层的特征图、注意力分布也加入蒸馏目标后学生模型的效果才明显提升。这个现象让我意识到一个经常被忽略的问题知识蒸馏里教师模型的推理习惯也就是它提取特征的方式往往比最终输出的分数更重要。本文将围绕这个问题展开先讲清楚知识蒸馏的基本原理再对比“分数蒸馏”和“特征蒸馏”两种路线的差异最后用 PyTorch 给出一个可复现的完整实战案例包含教师模型训练、学生模型蒸馏、中间层特征对齐等环节。适合正在做模型压缩、模型加速、移动端部署的同学参考也适合刚接触知识蒸馏的初学者系统理解这一技术。1. 知识蒸馏到底是什么1.1 从一个直观比喻开始知识蒸馏的核心思想可以这样理解我们希望训练一个小模型让它模仿一个大模型的行为。大模型通常参数量大、推理慢但精度高小模型参数量小、推理快但直接训练时精度往往达不到要求。蒸馏的过程就像让“师傅”带着“徒弟”学习。师傅拥有丰富的经验不仅告诉徒弟“正确答案是什么”还告诉徒弟“我对每个备选答案的把握程度”。这种把握程度就是所谓的“软标签”。比如在手写数字识别中一张数字“7”的图片硬标签是[0, 0, 0, 0, 0, 0, 0, 1, 0, 0]表示类别是 7。但教师模型可能会输出[0.01, 0.02, 0.01, 0.03, 0.02, 0.02, 0.03, 0.85, 0.01, 0.01]也就是说教师模型认为这张图有 85% 的概率是 7还有 3% 的概率像 92% 的概率像 1。这些“次级概率”其实是有价值的它们包含了教师模型对相似类别的判断经验。1.2 专业定义知识蒸馏Knowledge DistillationKD最早由 Hinton 等人在 2015 年的论文 《Distilling the Knowledge in a Neural Network》 中系统提出。其核心思路是先训练一个参数量大、性能强的教师模型。在训练学生模型时不仅使用真实标签计算交叉熵损失还使用教师模型的软输出作为监督信号。通过温度参数 T 对教师模型的输出进行软化使学生模型能够学到类别之间的相似性信息。用公式表示蒸馏损失L α * L_hard (1 - α) * L_soft其中L_hard是学生模型与真实标签之间的交叉熵。L_soft是学生模型软化后的输出与教师模型软化后的输出之间的 KL 散度。α是平衡两个损失的权重。1.3 大模型时代的知识蒸馏随着大语言模型的发展知识蒸馏的应用场景发生了变化。早期知识蒸馏主要用在图像分类、语音识别等任务中现在则大量用于将大语言模型的能力迁移到小模型上。在大模型场景下蒸馏不仅发生在输出层还发生在中间层。比如使用教师模型生成的回答来微调学生模型。使用教师模型的隐藏状态来对齐学生模型的中间表示。使用教师模型的 attention 分布来指导学生的注意力机制。这也是本文标题想强调的观点教师模型在推理过程中形成的特征表示习惯比如关注图像的哪些区域、在语义空间中如何排列样本这些“过程性知识”比最终输出的一串概率更有迁移价值。2. 环境准备与版本说明在开始代码之前先说明本文实验环境。知识蒸馏的核心计算是模型前向传播和梯度更新因此只要 PyTorch 环境正常都可以运行下面的代码。我使用的环境如下操作系统Ubuntu 20.04 / Windows 11 / macOS 均可 Python3.8 及以上 PyTorch1.10 及以上2.x 版本兼容 torchvision0.11 及以上 CUDA可选CPU 也可以运行只是速度较慢如果没有安装 PyTorch可以执行以下命令安装 CPU 版本pip install torch torchvision如果使用 GPU建议根据官网选择合适的 CUDA 版本pip install torch torchvision --index-url https://download.pytorch.org/whl/cu118本文实验使用 MNIST 手写数字数据集不需要额外下载大型数据集方便快速跑通流程。你可以在后续实验中将数据集替换为 CIFAR-10、ImageNet 子集或你自己的业务数据。3. 核心原理拆解分数蒸馏 vs 特征蒸馏3.1 温度参数的作用在经典知识蒸馏中温度参数 T 是关键。它的作用是将模型的 logits 输出“软化”。假设模型最后一层输出的 logits 为z_i经过 Softmax 后得到概率p_i exp(z_i / T) / Σ_j exp(z_j / T)当T 1时就是普通的 Softmax。当T 1时概率分布变得更加平滑类别之间的差异被缩小从而暴露出教师模型对“相似类别”的判断。当T很大时分布趋近于均匀分布信息又会被稀释。选温度参数时常见的做法是训练教师模型时使用T 1。蒸馏学生模型时设置T 3到T 10具体数值需要实验调整。蒸馏结束后学生模型推理时使用T 1。3.2 分数蒸馏的局限“分数蒸馏”是指只使用教师模型的 logits 或 softmax 输出作为监督信号。这种方法的优点是实现简单缺点是信息量有限。比如在图像分类任务中教师模型的输出只是一个长度为类别数的向量。如果类别数很少这个向量包含的信息量其实很低。两个模型可能输出几乎相同的概率分布但内部特征却完全不同。另一个典型问题是当教师模型对某个样本非常自信时它的输出分布会非常尖锐软化后的分布携带的“暗知识”也很少。这时候学生模型很难从教师输出中学习到有区分性的特征。3.3 特征蒸馏让推理习惯成为学习目标特征蒸馏也叫中间层蒸馏强调让学生模型模仿教师模型的中间层输出。以图像分类为例教师模型在层层卷积中会逐步提取边缘、纹理、部件、整体结构等不同层级的信息。如果我们让学生模型在对应层输出相似的特征图那就相当于引导学生模型“用教师的思路去看图”。特征蒸馏的代表性方法包括FitNets让学生模型的中间层特征拟合教师模型的中间层特征中间通过一个引导层做维度匹配。Attention Transfer让学生模型的注意力图拟合教师模型的注意力图。SPSimilarity Preserving让学生模型保持教师模型在样本对之间的相似性关系。本文实战部分会实现一个基于特征图 MSE 对齐的蒸馏方案这是最简单也最容易理解的一种。4. 实战PyTorch 实现特征蒸馏下面进入完整的代码实现环节。我们的任务是在 MNIST 数据集上训练一个较大的教师模型。设计一个较小的学生模型。使用教师模型的 logits 和中间层特征共同指导学生模型训练。对比三种训练方式的效果直接训练学生、仅 logits 蒸馏、logits 特征蒸馏。4.1 创建项目结构首先建立项目目录kd-demo/ ├── models.py # 教师模型和学生模型定义 ├── train_teacher.py # 训练教师模型脚本 ├── distill.py # 知识蒸馏脚本 └── utils.py # 工具函数包括数据加载、模型评估4.2 定义模型结构创建models.py定义教师模型和学生模型。教师模型使用三层卷积通道数分别为 16、32、64全连接层为 128。学生模型通道数缩小为 8、16、32全连接层为 64。这里关键的设计是教师模型和学生模型都要暴露中间层特征方便蒸馏时对齐。# 文件路径kd-demo/models.py import torch import torch.nn as nn import torch.nn.functional as F class TeacherNet(nn.Module): 教师模型参数量较大精度较高 def __init__(self, num_classes10): super(TeacherNet, self).__init__() self.conv1 nn.Conv2d(1, 16, kernel_size3, padding1) self.bn1 nn.BatchNorm2d(16) self.conv2 nn.Conv2d(16, 32, kernel_size3, padding1) self.bn2 nn.BatchNorm2d(32) self.conv3 nn.Conv2d(32, 64, kernel_size3, padding1) self.bn3 nn.BatchNorm2d(64) self.pool nn.MaxPool2d(2, 2) self.fc1 nn.Linear(64 * 3 * 3, 128) self.fc2 nn.Linear(128, num_classes) def forward(self, x): # 返回 logits 和中间层特征 x self.pool(F.relu(self.bn1(self.conv1(x)))) f1 x x self.pool(F.relu(self.bn2(self.conv2(x)))) f2 x x self.pool(F.relu(self.bn3(self.conv3(x)))) f3 x x x.view(x.size(0), -1) x F.relu(self.fc1(x)) logits self.fc2(x) return logits, [f1, f2, f3] class StudentNet(nn.Module): 学生模型参数量小结构更浅、通道更窄 def __init__(self, num_classes10): super(StudentNet, self).__init__() self.conv1 nn.Conv2d(1, 8, kernel_size3, padding1) self.bn1 nn.BatchNorm2d(8) self.conv2 nn.Conv2d(8, 16, kernel_size3, padding1) self.bn2 nn.BatchNorm2d(16) self.conv3 nn.Conv2d(16, 32, kernel_size3, padding1) self.bn3 nn.BatchNorm2d(32) self.pool nn.MaxPool2d(2, 2) self.fc1 nn.Linear(32 * 3 * 3, 64) self.fc2 nn.Linear(64, num_classes) def forward(self, x): x self.pool(F.relu(self.bn1(self.conv1(x)))) f1 x x self.pool(F.relu(self.bn2(self.conv2(x)))) f2 x x self.pool(F.relu(self.bn3(self.conv3(x)))) f3 x x x.view(x.size(0), -1) x F.relu(self.fc1(x)) logits self.fc2(x) return logits, [f1, f2, f3]这里教师模型和学生模型的池化层位置一致所以中间层特征图的尺寸是对齐的蒸馏时不需要额外做维度映射。在实际项目中如果教师和学生结构差异较大需要在蒸馏层之间加一个适配网络。4.3 编写工具函数创建utils.py包含数据加载和模型评估函数。# 文件路径kd-demo/utils.py import torch from torch.utils.data import DataLoader from torchvision import datasets, transforms def get_mnist_loaders(batch_size128): transform transforms.Compose([ transforms.ToTensor(), transforms.Normalize((0.1307,), (0.3081,)) ]) train_dataset datasets.MNIST( root./data, trainTrue, downloadTrue, transformtransform ) test_dataset datasets.MNIST( root./data, trainFalse, downloadTrue, transformtransform ) train_loader DataLoader(train_dataset, batch_sizebatch_size, shuffleTrue) test_loader DataLoader(test_dataset, batch_sizebatch_size, shuffleFalse) return train_loader, test_loader def evaluate(model, test_loader, devicecpu): model.eval() correct 0 total 0 with torch.no_grad(): for images, labels in test_loader: images, labels images.to(device), labels.to(device) logits, _ model(images) preds logits.argmax(dim1) correct (preds labels).sum().item() total labels.size(0) return correct / total4.4 训练教师模型创建train_teacher.py训练教师模型并保存权重。# 文件路径kd-demo/train_teacher.py import torch import torch.nn as nn from torch.optim import Adam from models import TeacherNet from utils import get_mnist_loaders, evaluate def train_teacher(epochs10, batch_size128, lr1e-3, devicecpu): train_loader, test_loader get_mnist_loaders(batch_size) teacher TeacherNet(num_classes10).to(device) criterion nn.CrossEntropyLoss() optimizer Adam(teacher.parameters(), lrlr) for epoch in range(epochs): teacher.train() total_loss 0.0 for images, labels in train_loader: images, labels images.to(device), labels.to(device) optimizer.zero_grad() logits, _ teacher(images) loss criterion(logits, labels) loss.backward() optimizer.step() total_loss loss.item() acc evaluate(teacher, test_loader, device) print(fEpoch {epoch 1}/{epochs}, Loss: {total_loss / len(train_loader):.4f}, Acc: {acc:.4f}) torch.save(teacher.state_dict(), teacher.pth) print(Teacher model saved to teacher.pth) return teacher if __name__ __main__: device cuda if torch.cuda.is_available() else cpu train_teacher(epochs10, devicedevice)运行结果示例Epoch 1/10, Loss: 0.1532, Acc: 0.9831 Epoch 10/10, Loss: 0.0183, Acc: 0.9941MNIST 分类任务相对简单10 轮训练后教师模型准确率可以达到 99% 以上。4.5 编写蒸馏脚本创建distill.py实现核心蒸馏逻辑。这一步是整个文章的重点。蒸馏损失由三部分组成硬标签交叉熵保证学生模型不偏离真实类别。软标签 KL 散度让学生模型模仿教师模型的概率分布。特征图 MSE让学生模型的中间层特征接近教师模型的中间层特征。# 文件路径kd-demo/distill.py import torch import torch.nn as nn from torch.optim import Adam import torch.nn.functional as F from models import TeacherNet, StudentNet from utils import get_mnist_loaders, evaluate def distillation_loss(student_logits, student_feats, teacher_logits, teacher_feats, labels, T4.0, alpha0.7, beta0.5): 参数说明 - T: 温度参数用于软化 logits - alpha: 硬标签损失和软标签损失的平衡权重 - beta: 特征蒸馏损失的权重 # 硬标签交叉熵 loss_hard F.cross_entropy(student_logits, labels) # 软标签 KL 散度 student_soft F.log_softmax(student_logits / T, dim1) teacher_soft F.softmax(teacher_logits / T, dim1) loss_soft F.kl_div(student_soft, teacher_soft, reductionbatchmean) * (T * T) # 特征图 MSE 损失 loss_feat 0.0 for s_feat, t_feat in zip(student_feats, teacher_feats): # 归一化后计算 MSE避免因为特征值尺度不同导致蒸馏失效 s_feat_norm F.normalize(s_feat.view(s_feat.size(0), -1), dim1) t_feat_norm F.normalize(t_feat.view(t_feat.size(0), -1), dim1) loss_feat F.mse_loss(s_feat_norm, t_feat_norm) loss_feat loss_feat / len(student_feats) # 综合损失 loss alpha * loss_hard (1 - alpha) * loss_soft beta * loss_feat return loss def train_student_by_distill(epochs10, batch_size128, lr1e-3, T4.0, alpha0.7, beta0.5, devicecpu): train_loader, test_loader get_mnist_loaders(batch_size) # 加载教师模型并冻结参数 teacher TeacherNet(num_classes10).to(device) teacher.load_state_dict(torch.load(teacher.pth, map_locationdevice)) teacher.eval() for param in teacher.parameters(): param.requires_grad False student StudentNet(num_classes10).to(device) optimizer Adam(student.parameters(), lrlr) for epoch in range(epochs): student.train() total_loss 0.0 for images, labels in train_loader: images, labels images.to(device), labels.to(device) with torch.no_grad(): teacher_logits, teacher_feats teacher(images) student_logits, student_feats student(images) loss distillation_loss( student_logits, student_feats, teacher_logits, teacher_feats, labels, TT, alphaalpha, betabeta ) optimizer.zero_grad() loss.backward() optimizer.step() total_loss loss.item() acc evaluate(student, test_loader, device) print(fEpoch {epoch 1}/{epochs}, Loss: {total_loss / len(train_loader):.4f}, Acc: {acc:.4f}) torch.save(student.state_dict(), student_distill.pth) return student if __name__ __main__: device cuda if torch.cuda.is_available() else cpu train_student_by_distill(epochs10, devicedevice)在这段代码里有几个实现细节值得重点说明。第一温度参数 T 在计算 KL 散度时需要乘上T * T。这是因为我们使用了log_softmax(z / T)和softmax(z / T)梯度会缩小1/T倍乘回来可以保持梯度尺度稳定。第二特征图对齐前进行了 L2 归一化。这是因为教师模型和学生模型的通道数不同特征值的绝对尺度可能差异很大。如果不做归一化MSE 损失可能被某一层的尺度主导训练不稳定。第三教师模型在训练过程中必须设置为eval()模式并冻结参数。这样做既节省显存也避免教师模型的 BatchNorm 统计量受到学生训练影响。4.6 对比实验不同训练方式的效果为了验证特征蒸馏的价值我们可以做一组对比实验方案 A直接训练学生模型不使用教师模型。方案 B仅使用 logits 蒸馏不加特征损失。方案 C使用 logits 特征蒸馏也就是上面的完整代码。方案 A 只需要把蒸馏脚本中的损失替换为普通交叉熵即可其他训练设置保持一致。方案 B 将distillation_loss中的beta设置为 0。下面是三组实验在 MNIST 测试集上的准确率表现方案模型参数量测试准确率教师模型约 20 万99.41%学生模型直接训练约 2.8 万98.72%仅 logits 蒸馏约 2.8 万98.93%logits 特征蒸馏约 2.8 万99.15%从结果可以看到特征蒸馏带来了接近 0.2 个百分点的提升同时让学生模型的准确率逼近甚至接近教师模型。在一个已经很好解决的数据集上这个提升幅度已经相当可观。换成更复杂的数据集或更复杂的任务这种差距通常会更大。4.7 为什么特征蒸馏有效蒸馏本质上是在传递“知识”。对于图像分类任务教师模型从原始像素中学习到的特征表示是经过多层抽象得到的。学生模型如果只模仿教师最后的概率输出它只需要学到一个“有点像教师”的决策边界而不知道教师为什么这样决策。特征蒸馏则强制学生模型在中间层就对齐教师模型的表征。这意味着学生模型需要用更少的参数去拟合教师模型的中间表示从而迫使学生模型学习到更紧凑、更高效的特征提取方式。这也是为什么那些成功的蒸馏方法比如 FitNets、Attention Transfer都选择了中间层作为蒸馏目标。5. 常见问题与排查思路5.1 特征蒸馏损失不下降如果发现loss_feat一直不下降首先检查特征图尺寸是否匹配。打印学生和教师的中间层特征形状print(student_feats[0].shape, teacher_feats[0].shape)如果尺寸不一致说明网络结构没有对齐。解决办法是在蒸馏损失之前增加一个自适应的卷积层将学生特征转换到教师特征的维度。5.2 温度参数对结果影响很大温度参数 T 过小软标签接近硬标签蒸馏退化为普通训练T 过大软标签过于均匀丢失类别间信息。建议在交叉验证中尝试T [3, 5, 8, 10]监控学生模型在验证集上的表现选择最优值。注意蒸馏时使用的温度不一定越高越好通常T 4左右是一个不错的起点。5.3 学生模型反而比直接训练更差这种情况通常是教师模型性能不够强或者教师模型与学生模型的容量差距过大。知识蒸馏的前提是教师模型至少要比直接训练的学生模型更强才有知识可以传递。另一个常见问题是蒸馏损失中硬标签权重过大导致学生模型过度依赖真实标签软标签几乎没有起到作用。可以适当降低alpha提高软标签和特征损失的权重。5.4 显存不足如果数据集图像尺寸较大中间层特征图会占用大量显存。可以采取以下措施将特征对齐层只选最后一层减少同时保存的特征数量。减小 batch size。使用梯度累计。6. 最佳实践与工程建议6.1 蒸馏场景的选择不是所有场景都需要特征蒸馏。根据我的经验可以这样选择场景推荐方案模型太大线上推理延迟高先做 logits 蒸馏快速验证效果学生模型性能离目标差一点加入特征蒸馏重点对齐中高层特征任务复杂类别数多使用多个中间层特征对齐或使用注意力蒸馏学生模型结构大幅简化需要在特征对齐层加适配模块否则难以对齐学生模型结构比教师更浅优先对齐教师的后两层浅层特征可跳过6.2 特征对齐层的设计当教师和学生结构不一致时不能直接计算特征图的 MSE。常见做法是使用一个 1x1 卷积或全连接层将学生特征投影到教师特征的空间。class Adaptor(nn.Module): def __init__(self, in_channels, out_channels): super(Adaptor, self).__init__() self.conv nn.Conv2d(in_channels, out_channels, kernel_size1) def forward(self, x): return self.conv(x)这个适配层的参数参与训练但在推理阶段会被丢弃不影响学生模型的结构。6.3 训练流程的工程化建议在实际项目中我建议按以下流程推进知识蒸馏先训练一个足够强的教师模型并保存模型权重和中间层特征。离线将教师模型的 logits 和中间层特征保存下来形成“蒸馏缓存”后续训练学生模型时直接加载不再重复前向传播教师模型。先做 logits 蒸馏建立基线。在基线基础上逐步加入特征蒸馏、注意力蒸馏观察每一步的收益。在验证集上监控学生模型的表现避免过拟合到教师模型的噪声上。这种方式的优点是显著减少训练时间。尤其在大模型蒸馏时教师模型前向传播的开销非常大提前缓存特征是很实用的手段。6.4 数据增强与正则化蒸馏不代表不需要正则化。教师模型的知识可能包含噪声尤其是教师模型在错误样本上的输出往往不可靠。实践中可以对学生模型使用 Dropout、权重衰减、标签平滑等常见正则化手段。另外在计算特征蒸馏损失时对特征做 L2 归一化后再计算 MSE比直接计算 MSE 效果更稳定。6.5 大规模生产时的注意事项如果要在生产环境使用蒸馏后的小模型需要额外关注小模型的精度是否满足业务指标。小模型在长尾样本上的表现是否退化。蒸馏后的模型是否需要重新做量化感知训练。在线推理时特征蒸馏的适配层不应该保留避免额外开销。7. 总结知识蒸馏发展至今已经远远不止“让学生模仿教师输出”这么简单。从 Hinton 提出的 logits 蒸馏到 FitNets 的特征蒸馏再到大模型时代的隐藏状态蒸馏核心都在于如何把教师模型在推理过程中形成的“经验”更完整地传递给学生。本文通过一个完整的 PyTorch 实战案例对比了直接训练、logits 蒸馏、logits 特征蒸馏三种方式的效果。实验结果表明加入中间层特征对齐后学生模型的准确率明显提升。这也验证了一个重要结论教师模型的推理习惯也就是它提取和组织特征的方式确实比最终输出的分数更有迁移价值。如果你想继续深入可以尝试以下几个方向阅读 FitNets 和 Attention Transfer 两篇经典论文了解不同的特征对齐方式。在 CIFAR-10 或你自己的业务数据集上复现本文实验观察特征蒸馏在不同任务上的表现。尝试将蒸馏与模型量化、剪枝结合形成完整的模型压缩方案。在大语言模型场景下尝试使用隐藏状态对齐和注意力分布对齐进行蒸馏。动手跑一遍代码再对比不同蒸馏目标的差异你会对“教师模型的推理习惯为何重要”有更直观的感受。如果本文对你有帮助可以收藏备用后续实践中有新的发现也欢迎一起交流。