
“什么时候也能蒸馏我自己”如果你在训练完模型之后冒出过这个念头说明你已经对知识蒸馏的固定剧本产生了怀疑为什么每一次搬知识都得先造一个更大的老师这个想法并不是段子。研究者早就把它做成了正经课题也就是自蒸馏Self-Distillation。它不额外训练一个巨大的教师模型而是让同一个模型的不同部分、不同训练阶段互相扮演师生。换句话说蒸馏关系中最重要的师生角色可以由一个主体自己包圆。我的判断是自蒸馏看起来像是“小模型单靠意志力变强”本质上却是一种更高效的训练监督机制。它在推理阶段不增加参数量、不引入部署成本只在训练时给网络注入更丰富的监督信号。对于没有大算力、没有预训练超大模型可用的团队来说这是一个性价比非常高的优化手段。这篇文章会从知识蒸馏的原理开始梳理三种常见的自蒸馏范式再给出一套可以直接运行的 PyTorch 自蒸馏训练代码以 CIFAR-10 分类为例跑通“自己蒸馏自己”的完整流程。最后会重点讨论哪些场景真正适合自蒸馏哪些场景不要想当然地用。这部分判断往往比代码本身更有价值。1. 为什么会有“蒸馏我自己”这个奇怪问题知识蒸馏的标准流程是两阶段师生训练。先准备一个表达能力很强的教师模型通常是一个大网络或者大模型然后用教师输出的软标签去监督学生模型的训练让学生用更少的参数逼近教师的预测能力。这套流程在移动端模型、边缘推理模型上已经非常常见比如把一个大分类模型蒸馏成几 MB 的量化小模型。但真正落到工程里这套流程有几个并不舒服的前提。第一个痛点是没有大老师可用。很多中小团队的业务场景是自研小模型而不是站在开源大模型肩膀上。如果团队没有训练过大模型也不可能让大模型对这个业务领域产生专门认知就根本没有高质量的教师输出。蒸馏这件事无从谈起。第二个痛点是大老师太贵。即使团队有实力训练大模型也要算清这笔账训练教师的算力成本、存储成本、教师推理生成软标签的时间成本以及中间如果教师更新版本学生要不要重新蒸馏一遍。对于快速迭代的线上模型这套路径的维护成本不小。第三个痛点是教师的“知识”未必是学生能消化的。容量差距过大的师生对学生很容易丢失细粒度信息。教师模型的输出分布再准学生模型容量不够时也只能学到一部分。研究者后来发现当教师和学生结构一致甚至同一个网络分支时软标签传递效率反而更高因为特征的抽象层次和分布更接近。这三个痛点叠加产生了一个很自然的动机如果教师不是外部的庞然大物而是模型自身更深的分支或上一轮的自己就能绕开大部分成本问题。“自蒸馏”这个概念就是在这种背景下被正式命名的。所以自蒸馏并不是简单地把知识蒸馏的教师去掉而是重新设计了一条不需要额外教师的监督路径。2. 先搞懂知识蒸馏老师-学生的基本框架要理解自蒸馏得先把知识蒸馏的师生框架看透。普通分类任务训练时损失函数一般采用交叉熵监督信号是 one-hot 硬标签。比如一张图是猫标签就是[0, 1, 0]这种形式模型只需要学会把正确类别的分数拉高。知识蒸馏的核心变化是引入了软标签。教师模型对样本输出的不是 one-hot而是一个概率分布。比如某张猫的图片教师可能输出“猫 0.6、狗 0.3、汽车 0.1”。这个分布中藏着更丰富的知识猫和狗的视觉特征比猫和汽车更接近。one-hot 标签完全丢失了这类类间关系信息而软标签把它们保留下来了。为了让软标签更有区分度蒸馏通常会引入温度系数 T。计算时先把逻辑值除以 T再做 softmax分布就会变得更平滑在损失计算中再把梯度量级乘回 T 的平方保证学习尺度不失控。蒸馏损失函数的基本形式如下L CE(student_logits, hard_label) alpha * KL(softmax(teacher_logits / T) || softmax(student_logits / T)) * T^2第一部分是标准的硬标签交叉熵保证学生任务方向正确第二部分是蒸馏损失让学生向教师输出的软分布靠拢alpha 用来控制两部分权重。下面这张表可以快速对比普通训练和知识蒸馏训练的关键差异对比维度标准监督训练知识蒸馏训练标签形式one-hot 硬标签硬标签 教师软标签信息量只含样本类别答案额外包含类别间相似关系教师需求不需要需要预训练教师模型训练成本低额外训练和推理教师典型效果取决于模型容量小模型可逼近大模型能力理解了这条主线再看自蒸馏就容易得多只要把上面公式中的 teacher 换成模型自身的一部分或历史版本蒸馏从“老师-学生”结构就变成了“自己-自己”结构。3. 自蒸馏到底在做什么自蒸馏的定义并不复杂训练过程中不需要外部教师模型而是由目标模型自身或其自身衍生物产生软标签用来指导网络某个部分的训练。真正复杂的是它的实现形态因为“自身”可以有不同的切分方式。目前工程和研究中最常见的自蒸馏有三种第一种是深度监督式自蒸馏。在隐藏层后面挂辅助分类器让网络更深层的主分类器输出软标签去教浅层的辅助分类器。此时教师是同一个网络的主干部分学生是同一个网络的浅层分支。第二种是重生网络范式英文叫 Born-Again Networks。分两个阶段执行先用常规监督训练训好一个模型然后用这个已训好的模型作为教师重新初始化一个同结构模型作为学生再用软标签加上硬标签重训一遍。这里的教师是“上一轮的自己”。第三种是时序自蒸馏或移动平均教师。训练过程中维护一组模型参数的指数移动平均EMA把 EMA 版本当作教师用它的输出去教当前训练中的模型参数。教师是“时间维度上的自己”。三种范式可以汇总成一张对比表范式教师来源学生是谁训练形态传统知识蒸馏外部大模型小模型离线两阶段深度监督自蒸馏同网络深层分支同网络浅层分支单次训练重生网络上一轮自己同架构新模型两阶段重训时序自蒸馏(EMA)历史参数平均值当前模型参数在线训练用学习来类比的话传统蒸馏是“请名师补课”自蒸馏更像是“整理错题本和分层自查”。一个人没有额外家教也可以通过反复回顾自己的做题路径、让不同层次的思维互相校准获得提升。自蒸馏纠偏的重点是让梯度信号能够更早、更平滑地注入网络的浅层部分。这里有一个很容易产生的误解需要提前澄清自蒸馏并不意味着模型会突然获得外部知识。它的提升空间始终受制于模型本身的表达能力和训练数据的覆盖范围。自蒸馏能做的是把监督信号利用得更充分而不是无中生有地造出更强的模型。4. 深度监督式自蒸馏的完整原理这篇文章的实操部分选择深度监督式自蒸馏因为它只需要一次训练不需要额外保存和加载教师模型最适合作为入门范例。深度监督本身并不是新东西它最早是为了解决深层网络梯度消失问题。当网络层数很深时反向传播的梯度从最后一层传到前面的层已经很微弱浅层特征得不到有效更新。早期的做法是在中间层挂辅助分类器给浅层单独补一份交叉熵损失这就是深度监督。自蒸馏在深度监督基础上做了关键升级辅助分类器不只吃硬标签还吃主分类器输出的软标签。深层的主分类器拥有更抽象的特征它生成的软分布中包含着类别关系这部分知识没法直接通过硬标签传回浅层现在通过蒸馏损失软标签也能沿着更短的路径影响浅层特征。损失函数的设计顺势变成了三部分L CE(main_logits, hard_label) CE(aux_logits, hard_label) alpha * KL(P_main / T, P_aux / T) * T^2主分类器负责承担主要任务辅助分类器继续保留自己的硬标签任务防止它变成纯粹模仿主分类器的摆设蒸馏损失则强制辅助分类器向主分类器的软分布对齐。这个结构为什么有效可以从三个角度理解。第一是梯度路径缩短。辅助分类器从主干中间层引出它的梯度可以直接作用于 conv2 等浅层参数不再需要跨越整条主干链。网络越深这条短路径的收益越明显。第二是监督信息更平滑。one-hot 标签把非正确类别的信息全部归零而主分类器的软标签保留着“哪些类别更像”的信息。浅层特征从这里学到的约束比硬标签更细腻。第三是隐式的多任务学习。辅助分支相当于给共享主干增加了一个弱分类任务这种多任务压力会让特征提取器学习到更通用的表示而不是只服务单一分类头。在实际实现中主分类器和辅助分类器共享特征提取层。两者输出 softmax 之后得到的不是观测值而是同一个前向过程的两个视角自蒸馏的“自”就体现在这里。5. 环境准备与数据集选择开始写代码之前先确认环境。本文的示例基于 PyTorch选用 CIFAR-10 分类数据集原因是它规模适中、类别明确对小网络来说既不会在几分钟内过拟合也能在普通 GPU 上快速跑出结果。系统方面Linux、macOS、Windows 都可以操作GPU 不是必须纯 CPU 也可以完成一次小规模演示只是训练时间会拉长。会更推荐有 NVIDIA GPU 的环境配合 CUDA 训练体验正常得多。Python 版本建议 3.8 及以上PyTorch 安装 2.x 版本同时安装配套的 torchvision。需要注意不同版本之间 API 基本兼容但为了减少环境问题建议参考 PyTorch 官网的命令安装可以根据操作系统选择 CPU 或 CUDA 版本本文不会刻意依赖某个新版本专属接口。数据集方面代码中会配置downloadTrue首次运行自动下载 CIFAR-10 到本地./data目录。如果服务器访问国外下载源较慢可以先用其他方式下载 CIFAR-10 的压缩包放到./data目录并保持目录结构再用代码加载。建议的工程目录结构如下self_distill_project/ ├── self_distill_net.py ├── self_distill_loss.py ├── train_self_distill.py └── data/三个 Python 文件分别负责网络结构、蒸馏损失函数和训练流程。这种拆分方式在后续换数据集、换网络时会让改动集中在单一文件内比较适合做实验迭代。6. 核心代码实现深度监督式自蒸馏训练下面进入可以运行的部分。整个示例拆成三个文件先写网络结构再写蒸馏损失最后写训练脚本。6.1 网络结构定义文件路径self_distill_net.py为了让自蒸馏的效果直观网络必须包含两条分类路径一条是主分类器一条是辅助分类器。主分类器从网络最后一层特征引出辅助分类器从第二层卷积之后引出。import torch import torch.nn as nn import torch.nn.functional as F class SelfDistillNet(nn.Module): def __init__(self, num_classes10, temperature4.0): super(SelfDistillNet, self).__init__() self.temperature temperature # 共享主干特征提取层 self.conv1 nn.Conv2d(3, 32, 3, padding1) self.bn1 nn.BatchNorm2d(32) self.conv2 nn.Conv2d(32, 64, 3, padding1) self.bn2 nn.BatchNorm2d(64) self.pool1 nn.MaxPool2d(2) # 32x32 - 16x16 self.conv3 nn.Conv2d(64, 128, 3, padding1) self.bn3 nn.BatchNorm2d(128) self.pool2 nn.MaxPool2d(2) # 16x16 - 8x8 # 辅助分类器挂在第二层卷积之后 self.aux_pool nn.AdaptiveAvgPool2d((4, 4)) self.aux_head nn.Sequential( nn.Flatten(), nn.Linear(64 * 4 * 4, 256), nn.ReLU(inplaceTrue), nn.Linear(256, num_classes) ) # 主分类器挂在最后一层特征之后 self.main_head nn.Sequential( nn.Flatten(), nn.Linear(128 * 8 * 8, 256), nn.ReLU(inplaceTrue), nn.Linear(256, num_classes) ) def forward(self, x): x F.relu(self.bn1(self.conv1(x))) x F.relu(self.bn2(self.conv2(x))) aux_feat x # 辅助分类器从这里分支出来 x self.pool1(x) x F.relu(self.bn3(self.conv3(x))) main_feat self.pool2(x) main_logits self.main_head(main_feat) aux_feat self.aux_pool(aux_feat) aux_logits self.aux_head(aux_feat) return main_logits, aux_logits这段代码的关键点在于forward的两次返回main_logits和aux_logits分别代表主路径和辅助路径的分类输出。训练时两个输出都会使用但在推理阶段我们只关心main_logits辅助分类器会在导出模型时移除。6.2 蒸馏损失函数文件路径self_distill_loss.py蒸馏损失按照第 2 节讲的原则实现先把双方的 logits 除以温度再做 KL 散度最后乘回温度平方。import torch.nn.functional as F def distillation_loss(student_logits, teacher_logits, temperature4.0): # 温度缩放 student_logits student_logits / temperature teacher_logits teacher_logits / temperature loss F.kl_div( F.log_softmax(student_logits, dim1), F.softmax(teacher_logits, dim1), reductionbatchmean ) # 温度平方补偿保持梯度量级稳定 return loss * temperature * temperature这里使用batchmean而不是sum或mean是为了让 KL 损失在 batch 维度上取平均数值相对稳定这也是 PyTorch 官方在实现蒸馏时常选择的 reduction 方式。6.3 完整训练脚本文件路径train_self_distill.py训练脚本包含数据加载、模型实例化、训练循环和测试评估。代码中设置了一个use_distill开关把它设为True就是自蒸馏训练设为False就是普通的“主分类器 辅助分类器”多损失训练方便做消融对比。import torch import torch.nn as nn import torch.optim as optim import torchvision import torchvision.transforms as transforms from torch.utils.data import DataLoader from self_distill_net import SelfDistillNet from self_distill_loss import distillation_loss def evaluate(model, test_loader, device): model.eval() correct 0 total 0 with torch.no_grad(): for images, labels in test_loader: images, labels images.to(device), labels.to(device) main_logits, _ model(images) _, predicted torch.max(main_logits, 1) total labels.size(0) correct (predicted labels).sum().item() return correct / total def main(): device torch.device(cuda if torch.cuda.is_available() else cpu) print(fUsing device: {device}) # 训练开关True 为自蒸馏False 为普通多损失训练 use_distill True alpha 0.7 num_epochs 20 batch_size 128 transform_train transforms.Compose([ transforms.RandomCrop(32, padding4), transforms.RandomHorizontalFlip(), transforms.ToTensor(), transforms.Normalize((0.4914, 0.4822, 0.4465), (0.2023, 0.1994, 0.2010)), ]) transform_test transforms.Compose([ transforms.ToTensor(), transforms.Normalize((0.4914, 0.4822, 0.4465), (0.2023, 0.1994, 0.2010)), ]) train_set torchvision.datasets.CIFAR10( root./data, trainTrue, downloadTrue, transformtransform_train) test_set torchvision.datasets.CIFAR10( root./data, trainFalse, downloadTrue, transformtransform_test) train_loader DataLoader(train_set, batch_sizebatch_size, shuffleTrue, num_workers2) test_loader DataLoader(test_set, batch_size256, shuffleFalse, num_workers2) model SelfDistillNet(num_classes10, temperature4.0).to(device) optimizer optim.SGD(model.parameters(), lr0.1, momentum0.9, weight_decay5e-4) scheduler optim.lr_scheduler.CosineAnnealingLR( optimizer, T_maxnum_epochs) ce_loss nn.CrossEntropyLoss() for epoch in range(num_epochs): model.train() running_loss 0.0 correct_main 0 correct_aux 0 total 0 for images, labels in train_loader: images, labels images.to(device), labels.to(device) main_logits, aux_logits model(images) # 硬标签损失 loss_main ce_loss(main_logits, labels) loss_aux ce_loss(aux_logits, labels) if use_distill: # 辅助分类器学习主分类器输出的软标签 loss_distill distillation_loss( aux_logits, main_logits.detach(), model.temperature ) loss loss_main loss_aux alpha * loss_distill else: loss loss_main loss_aux optimizer.zero_grad() loss.backward() optimizer.step() running_loss loss.item() * images.size(0) _, pred_main torch.max(main_logits, 1) _, pred_aux torch.max(aux_logits, 1) total labels.size(0) correct_main (pred_main labels).sum().item() correct_aux (pred_aux labels).sum().item() scheduler.step() test_acc evaluate(model, test_loader, device) print( fEpoch [{epoch 1}/{num_epochs}] fLoss: {running_loss / total:.4f} fTrain Main Acc: {correct_main / total:.4f} fTrain Aux Acc: {correct_aux / total:.4f} fTest Acc: {test_acc:.4f} ) torch.save(model.state_dict(), self_distill_model.pth) print(Model saved to self_distill_model.pth) if __name__ __main__: main()6.4 代码逻辑说明整个训练脚本的核心逻辑其实只有三步。第一步是前向计算。图像进入网络后同时得到主分类器输出和辅助分类器输出。主输出来自更深的特征理论上特征更抽象辅助输出来自中间层特征训练初期的稳定性会略差。第二步是损失组装。主分类器使用硬标签交叉熵辅助分类器也使用硬标签交叉熵如果启用了自蒸馏则额外加入辅助分类器向主分类器软分布对齐的蒸馏损失。因为辅助分类器同时承担硬标签和软标签两个任务它不会变成一个只模仿主分类器的空壳。第三步是梯度更新与评估。优化器用标准的 SGD 加余弦退火学习率。每个 epoch 结束后用当前模型评估测试集打印主分支和辅助分支的训练精度。主分支的训练精度和测试精度是判断模型最终质量的核心指标。值得强调的是main_logits.detach()这个操作。蒸馏损失要改变的是学生辅助分类器和共享主干的表示不希望梯度直接通过主分类器的输出回流到主分支导致教师信号在优化过程中不断漂移。虽然主分类器和辅助分类器共享大部分卷积参数本身会对教师产生影响但从单步损失计算的角度detach()能避免“教师在训练中被自身学生反复拉扯”的震荡问题是一个简洁而稳妥的做法。7. 运行结果与验证方式在项目根目录执行python train_self_distill.py如果环境正常首先会看到设备信息和 CIFAR-10 下载提示然后开始训练。运行日志的形态大致如下不同环境数值会有差异Using device: cuda Files already downloaded and verified Epoch [1/20] Loss: 2.0310 Train Main Acc: 0.3814 Train Aux Acc: 0.3092 Test Acc: 0.5123 Epoch [5/20] Loss: 1.2865 Train Main Acc: 0.6337 Train Aux Acc: 0.5810 Test Acc: 0.6745 Epoch [10/20] Loss: 0.9412 Train Main Acc: 0.7804 Train Aux Acc: 0.7218 Test Acc: 0.7496 Epoch [15/20] Loss: 0.7412 Train Main Acc: 0.8664 Train Aux Acc: 0.8103 Test Acc: 0.7831 Epoch [20/20] Loss: 0.6887 Train Main Acc: 0.9105 Train Aux Acc: 0.8672 Test Acc: 0.8221判定自蒸馏是否生效可以看两个信号。第一个信号是辅助分支精度的上升曲线。如果自蒸馏真正起作用辅助分支的训练精度会在训练中期明显追向主分支而不是一直停留在较低水平。这说明深层信息通过软标签有效回流到了浅层分类器。第二个信号是与普通训练的横向对比。把use_distill改为False其他配置完全不变再跑一遍同样的 20 个 epoch比较两次测试集精度。自蒸馏的价值不在于每次都必须“碾压”普通训练而是通常会在收敛速度或最终精度上带来正向收益。如果两种模式结果差别很小说明当前数据量和模型容量对这个小任务已经足够自蒸馏这种正则化手段没有发挥空间并不代表代码有问题。更精细的验证方式是保存训练过程中主分支和辅助分支的 loss画出学习曲线。正常情况下启用自蒸馏后辅助分支 loss 下降会更平缓主分支测试精度曲线前期的震荡也会更小因为软标签本质上给训练提供了一个更平滑的监督目标。8. 常见问题与排查思路自蒸馏代码本身不复杂但实际运行中仍然有不少容易踩的坑。下面这些问题是我认为最值得提前掌握的问题现象可能原因排查方式解决方案训练 loss 出现 NaN温度 T 过小导致 softmax 数值溢出或学习率过高检查 loss 数值在哪个 epoch 开始异常打印 logits 范围提高温度至 3 以上降低学习率确认输入数据没有 NaN辅助分支精度一直很低辅助分类器容量不足或分支位置太浅观察 aux 分支是否随 epoch 缓慢上升增加辅助分类器隐层维度或把分支位置后移到更深特征层自蒸馏和普通训练效果几乎一样模型对当前任务容量过剩数据偏简单对比两条学习曲线确认已经跑满足够 epoch换更难的数据集或加深主干也可以增大 alpha 观察敏感性显存不足辅助头增加了中间特征使用量看报错发生在哪个 forward 阶段减小 batch size减少辅助分类器通道数关闭 grad checkpoint推理模型权重变大部署时把辅助分类器也保留了下来导出模型前查看 state_dict 键名导出前过滤掉aux_开头的参数或重建网络后只加载main相关权重Windows 下 DataLoader 报错num_workers在多进程环境下有问题查看进程启动和线程报错信息将num_workers调为 0或把训练脚本放到if __name__ __main__保护块中第一个问题最常见也是初学者最头痛的。KL 散度中的log_softmax在温度很小的时候输入的 logits 数值会被放大到一个很大的范围特别是 CIFAR-10 这种多分类任务极端 logits 经过指数运算后很容易溢出。所以温度一般从 3 或者 4 开始尝试而不是从 1 开始强行调小。第二个问题需要结合任务判断。辅助分类器本质上是给主干“搭”出来的一个额外监督头如果它太浅学到的特征本身就不具备判别性哪怕有软标签也提升有限。一个实用的方案是辅助分支的输入特征至少要有两个卷积块的抽象程度并且其全连接层容量不能比主分类头小太多。第三个问题经常被误判为“自蒸馏没用”其实恰恰是自蒸馏的适用范围问题。当模型容量大于数据复杂度时普通训练已经可以把有效信息完全吸收自蒸馏这种监督信号增强手段自然看不到收益。这种情况下不需要强行调参换成更难的数据集或更深的网络再对比效果差异就会明显。9. 最佳实践与工程建议自蒸馏虽然听起来很轻巧但要用到实际项目里还是有不少工程细节值得讲究。第一辅助头的位置不要贪多。对常见的卷积网络在总深度三分之一到二分之一的位置挂 1 到 2 个辅助分类器比较稳妥。辅助头挂得太多不仅显存开销上升也会让主干在多个方向被拉扯出现训练不稳定。对 ResNet 这类带残差结构的网络可以在第 2 个或第 3 个 stage 之后插入辅助头。第二温度 T 和蒸馏权重 alpha 要一起调。温度影响软标签的平滑程度alpha 影响蒸馏损失在总损失中的比例。通常 T 在 3 到 6 之间比较常用alpha 从 0.5 起步观察辅助分支的收敛情况后再调整。如果 alpha 太低辅助分支学不到软标签信息太高则会压制硬标签导致辅助分支在类别边界处过度平滑、精度下降。第三推理阶段一定要裁剪辅助分支。自蒸馏的价值集中在训练阶段部署时模型只需要主分类器。保存模型时可以先把辅助分支从 state_dict 中过滤掉或者新建网络结构后只加载主分类器和主干参数。不裁剪的话推理时会白白多计算辅助分支部分的 FLOPs功耗和延迟都不划算。第四可以和其他知识蒸馏叠加使用。自蒸馏不是只能单独出现。如果团队恰好有合法、合规的大模型教师也可以把外部教师的软标签用于主分支同时让主分支继续向辅助分支传递软标签。这样形成“外部教师教主分支、主分支教辅助分支”的多层监督结构。不过实践中的收益并不一定随层数线性增加更推荐先做内部自蒸馏消融确认收益后再引入外部教师。第五合理设置随机种子和日志记录。自蒸馏实验对比的是同一个网络在“有没有自蒸馏”下的差异而不是对比不同随机种子下的运气。训练前固定random、numpy、torch的随机种子并把每个 epoch 的 loss、主分支精度、辅助分支精度都写入日志文件。这样后续判断调参方向时才不会靠回忆。第六注意自蒸馏的适用边界。它的本质决定了它只能优化表达和监督不能让一个欠拟合的小网络凭空获得大模型级别的能力。以下情况不要指望自蒸馏训练数据本身噪声极大、标签质量不可靠任务对单个特征分组极度敏感需要外部先验知识模型已经严重欠拟合且需要增加容量而不仅仅改善监督信号。从工程上看自蒸馏真正适合的角色是“在既定训练管线中低成本增加训练效率”。它不是模型架构大改不需要额外准备大模型和算力只增加一块蒸馏损失和一个小小的辅助分支几十行代码就能接入现有训练脚本属于典型的低成本高回报尝试。10. 总结与下一步实践回到最开始那个问题什么时候可以“蒸馏我自己”答案是现在就可以而且只需要一部训练脚本和一张数据集。这篇文章把知识蒸馏的基本框架、自蒸馏的三种经典范式以及深度监督式自蒸馏的完整实现都走了一遍。代码层面我们给网络加了一个辅助分类器用主分支的软标签去教辅助分支并且保留硬标签交叉熵形成“自蒸馏 深度监督”的混合训练模式。打开use_distill开关就能跑通完整训练关闭它就能做对比实验。下一步的实践路径可以从三个方向展开。第一个方向是复现和消融。先把use_distill保留为可配置参数在 CIFAR-10 上分别启用和关闭跑 20 个 epoch记录测试精度和收敛曲线。理解这个对比之后你才算真正掌握了自蒸馏的调参手感。第二个方向是换结构验证。辅助分类器不只适用于小 CNN。可以把它接到 ResNet 的某个 stage 后或者在 Transformer 编码器的中间层后面加分类头观察自蒸馏在更现代网络上的表现。原理完全一致效果会有差异值得在不同数据集上做消融。第三个方向是结合部署需求。把裁剪辅助分支、导出主分支、量化小模型这三步串起来用自蒸馏作为训练阶段优化用裁剪和量化作为推理阶段优化。这种“训练优化 推理压缩”的组合是中小团队很实用的模型瘦身路径。最后给个直接可操作的建议把你手头正在训练的小模型先找一张不敏感的数据集复现一遍上面的脚本保存一组带自蒸馏和一组不带自蒸馏的日志。实践一次之后你会很清楚“蒸馏我自己”到底能带来多少收益也能判断它在你自己的业务场景中到底该不该用。