☰
自蒸馏原理与PyTorch实战:自己蒸馏自己
2026/9/28 1:14:59 网站建设 项目流程

“什么时候,也能蒸馏我自己?”如果你在训练完模型之后冒出过这个念头,说明你已经对知识蒸馏的固定剧本产生了怀疑:为什么每一次搬知识,都得先造一个更大的老师?

这个想法并不是段子。研究者早就把它做成了正经课题,也就是自蒸馏(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 版本,本文不会刻意依赖某个新版本专属接口。

数据集方面,代码中会配置download=True,首次运行自动下载 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_classes=10, temperature=4.0): super(SelfDistillNet, self).__init__() self.temperature = temperature # 共享主干特征提取层 self.conv1 = nn.Conv2d(3, 32, 3, padding=1) self.bn1 = nn.BatchNorm2d(32) self.conv2 = nn.Conv2d(32, 64, 3, padding=1) self.bn2 = nn.BatchNorm2d(64) self.pool1 = nn.MaxPool2d(2) # 32x32 -> 16x16 self.conv3 = nn.Conv2d(64, 128, 3, padding=1) 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(inplace=True), nn.Linear(256, num_classes) ) # 主分类器,挂在最后一层特征之后 self.main_head = nn.Sequential( nn.Flatten(), nn.Linear(128 * 8 * 8, 256), nn.ReLU(inplace=True), 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, temperature=4.0): # 温度缩放 student_logits = student_logits / temperature teacher_logits = teacher_logits / temperature loss = F.kl_div( F.log_softmax(student_logits, dim=1), F.softmax(teacher_logits, dim=1), reduction="batchmean" ) # 温度平方补偿,保持梯度量级稳定 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(f"Using device: {device}") # 训练开关:True 为自蒸馏,False 为普通多损失训练 use_distill = True alpha = 0.7 num_epochs = 20 batch_size = 128 transform_train = transforms.Compose([ transforms.RandomCrop(32, padding=4), 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", train=True, download=True, transform=transform_train) test_set = torchvision.datasets.CIFAR10( root="./data", train=False, download=True, transform=transform_test) train_loader = DataLoader(train_set, batch_size=batch_size, shuffle=True, num_workers=2) test_loader = DataLoader(test_set, batch_size=256, shuffle=False, num_workers=2) model = SelfDistillNet(num_classes=10, temperature=4.0).to(device) optimizer = optim.SGD(model.parameters(), lr=0.1, momentum=0.9, weight_decay=5e-4) scheduler = optim.lr_scheduler.CosineAnnealingLR( optimizer, T_max=num_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( f"Epoch [{epoch + 1}/{num_epochs}] " f"Loss: {running_loss / total:.4f} " f"Train Main Acc: {correct_main / total:.4f} " f"Train Aux Acc: {correct_aux / total:.4f} " f"Test 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 编码器的中间层后面加分类头,观察自蒸馏在更现代网络上的表现。原理完全一致,效果会有差异,值得在不同数据集上做消融。

第三个方向是结合部署需求。把裁剪辅助分支、导出主分支、量化小模型这三步串起来,用自蒸馏作为训练阶段优化,用裁剪和量化作为推理阶段优化。这种“训练优化 + 推理压缩”的组合是中小团队很实用的模型瘦身路径。

最后给个直接可操作的建议:把你手头正在训练的小模型先找一张不敏感的数据集,复现一遍上面的脚本,保存一组带自蒸馏和一组不带自蒸馏的日志。实践一次之后,你会很清楚“蒸馏我自己”到底能带来多少收益,也能判断它在你自己的业务场景中到底该不该用。

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询