最小复现:从设计思路到代码逐行剖析)
知识蒸馏KD最小复现从设计思路到代码逐行剖析本文承接上一篇《知识蒸馏Knowledge Distillation入门》理论篇把 Hinton 蒸馏落地成一个可跑的 PyTorch 最小复现项目。先讲清楚整个实验的代码设计思路再逐文件剖析代码最后附上实验运行结果。完整代码仓库结构models.py网络distill.py损失train.py训练run_all.py编排。目录一、实验设计思路四文件三层架构二、代码逐文件剖析2.1 models.py · 结构层2.2 distill.py · 损失层2.3 train.py · 训练引擎层2.4 run_all.py · 编排层三、实验运行与结果四、总结一、实验设计思路四文件三层架构整个项目只有一个目标验证「蒸馏学生 裸训学生」这个不等式。为了让代码清晰、可复用、可扩展分成 4 个文件遵循三条设计原则1.1 三层架构┌─────────────────────────────────────┐ │ run_all.py · 编排层 │ │ 决定训什么、按什么顺序、如何验收 │ └───────────────┬─────────────────────┘ │ import ┌───────────────▼─────────────────────┐ │ train.py · 训练引擎层 │ │ train_teacher / train_student │ └───────┬──────────────────┬──────────┘ │ import │ import ┌───────▼──────┐ ┌──────▼─────────┐ │ models.py │ │ distill.py │ │ 结构层 │ │ 损失层 │ │ TeacherNet │ │ distillation │ │ StudentNet │ │ _loss │ └──────────────┘ └────────────────┘1.2 三条设计哲学单一职责结构归结构损失归损失训练归训练编排归编排。改网络结构绝不会碰乱损失函数。复用train_student一个函数同时承担裸训和蒸馏两种任务靠teacher_path参数切换。可扩展阶段 2 换 ResNet-56/20只需改models.py其余文件零改动。1.3 数据流CIFAR-10 数据 │ ├─[1] 训练教师 ──→ teacher.pt ├─[2] 学生裸训 ──→ student_baseline.pt ├─[3] 蒸馏训练 ──→ student_distill.pt └─[4] 验收distill_acc baseline_acc ?二、代码逐文件剖析2.1 models.py · 结构层教师/学生网络教师和学生结构同构都是卷积块 分类头区别只在通道数——教师用 64/128/256学生用 32/64/128。这就是蒸馏的前提学生是教师的缩水版。classTeacherNet(nn.Module):def__init__(self,num_classes10):super().__init__()self.featuresnn.Sequential(nn.Conv2d(3,64,3,padding1),nn.BatchNorm2d(64),nn.ReLU(inplaceTrue),nn.Conv2d(64,64,3,padding1),nn.BatchNorm2d(64),nn.ReLU(inplaceTrue),nn.MaxPool2d(2),# ... 128、256 通道的卷积块逐层加宽)self.classifiernn.Sequential(nn.AdaptiveAvgPool2d(1),nn.Flatten(),nn.Linear(256,num_classes),# 输出 logits不做 softmax)defforward(self,x):returnself.classifier(self.features(x))关键点数据流3×32×32→ 卷积提特征通道 3→64→128→256→ 池化缩尺寸32→16→8→4→AdaptiveAvgPool2d(1)压成 1×1 →Flatten拉平 →Linear输出 10 维 logitsforward 输出 logits不提前 softmax——因为蒸馏要控制温度 Tsoftmax 留给损失函数get_model(name)工厂函数按字符串返回网络方便阶段 2 换 ResNet2.2 distill.py · 损失层论文公式(2) 的落地这是整个项目最有含金量的一个文件5 行代码承载论文全部思想Lα⋅T2⋅KL(qsT∥qtT)(1−α)⋅CE(qs,y) L \alpha \cdot T^2 \cdot \text{KL}(q_s^T \| q_t^T) (1-\alpha) \cdot \text{CE}(q_s, y)Lα⋅T2⋅KL(qsT∥qtT)(1−α)⋅CE(qs,y)defdistillation_loss(student_logits,teacher_logits,labels,T4.0,alpha0.7):# 目标 A软标签蒸馏soft_targetsF.softmax(teacher_logits.detach()/T,dim1)# 教师软化概率student_log_probsF.log_softmax(student_logits/T,dim1)# 学生软化 log 概率kd_lossF.kl_div(student_log_probs,soft_targets,reductionbatchmean)*(T*T)# 目标 B硬标签监督ce_lossF.cross_entropy(student_logits,labels)returnalpha*kd_loss(1.0-alpha)*ce_loss三个致命细节teacher_logits.detach()——教师冻结不参与梯度* (T * T)乘在kl_div整体上——补偿温度对梯度的 1/T² 稀释ce_loss用原始 logits 不放 T——硬标签这路是标准分类2.3 train.py · 训练引擎层提供 7 个函数set_seed锁随机数、get_loaders数据加载、evaluate算精度、build_optimizer造优化器、train_teacher、train_student、main。数据加载的关键训练集做数据增强随机裁剪翻转防止过拟合测试集只做归一化公平测试。优化器SGD(lr0.1, momentum0.9, weight_decay5e-4) 余弦退火学习率。训练循环四行骨架每个 epoch 重复这是所有神经网络训练的本质forx,yintrain_loader:optimizer.zero_grad()# ① 清空上一步梯度梯度会累加losscriterion(model(x),y)# ② 前向传播 算损失loss.backward()# ③ 反向传播自动求导optimizer.step()# ④ 更新参数蒸馏分支train_student的核心靠teacher_path切换outmodel(x)# 学生前向ifteacherisnotNone:withtorch.no_grad():# 教师前向不参与梯度t_outteacher(x)lossdistillation_loss(out,t_out,y,TT,alphaalpha)# 蒸馏else:losscriterion(out,y)# 裸训普通交叉熵2.4 run_all.py · 编排层teacher_acctrain_teacher(EPOCHS,BATCH_SIZE,device,save_pathteacher.pt)baseline_acctrain_student(EPOCHS,BATCH_SIZE,device,save_pathstudent_baseline.pt)distill_acctrain_student(EPOCHS,BATCH_SIZE,device,teacher_pathteacher.pt,TT,alphaALPHA,save_pathstudent_distill.pt)print(f蒸馏提升: {distill_acc-baseline_acc:.2f}个百分点)ifdistill_accbaseline_acc:print(验收通过蒸馏学生 裸训学生复现成功)它自己不训练任何模型只决定顺序训教师 → 裸训学生 → 蒸馏学生 → 打印三行精度做验收。三、实验运行与结果环境PyTorch torchvisionCIFAR-10CPU 训练EPOCHS7快速验证版。3.1 训练教师网络python train.py --mode teacher的训练过程输出展示每个 epoch 的 loss 和 accuracy 变化最终保存 teacher.pt3.2 学生裸训baselinepython train.py --mode baseline的输出学生不蒸馏、纯交叉熵训练的精度作为对照基线3.3 蒸馏训练学生python train.py --mode distill的输出加载教师、用蒸馏损失训练学生的精度变化3.4 验收对比python run_all.py最后的复现验收四行输出教师/裸训/蒸馏三行精度 蒸馏提升幅度这是整个实验的最终结论验收判据只要「学生蒸馏精度 学生裸训精度」成立就证明蒸馏有效复现成功。四、总结本文把 Hinton 蒸馏落地为一个最小可跑的 PyTorch 项目核心脉络架构四文件三层结构 / 损失 / 训练 / 编排单一职责、可复用、可扩展灵魂distill.py里 5 行代码承载公式(2) 的全部思想本质训练循环的四行骨架zero_grad → forward → backward → step验收蒸馏精度 裸训精度理解这套代码后阶段 2 换 ResNet、阶段 3 做改进改 distill.py 的损失都是在这个骨架上换零件主线不会变。本文是作者复现 Hinton 2015 论文的代码学习笔记如有理解不当之处欢迎指正交流。