ARTICLE DETAIL

资讯详情

深耕编程入门与网站建设的一线实战洞察。

IRG知识蒸馏实战:通道关系图让ResNet18继承ResNet50能力

IRG知识蒸馏实战:通道关系图让ResNet18继承ResNet50能力 简介知识蒸馏IRG算法实战项目面向深度学习初学者与模型压缩研究者聚焦利用ResNet50作为教师网络、ResNet18作为学生网络完成知识蒸馏的全流程实现帮助读者理解教师—学生框架下的特征图关系对齐与损失函数设计。压缩包约930.95MB含2000余个文件其中2406张png可视化图片覆盖训练过程与结果对比7个py文件提供完整工程代码4个json文件保存学生模型、教师模型及蒸馏结果指标另有pyc与txt辅助文件便于直接运行和结果核对。已有720人学习下载适合想亲手跑通IRG算法并深入原理的读者。获取后可得到可复现的Python工程、训练配置与蒸馏结果记录还能借助可视化图片直观分析特征关系建模过程并以此为基础迁移到其他网络结构或数据集继续扩展自己的蒸馏实验。1. 从 ResNet50 到 ResNet18为什么 IRG 是视觉蒸馏里最值得先试的算法先看一个很反直觉的实验结论你用 ResNet50 当老师直接拿 logits 蒸馏学生的准确率可能只涨 0.5 到 1 个百分点但如果你把 ResNet50 中间层输出的特征图抽出来算成一张张「通道关系图」再拿这些关系图去约束 ResNet18效果往往能再翻一倍。这就是 IRGInter-channel Relation Graph算法最核心的卖点它不走传统的 soft label 路线而是让学生的神经网络内部结构去模仿老师神经网络内部结构的组织方式。知识蒸馏并不是简单的「大模型教小模型抄答案」真正有价值的信息藏在特征通道之间的关系里。这篇文章就用一个完整可复现的最小案例把 IRG 算法从原理、特征出口对齐、损失函数到训练与验证全部跑通。适合正在做模型压缩、端侧部署或者想把 ResNet 系模型换成轻量化版本但不想掉点的人。2. IRG 算法的核心拆解通道关系图与蒸馏损失设计2.1 单层特征图上的关系矩阵计算IRG 这个名字里的 Inter-channel 指的是特征图通道之间的关联Relation Graph 则指把这些关联组织成图结构。要理解它先看一张特征图张量长什么样假设输入一张图片经过某个 stage 之后我们拿到一个形状为[B, C, H, W]的特征图。传统特征蒸馏会直接让学生的[B, C, H, W]去逼近老师的[B, C, H, W]但这里有两个问题一是老师学生的通道数可能不一样比如 ResNet50 第一个 stage 输出 256 通道ResNet18 只有 64 通道没法直接对齐二是逐像素逼近会把大量噪声和背景纹理也强迫学生学到学生的容量本来就小全学只会变得僵硬。IRG 的做法是把特征图从[B, C, H, W]处理成[B, C, C]的通道关系矩阵。具体做法是先把 H 和 W 两个维度全局平均池化掉得到[B, C, 1]的向量然后拿这个向量和它自己的转置做矩阵乘法。用代码表示就是import torch import torch.nn.functional as F def compute_relation_map(feature_map): 输入: feature_map, 形状 [B, C, H, W] 输出: relation_map, 形状 [B, C, C] B, C, H, W feature_map.shape # 全局平均池化把空间维度压成 1 pooled F.adaptive_avg_pool2d(feature_map, (1, 1)) # [B, C, 1, 1] pooled pooled.view(B, C) # [B, C] # 计算 C x C 的相关性矩阵这里用点积相似度 relation_map torch.bmm(pooled.unsqueeze(2), pooled.unsqueeze(1)) # [B, C, C] return relation_map这段代码里torch.bmm是批量矩阵乘法pooled.unsqueeze(2)变成[B, C, 1]pooled.unsqueeze(1)变成[B, 1, C]两者相乘就得到每个通道两两之间的点积相似度。点积相似度虽然不是归一化的但在蒸馏场景里没关系因为我们后面用的损失函数是均方误差关注的是学生和老师的相对组织模式是否一致。为什么要加adaptive_avg_pool2d这一步因为 IRG 希望关系矩阵具备空间平移不变性一张猫的图片猫在左上角还是右下角通道之间的激活关系应该大致一样。如果直接用 HxW 的所有像素去算关系矩阵计算量会爆炸而且会引入太多空间位置信息。池化到 1x1 之后计算一个 batch 的关系矩阵成本可以忽略。这一层输出的[B, C, C]矩阵就是单层的 relation graph。2.2 多层特征对齐与蒸馏损失的加权单层关系矩阵只是 IRG 的一部分。实际使用中我们不会只抽取某一层而是从老师和学生的网络中对应位置抽出多组特征图逐层计算关系矩阵并把每一层的蒸馏损失按权重加起来。在 ResNet50 蒸馏 ResNet18 的场景里我一般会在每个 stage 的最后一个残差块之后抽特征也就是 stage1、stage2、stage3、stage4 各抽一次形成四个蒸锚点。损失函数由三部分组成但权重分配是关键。基础的损失是「学生自己特征图经过映射后和老师特征图经过映射后的关系矩阵之间的 L2 距离」。这里有个细节关系矩阵不一定是方阵才能比较如果老师和学生通道数不同关系矩阵尺寸确实不同所以需要把老师或学生的关系矩阵先映射到同一维度。常见做法是加一个可学习的 1x1 卷积或全连接层把学生的关系矩阵投影到老师的维度。我常用的损失形式是L_irg lambda_1 * || R_T - R_S ||_2 lambda_2 * || R_T - R_T_hat ||_2 lambda_3 * || R_S - R_S_hat ||_2中间两项是针对同一网络内部的关系一致性约束。R_T_hat是把老师特征图经过另一个投影头再算一次关系矩阵这样做的目的是约束老师自己的特征在投影前后关系保持一致防止投影头把学生特征强行扭曲到完全贴合老师而破坏了学生自身特征的语义。很多人第一次跑 IRG 时只用了第一项发现学生准确率还不如直接训就是少了后两项的约束。权重比例上我的经验是lambda_1在 0.1 到 0.5 之间lambda_2和lambda_3加起来不要超过 1.0具体取值要看数据集的难度CIFAR 上可以取小一点ImageNet 这样的数据量大任务师生对齐项要更保守否则学生过拟合到老师的噪声上。下面给出一个组合损失的完整写法class IRGLoss(nn.Module): def __init__(self, lambda_10.25, lambda_20.5, lambda_30.25): super().__init__() self.l1 lambda_1 self.l2 lambda_2 self.l3 lambda_3 def forward(self, f_t_feats, f_s_feats, proj_s_feats): # 假设传入的是某个 stage 的特征对 R_t compute_relation_map(f_t_feats) R_s compute_relation_map(f_s_feats) R_t_hat compute_relation_map(proj_s_feats) # 学生投影后再算关系图 R_s_hat compute_relation_map(f_s_feats) loss (self.l1 * F.mse_loss(R_t, R_s) self.l2 * F.mse_loss(R_t, R_t_hat) self.l3 * F.mse_loss(R_s, R_s_hat)) return loss注意proj_s_feats的来源它不是老师的特征而是学生特征通过一个投影头之后的结果。投影头的存在解决了师生通道维度不一致的问题同时又不会像直接线性层那样破坏学生的原生表达。2.3 IRG 和 logits 蒸馏的关键差别很多第一次接触 IRG 的人会犯一个定位错误把 IRG 当成 logits 蒸馏的替代品用了 IRG 就丢掉软化概率损失。实际两者并不冲突而且一起用效果更好。logits 蒸馏管的是一张图最终属于哪个类别IRG 管的是特征空间里各个维度之间的组织和相互作用。打个比方logits 是告诉学生这道题的正确答案IRG 是告诉学生解题时各个步骤之间的逻辑关系。一个学生如果只看答案可能在数据分布比较集中的测试集上表现还行但面对长尾分布或者输入图像加了点仿射变换立刻崩。IRG 约束过的学生特征通道之间的激活模式更像老师天然更鲁棒。但如果两者都用损失的融合比重要控制。一种常见做法是把 logits 蒸馏的 KL 散度损失和 IRG 损失相加logits 部分的权重在 0.5 到 1.0 之间IRG 部分按前一节的比例配置。我见过不少工程同学把 IRG 损失权重拉到 5 甚至 10结果学生训练不稳定验证准确率波动很大最后又把 IRG 删了。问题的根源不在 IRG 本身而在权重失衡关系图损失的量纲和交叉熵不同需要在训练初期打印出两者的数值再按数量级调整。3. 数据与模型搭桥ResNet50 和 ResNet18 的特征出口对齐3.1 数据增强策略蒸馏场景下的增强要保守克制在普通分类训练里我们习惯用 RandomResizedCrop、ColorJitter、MixUp 一堆增强。但在蒸馏训练里数据增强太猛会让老师学生的特征对齐变得困难。原因是 IRG 在特征层面做像素级关系匹配图像经过剧烈裁剪或颜色扰动之后老师的关系矩阵本身也会发生很大变化学生学到的关系模式就会不稳定。所以我的做法是蒸馏阶段只保留基础增强RandomCrop 加 RandomHorizontalFlip归一化参数沿用 ImageNet 的 mean 和 std。代码上可以直接复用 torchvision 的 transforms。from torchvision import transforms train_transform transforms.Compose([ transforms.Resize(256), transforms.RandomCrop(224), transforms.RandomHorizontalFlip(), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]) ]) val_transform transforms.Compose([ transforms.Resize(256), transforms.CenterCrop(224), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]) ])这里RandomResizedCrop没有使用因为它会改变物体尺度导致师生特征关系矩阵的空间分布不一致。如果你用 CIFAR-10 这类输入只有 32x32 的小图可以把 Resize 去掉直接 RandomCroppadding 设为 4。3.2 模型改造从 ResNet50 和 ResNet18 中间抽取特征ResNet 的结构是 block 序列每个 stage 内部由多个残差块堆叠。我们要做的是把每个 stage 的输出接出来。用 torchvision 的 resnet50 和 resnet18可以直接修改 forward返回中间特征。import torchvision.models as models import torch.nn as nn class ResNetWithMiddleOutputs(nn.Module): def __init__(self, archresnet18, pretrainedFalse): super().__init__() if arch resnet50: self.model models.resnet50(pretrainedpretrained) else: self.model models.resnet18(pretrainedpretrained) # 取出各stage的关键层引用 self.layer1 self.model.layer1 self.layer2 self.model.layer2 self.layer3 self.model.layer3 self.layer4 self.model.layer4 self.fc self.model.fc def forward_feats(self, x): x self.model.conv1(x) x self.model.bn1(x) x self.model.relu(x) x self.model.maxpool(x) f1 self.layer1(x) f2 self.layer2(f1) f3 self.layer3(f2) f4 self.layer4(f3) return [f1, f2, f3, f4]这一层改造的关键是拿到f1到f4也就是四个 stage 的末层输出。ResNet50 的layer1输出是[B, 256, 56, 56]ResNet18 的layer1输出是[B, 64, 56, 56]空间尺寸相同但通道数差四倍正好用来验证 IRG 算法在通道数不一致时的对齐能力。教师网络用pretrainedTrue加载 ImageNet 预训练权重学生网络用pretrainedFalse从零开始训练。3.3 投影头的设置与通道数匹配因为师生关系矩阵维度不同每个 stage 都需要一个投影头把学生的特征转到老师的通道数。投影头用 1x1 卷积最直接或者用一个两层 conv 加 BN。我建议头不要设计得太重否则学生主要精力放在模仿老师的输出上自身的学习能力会被削弱。class ProjectionHead(nn.Module): def __init__(self, in_channels, out_channels): super().__init__() self.conv nn.Conv2d(in_channels, out_channels, kernel_size1, biasFalse) self.bn nn.BatchNorm2d(out_channels) def forward(self, x): return self.bn(self.conv(x))四个 stage 对应的投影头模块可以从一个 list 里初始化teacher ResNetWithMiddleOutputs(resnet50, pretrainedTrue) student ResNetWithMiddleOutputs(resnet18, pretrainedFalse) C_teacher [256, 512, 1024, 2048] C_student [64, 128, 256, 512] projectors nn.ModuleList([ ProjectionHead(cs, ct) for cs, ct in zip(C_student, C_teacher) ])注意这里投影头输入是学生通道数输出是老师通道数。训练时学生原始特征和投影后的特征都要传入损失函数。训练前半段投影头的参数变化较快如果发现 IRG 损失在 epoch 1 到 epoch 5 之间剧烈下降是正常现象如果一直不降先检查投影头前面有没有接 ReLU有些实现里残留了 ReLU导致关系矩阵被截断。4. 蒸馏训练实战损失函数、超参与训练日志4.1 完整的训练流程与损失组合先把训练流程的主干写出来包含 IRG 损失、logits 蒸馏损失和交叉熵。用 ImageNet 数据太大这里用 CIFAR-100 做演示你可以直接换成自己的数据集。模型结构是上一节改造过的。import torch.optim as optim teacher.eval() for p in teacher.parameters(): p.requires_grad False student.train() optimizer optim.SGD(student.parameters(), lr0.01, momentum0.9, weight_decay5e-4) projectors_optimizer optim.SGD(projectors.parameters(), lr0.01, momentum0.9, weight_decay5e-4) irg_criterion IRGLoss(lambda_10.25, lambda_20.5, lambda_30.25) ce_criterion nn.CrossEntropyLoss() kl_criterion nn.KLDivLoss(reductionbatchmean) def kl_loss_with_temperature(s_logits, t_logits, T4.0): s_logits F.log_softmax(s_logits / T, dim1) t_probs F.softmax(t_logits / T, dim1) return kl_criterion(s_logits, t_probs) * (T * T)训练一个 batch 的完整流程如下。for images, labels in train_loader: images, labels images.cuda(), labels.cuda() with torch.no_grad(): t_feats teacher.forward_feats(images) t_logits teacher.model.fc(teacher.model.avgpool(t_feats[-1]).flatten(1)) s_feats student.forward_feats(images) s_logits student.model.fc(student.model.avgpool(s_feats[-1]).flatten(1)) # 投影头输出 proj_feats [proj(sf) for proj, sf in zip(projectors, s_feats)] # IRG 损失每个 stage 的损失累加 irg_loss 0.0 for tf, sf, pf in zip(t_feats, s_feats, proj_feats): irg_loss irg_criterion(tf, sf, pf) irg_loss irg_loss / len(t_feats) ce_loss ce_criterion(s_logits, labels) distill_loss kl_loss_with_temperature(s_logits, t_logits) total_loss ce_loss 1.0 * distill_loss 0.5 * irg_loss optimizer.zero_grad() projectors_optimizer.zero_grad() total_loss.backward() optimizer.step() projectors_optimizer.step()这段代码最需要注意的是student.model.fc被调用了两次一次在forward_feats里没有经过另一次手动拿s_feats[-1]去接 avgpool 和 fc。这里的t_feats[-1]和s_feats[-1]都是 stage4 的输出对于 224x224 输入尺寸是[B, C, 7, 7]先 avgpool 到[B, C, 1, 1]再 flatten 就能接全连接层。如果你想省事可以把forward_feats里也返回最终的 logits。4.2 超参选择与调整方法蒸馏训练的初始化策略和普通训练有区别。学生网络初始化为随机权重教师网络权重保持预训练状态不变。学习率从 0.01 开始每 30 个 epoch 乘 0.1。温度参数 T 我建议设 4.0这是比较常见的中间值。温度太低软标签接近 one-hotlogits 蒸馏的信息量下降温度太高类别间的细粒度差异被抹平学生学到的区分能力变弱。IRG 损失的权重是目前影响最大的超参数它不是一个无量纲的值而是与数据集、教师网络的表征质量相关。可以参考下面的调参顺序权重项初始值调整方向lambda_1 师生关系对齐0.25学生模型表现弱时小步加到 0.5lambda_2 教师自关系一致0.5特征可视化逐渐收敛时不需动lambda_3 学生自关系一致0.25加太大会压缩学生的表达能力不建议超过 0.5logits 蒸馏温度 T4.0数据集越小温度越低2.0 也可训练过程中观察 loss 的下降曲线如果irg_loss在 epoch 5 后还在剧烈波动说明投影头的学习率太高把投影头的 lr 单独降到学生的十分之一。投影头和学生的优化器分开就是这个目的。如果 total_loss 一直降但验证准确率不动检查是不是 IRG 权重太大学生的表征被过度约束到老师的空间中丧失了自适应能力。4.3 训练日志里必须监控的四个指标不要只看训练 loss。蒸馏训练里最值得监控的是师生关系矩阵之间的距离、学生自身的交叉熵、蒸馏后的验证准确率、以及学生关系矩阵的稳定性。第四个指标往往是初坑所在如果你发现训练几千步后学生不同 batch 之间的关系矩阵波动非常大但 IRG 损失已经在下降说明关系矩阵的绝对值大小可能是对的但分布形状不对原因通常是compute_relation_map里没有对矩阵做任何归一化通道激活绝对值大的点积结果也大数值范围失真。在计算关系矩阵时建议对特征向量做一次 L2 归一化也就是把pooled除以它的 L2 范数这样关系矩阵的元素范围落在 -1 到 1 之间训练更稳定。代价是蒸馏的信息会损失一点方向性强度但对大多数场景来说这个交换是值得的。修改方式很简单pooled F.normalize(pooled, p2, dim1)训练日志建议每 200 步打印一次格式写成step | CE | KD | IRG | acc这里acc是学生模型在当前 batch 上的 top1 准确率。代码实现时这个 acc 在训练的前几个 epoch 会很低可能只有 10% 左右但千万别因为 acc 低就以为蒸馏失败学生还在拟合阶段。5. 用关系矩阵热力图验证学生是否真正学到了老师的内部结构模型训练结束后验证准确率只能说明学生知道了正确答案不能证明学生学会了老师的组织结构。IRG 算法的成功标准应该是学生网络的关系矩阵和老师网络的关系矩阵在结构上高度相似。这时候可以写一个小脚本把关系矩阵可视化。取验证集的一批图片传入师生网络固定取layer3的输出计算关系矩阵然后画热力图。对 ResNet50 蒸馏 ResNet18 来说layer3师生关系矩阵尺寸分别是 1024x1024 和 512x512太大不适合直接可视化所以要降采样要么只取前 64 个通道要么把关系矩阵做 top-k 稀疏化只保留每行最大的 8 个值其余置零。后者更能体现关系结构。import matplotlib.pyplot as plt teacher.eval() student.eval() with torch.no_grad(): batch, _ next(iter(val_loader)) batch batch.cuda() t_feats teacher.forward_feats(batch) s_feats student.forward_feats(batch) R_t compute_relation_map(t_feats[2])[0].cpu().numpy() R_s compute_relation_map(s_feats[2])[0].cpu().numpy() # 取前 64 个通道避免画布过大 R_t R_t[:64, :64] R_s R_s[:64, :64] fig, axes plt.subplots(1, 2, figsize(10, 4)) axes[0].imshow(R_t, cmapviridis) axes[0].set_title(Teacher ResNet50 layer3) axes[1].imshow(R_s, cmapviridis) axes[1].set_title(Student ResNet18 layer3) plt.savefig(irg_relation_map.png, dpi150)从热力图里可以直观看到如果学生的关系矩阵出现了和老师相似的块状结构说明 IRG 约束生效了。如果学生矩阵呈现成一片均匀的亮色说明学生没有学到通道间的差异化关系很可能关系矩阵没有做 L2 归一化或者投影头过强把学生的特征全部拉到了同一个尺度上。接着做一个量化验证计算师生关系矩阵的余弦相似度。不需要逐元素比较计算矩阵展平后的余弦相似度数值越高说明结构越接近。正常训练下这个值应该接近 0.8 以上。from numpy.linalg import norm cos_sim (R_t.flatten() R_s.flatten()) / (norm(R_t.flatten()) * norm(R_s.flatten())) print(frelation map cosine similarity: {cos_sim:.4f})如果这个相似度低于 0.6即使准确率到了预期值也要回头检查 IRG 损失是否真的在反向传播中生效大概率是requires_grad没有设为 True或者投影头的输出没有接入损失计算。最后还有一个通用技巧把训练好的学生模型固定下来只重新训练投影头观察 IRG 损失能不能降到很低如果投影头在冻结学生参数时还能把关系矩阵对上说明学生的内部结构已经固化可以放心部署。本文还有配套的精品资源点击获取
返回列表