ARTICLE DETAIL

资讯详情

深耕网站视觉设计与运营推广的一线实战洞察。

RKD关系知识蒸馏实战:用CoatNet蒸馏ResNet的Loss实现与调参

RKD关系知识蒸馏实战:用CoatNet蒸馏ResNet的Loss实现与调参 简介RKD关系知识蒸馏实战代码包面向有一定深度学习基础、希望系统掌握用CoatNet蒸馏ResNet的算法工程师与研究人员。方案围绕展平层特征展开通过样本间距离与角度关系构建Distance-wise Loss和Angle-wise Loss分别从二阶距离和三阶角度传递教师模型的结构化知识比单纯模仿logits更利于学生模型学习特征间相对关系。压缩包约2000个文件绝大多数为训练与推理过程的可视化png图表便于复盘损失曲线和特征分布变化另有7个Python脚本封装模型构建、蒸馏损失与训练流程并包含1个pyc辅助文件整包约930MB。目前已有621人学习下载适合与作者前一篇蒸馏方法对照阅读快速复现RKD并将相应模块迁移至自己的分类或检测任务中。1. RKD知识蒸馏实战用CoatNet蒸馏ResNet前先搞懂什么叫关系蒸馏做知识蒸馏的同学通常有个惯性把教师模型的输出 logits或者某一层特征图拉出来让学生照着抄。但 RKDRelational Knowledge Distillation关系知识蒸馏的做法反直觉——它不对齐单个样本的输出而是把模型展平层的特征提取出来蒸馏样本与样本之间的关系。蒸馏的 Loss 被拆成两部分二阶的 Distance-wise Loss 负责让样本对距离分布一致三阶的 Angle-wise Loss 负责让三元组夹角分布一致。这套方法用在 CoatNet 蒸馏 ResNet 上比逐样本特征蒸馏更容易收敛。下面从 RKD 与传统蒸馏的差异讲起把 Loss 实现、训练配置、参数调优和踩坑记录完整过一遍代码可以直接改到自己的分类任务上。适合正在做模型压缩、或者想对比不同知识蒸馏代码实现的人。2. RKD与传统蒸馏的差异为什么特征蒸馏不适合CoatNet到ResNet的迁移2.1 传统蒸馏在做什么logits逐样本对齐的局限Hinton 的知识蒸馏核心是让学生软化后的 logits 拟合教师的软化 logits。softmax 之前除以温度 T把概率分布压平类别之间的相对关系被视为教师学到的暗知识。温度 T 的取值直接影响蒸馏效果我一般从 T3 开始试。T 太小软化不明显学生只能学到教师自信的类别T 太大输出过度平滑类别间的精细排序被稀释学生拿到的监督信号就变成了模糊的排序关系。后续的 FitNet 把蒸馏点从输出层挪到中间层直接对齐教师和学生的 feature map。这类方法的共同点是逐样本对齐每个输入样本教师和学生各出一个表征然后强制这两个表征一致。问题在于它把每个样本当作独立的点处理教师内部真正有价值的结构——哪些样本在特征空间里靠得近、哪些离得远、不同类别在空间里怎么分布——被逐样本对齐这个操作丢掉了。还有一个更实际的工程问题logits 蒸馏是在类别维度上做对齐特征已经被压成类别数那么长的向量feature map 蒸馏则要求教师和学生的空间分辨率、通道数对齐异构网络下往往需要额外插值和通道变换。这些约束让逐样本方法在 CoatNet 蒸馏 ResNet 的场景里非常别扭后面会解释为什么别扭。2.2 RKD的核心二阶距离损失与三阶角度损失RKD 出自 CVPR 2019 的《Relational Knowledge Distillation》思路从蒸馏输出转向蒸馏关系。关系分两个层次。二阶关系是样本对之间的欧氏距离教师特征空间里batch 内所有样本对的距离构成一个分布RKD 要求学生特征空间里同一批样本对的距离分布去拟合这个分布。三阶关系是三元组的角度三个样本构成两个向量以其中一个样本为锚点向量之间会形成夹角RKD 让学生的夹角分布去拟合教师的夹角分布。两个 Loss 合起来等于在传递特征空间的几何结构。二阶信息约束样本间的远近三阶信息约束样本间的方位。实现上有一个关键约定RKD 蒸馏的不是原始 feature map而是展平层flattened layer的特征一般取全局池化后的 [B, D] 矩阵。B 是 batch sizeD 是特征维度所有距离和角度都基于这个矩阵计算。这一点决定了 RKD 对 batch size 有硬性要求batch 太小能构成的样本对和三元组数量太少loss 的统计意义被削弱梯度噪声会明显变大。我见过不少第一次跑 RKD 的人把 batch 设成 16然后来看我为什么 loss 震荡其实不是代码问题是统计量不足。2.3 为什么关系蒸馏更适合异构网络迁移模型压缩最常见的组合是同一架构的大小版本比如 ResNet-50 蒸馏 ResNet-18。同构组合的特征空间几何结构相似逐样本特征蒸馏还能凑合。但 CoatNet 蒸馏 ResNet 是典型的异构组合。CoatNet 是 CNN 和自注意力融合的网络既有卷积的局部建模能力又有自注意力的全局关系建模能力ResNet 是纯卷积网络感受野靠堆叠层数逐步扩大没有全局交互机制。结构差异直接反映在特征空间上。CoatNet 的高层特征偏粗粒度、全局化一个 token 能看到整张图的信息ResNet 的特征偏细粒度、局部化每个位置主要感知一个局部区域。两者的特征空间不存在清晰的逐点对应关系逐样本蒸馏依赖的那个映射本身就不成立。RKD 绕开了这个问题它不要求逐点对应只要求样本间的相对距离和相对角度保持一致。学生不需要长成教师的形状只需要在自己的特征空间里维持同样的几何关系。对比项logits蒸馏feature蒸馏RKD对齐对象软化概率分布中间层特征图样本间距离与角度特征形态类别数向量通道×高×宽全局池化后 [B, D]结构差异容忍度高低中高收敛稳定性稳定依赖层选择依赖 batch size计算开销极小中等距离矩阵较大什么时候别用 RKD如果教师和学生同构且结构接近feature 蒸馏配一点 logits 预热往往更省事RKD 的距离矩阵计算在 batch 较大时有额外开销。但要做异构迁移我的建议是直接用 RKD 起步省去逐层对齐的调参过程。3. 蒸馏流程与Loss实现Distance-wise和Angle-wise怎么落代码3.1 构建教师与学生模型CoatNet与ResNet的加载方式教师模型用 timm 加载预训练的 coatnet学生用 torchvision 的 resnet。这里选 coatnet_0 而不是更大的版本是因为蒸馏训练时教师只做前向大模型在前向上的时间开销同样会被放大小版本足够提供有效的关系监督。import torch import torch.nn as nn import torchvision.models as models from timm import create_model # 教师CoatNet加载 ImageNet 预训练权重之后全程冻结 teacher create_model(coatnet_0_rw_224, pretrainedTrue) teacher.eval() # 学生ResNet18从零训练按数据集修改 num_classes student models.resnet18(pretrainedFalse, num_classes100) # 冻结教师参数教师不参与梯度更新只提供特征 for p in teacher.parameters(): p.requires_grad False教师用create_model加载coatnet_0_rw_224是较小版本单卡能跑换coatnet_1或coatnet_2需要按显存调 batch size后面避坑章节会讲 batch 对 RKD 的影响。学生用resnet18num_classes按数据集改这里是 CIFAR-100 的示意。加载完之后的关键问题是RKD 要的是展平层特征也就是全局池化后的 [B, D] 向量而模型默认只返回 logits。timm 模型一般提供forward_features方法来拿 backbone 输出ResNet 则用avgpool的输出。做蒸馏时我习惯统一封装一层特征提取逻辑def extract_flatten_feature(model, x): if hasattr(model, forward_features): feat model.forward_features(x) # timm 模型通用入口 else: feat model.avgpool(model.layer4(x)) # torchvision ResNet if feat.dim() 4: feat feat.flatten(2).mean(-1) # 转成 [B, D] return feat这个封装把教师和学生的特征统一成 [B, D]。注意forward_features在不同 timm 版本里返回的可能是 GAP 之前的特征也可能是之后的所以保留一个dim() 4的判断确保进入 RKD 之前形状一定是 [B, D]。3.2 Distance-wise Loss样本对距离分布对齐Distance-wise Loss 的目标是让教师的样本对距离分布和学生的一致。先算 batch 内两两欧氏距离再归一化最后做平滑 L1 对齐。import torch.nn.functional as F def distance_wise_loss(t_feat, s_feat, alpha1.0, kernel_t0.5): # 输入形状 [B, D]返回一个标量 loss t_dist torch.cdist(t_feat, t_feat, p2) # [B, B] 距离矩阵 s_dist torch.cdist(s_feat, s_feat, p2) B t_feat.size(0) mask ~torch.eye(B, dtypetorch.bool, devicet_feat.device) t_dist t_dist[mask].view(B, B - 1) # 去掉对角线自身距离 s_dist s_dist[mask].view(B, B - 1) # softmax 核归一化距离越近归一化权重越大 t_dist F.softmax(-t_dist / kernel_t, dim1) s_dist F.softmax(-s_dist / kernel_t, dim1) loss F.smooth_l1_loss(s_dist, t_dist) * alpha return losstorch.cdist直接算两个矩阵逐行之间的欧氏距离对角线是样本与自身的距离必为 0用 mask 去掉。归一化这步很关键教师和学生特征向量的模长往往不在一个量级如果不归一化loss 会被尺度主导学生只要缩小自己的特征范数就能骗过 loss根本学不到关系结构。这里用的是 softmax 核归一化也是 RKD 论文里常见做法的工程简化版距离越近的样本对权重越大kernel_t控制分布的锐度值越小越集中。实际调参时我一般把kernel_t固定成 0.5优先调外面的alpha。smooth_l1_loss对离群距离对的梯度比 MSE 温和距离分布里偶尔出现一个极大离群点L1 部分不会放大那个梯度训练更稳。3.3 Angle-wise Loss三元组角度对齐Angle-wise Loss 负责三阶角度关系。简化版本用相邻样本构成三元组第 i 个样本做锚点i1 和 i2 作为另外两个点这样 batch 内 B 个样本能构造 B-2 个三元组计算量小实现也直观。def angle_wise_loss(t_feat, s_feat, alpha1.0): def get_angle(feat): anchor feat[:-2] # 锚点样本 pos feat[1:-1] # 第二个样本 neg feat[2:] # 第三个样本 v1 anchor - pos # 锚点到 pos 的向量 v2 anchor - neg # 锚点到 neg 的向量 cos F.cosine_similarity(v1, v2, dim1) # [B-2] return cos t_angle get_angle(t_feat) s_angle get_angle(s_feat) loss F.mse_loss(s_angle, t_angle) * alpha return losscosine_similarity的输出范围是 [-1, 1]用 MSE 让学生的夹角余弦贴近教师的夹角余弦等价于对齐角度值还避免了角度的周期性边界问题。这里的简化写法用的是相邻三元组batch 内的样本顺序会略微影响结果。如果想让角度统计更充分常见做法是用torch.combinations枚举全部三元组但计算量是 O(B³)batch 一大就扛不住。我一般改成随机采样每步随机抽 256 个三元组索引用gather把对应特征捞出来再算余弦统计意义比相邻三元组好开销也可控。这个替换对最终效果的提升在几个小数据集上能稳定观察到。3.4 蒸馏训练主循环Loss加权与参数更新训练主循环的骨架和普通分类训练差不多多出来的就是把两个 RKD loss 和 CE 按权重加在一起。教师特征必须detach()否则反向传播会走回教师模型的计算图显存直接翻倍。optimizer torch.optim.AdamW(student.parameters(), lr1e-3, weight_decay5e-4) scheduler torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max80) criterion nn.CrossEntropyLoss() beta 0.5 for epoch in range(80): for x, target in train_loader: x, target x.cuda(), target.cuda() with torch.no_grad(): teacher_logits teacher(x) teacher_feat extract_flatten_feature(teacher, x).detach() student_logits student(x) student_feat extract_flatten_feature(student, x) ce_loss criterion(student_logits, target) d_loss distance_wise_loss(teacher_feat, student_feat, alpha1.0) a_loss angle_wise_loss(teacher_feat, student_feat, alpha1.0) loss ce_loss beta * (d_loss a_loss) optimizer.zero_grad() loss.backward() optimizer.step() scheduler.step()这里教师在torch.no_grad()块里前向特征再detach()一道双保险。ce_loss不能去掉RKD 只约束特征空间的关系不约束学生输出和真实标签的对应关系没有 CE loss学生可能学出一个几何结构正确但类别错乱的特征空间。beta控制 RKD 总体强度一般 0.1 到 1.0 之间。调节顺序建议是先固定beta调两个alpha两个alpha稳定后再动beta。调度器用余弦退火蒸馏训练建议比普通训练多跑一些 epoch关系蒸馏的收敛通常比 logits 蒸馏慢半拍epoch 太少容易刚看到下降趋势就停了。4. 训练配置与参数调优从损失权重到特征层选择4.1 损失权重怎么配alpha和beta的调节顺序RKD 的最终 loss 由三部分组成CE、Distance-wise Loss、Angle-wise Loss。我用beta管 RKD 总体强度d_alpha和a_alpha分管距离和角度两路。这个分层结构的好处是可以先定位是关系监督整体太弱还是某一路关系没对齐。权重组合特点适用场景beta0.5, d_alpha1.0, a_alpha1.0均衡两路 loss 量级接近默认起点beta1.0, d_alpha2.0, a_alpha0.5距离主导教师学生特征空间规模差异大beta0.3, d_alpha0.5, a_alpha1.0角度主导样本对距离在异构网络下噪声较大调参顺序比我见过很多人直接乱试要重要先跑一版默认权重看验证集精度。距离 loss 下降慢就提高d_alpha角度 loss 震荡先降a_alpha再看 batch size。如果两个 loss 都正常但精度不动问题不在权重在特征层选择去查 4.2 的内容。还要注意一个常见误判RKD 两个 loss 的绝对数值通常不一样不能直接比较谁大谁小。我用的是各自相对训练初值的下降幅度来判断哪一路没学好。下降幅度小的那一路才是需要加权的对象。4.2 展平层的选择全局池化特征还是中间层特征RKD 对展平层特征做蒸馏但展平层具体取哪一层很多人没细想。我的经验是取全局池化之后的 [B, D] 向量理由有三点。第一GAP 后的特征是模型的全局表征聚合了整个空间的信息和 RKD 关注的样本间关系结构匹配第二中间层特征 [B, C, H, W] 直接拉平维度太高cdist矩阵计算开销大而且大量空间位置是噪声第三中间层空间分辨率在教师和学生之间往往不一致强行拉平会导致距离和角度的语义错位。粗粒度特征和细粒度特征在这个选择上的差异很明显。细粒度特征保留空间位置信息适合检测、分割这类像素级任务但关系蒸馏关心的是样本 A 和样本 B 在语义上多接近这个语义距离在细粒度特征上被空间噪声干扰得很厉害。反而是经过 GAP 聚合后的粗粒度特征每个维度都有明确的全局语义距离和角度的稳定性更好。如果学生模型太小、GAP 特征维度不够表达可以取最后一个 stage 的输出做自适应池化再拉平但要注意让教师和学生取的位置保持在同一语义层级。常见错误是教师取倒数第二层、学生取倒数第一层两边表征的抽象程度不一致关系对齐学到的东西对学生最终分类没有帮助。4.3 优化器、Batch Size与学习率RKD 对 batch size 的敏感度比常规蒸馏高很多。原因在于距离和角度都是在 batch 内构造样本对和三元组batch 越大关系统计越稳定。多卡训练时 batch256 是最理想的单卡建议至少 64低于这个值就要做好 angle loss 震荡的心理准备。学生从零训练时优化器我习惯用 AdamW学习率 1e-3配合余弦退火学生加载了预训练权重时学习率降到 1e-4 做微调。weight_decay按分类任务常规值 5e-4 设置。RKD 不需要像 feature distillation 那样做特征归一化对齐所以不用额外的 BN 同步技巧。有个判断经验可以分享训练前 10 个 epochdistance loss 应该稳步下降angle loss 波动是正常的如果 distance loss 一开始就剧烈震荡优先降学习率再看 batch size 是否太小。学习率过大时关系 loss 的梯度会把学生特征空间推来推去样本间距离结构根本稳定不下来。4.4 CIFAR-100上的一组参考配置给一组在 4 卡 GPU 上跑过的参考配置用于复现基线和对比。CIFAR-100 数据集较小但 224 分辨率输入能保证教师模型的特征质量不会因为输入裁剪损失太多信息。配置项取值数据集CIFAR-100分辨率 224教师coatnet_0_rw_224 预训练学生resnet18 随机初始化优化器AdamWlr1e-3weight_decay5e-4调度器CosineAnnealingLRT_max80batch size1284 卡每卡 32蒸馏权重beta0.5d_alpha1.0a_alpha1.0训练轮数80精度对比的趋势如下数字是示意性的依赖数据增强和随机种子但相对关系稳定。RKD 在异构蒸馏上确实能带来几个点的收益而且收敛更早。方法学生Top-1学生单独训练无蒸馏约 73%logits 蒸馏Hinton KD约 74%RKD本文配置约 75%如果学生换成 resnet34收益空间还会更大因为学生容量上来了教师输出的关系结构有更多维度去承载。5. RKD蒸馏避坑指南四个必踩的坑和排查顺序先交代一下下面这四条全是我和同事实际踩过的。RKD 的原理不难代码也不长但训练过程中的坑比想象中隐蔽而且每个坑的共性都是loss 看起来正常结果不对。5.1 坑教师特征没有detach显存直接翻倍现象训练刚开始显存占用是预期的两倍跑几个 iteration 后报 CUDA out of memory。原因教师参数虽然设了requires_gradFalse但 RKD 的 loss 回传路径上如果教师特征没有detach反向传播仍然会走进教师模型的计算图教师那一大堆中间激活全部被保留下来显存自然翻倍。这个问题在 ResNet 系列里不明显因为结构浅CoatNet 这种融合了自注意力的模型中间激活的显存占用非常高。解决传给两个 loss 的教师特征统一teacher_feat.detach()再保险一点把教师前向放进torch.no_grad()块里。我后来的习惯是每次定义 loss 输入时检查一眼张量的grad_fn如果带反向传播入口立刻意识到教师特征没有切断。5.2 坑展平层选错loss正常但验证集不掉现象distance loss 在下降angle loss 也在下降学生验证精度就是不动整个训练过程像个黑匣子。原因展平层选在了教师和学生对不齐的位置。比如教师用全局池化后的特征学生取的是第二个 stage 的中间特征两者代表的语义层级差了好几个抽象级别。关系对齐确实在发生但对齐的是一对语义错位的表征学生学到的东西对最终分类没帮助。解决统一取各自最后一个 stage 之后的 GAP 特征。验证方法也简单把学生的 GAP 特征直接接一个 linear 分类头单独训 10 个 epoch看这个特征本身能不能达到合理精度。特征本身表达力不够RKD 再强也救不回来。5.3 坑batch size太小angle loss剧烈震荡现象batch size 设 16 或 32angle loss 在训练中上下乱跳验证精度也忽高忽低调低学习率只能缓解不能消除。原因角度 loss 基于三元组计算batch16 只能构造 14 个相邻三元组一个异常样本就能让夹角分布产生明显抖动。这个 loss 本质上是统计量统计量太小自然不稳定。距离 loss 用的是样本对batch16 时样本对数量是 240勉强够角度只有 14 个样本远低于稳定统计的下限。解决batch 至少 64最好 128 以上。单卡显存不够就降输入分辨率或者用梯度累积凑到大 batch 再更新参数。用梯度累积时注意 loss 要除以累积步数否则相当于学习率被放大了学习率曲线对不上。5.4 坑教师太强学生被闷死现象用 ImageNet 预训练的大号 CoatNet 蒸馏 CIFAR-100 上的 resnet18loss 下降正常最终精度反而不如不带蒸馏的学生。原因教师能力过强时特征空间里样本间的关系结构太复杂学生容量不足拟合不了这种复杂关系蒸馏 loss 就退化成噪声源。这是异构蒸馏的常见翻车点不是方法无效是教师和学生的能力差距超出了合理区间。解决换更小的教师模型比如从 coatnet_2 降到 coatnet_0或者给学生加一点预训练权重先把学生拉到一个合理的起点或者把beta降到 0.1让蒸馏从主导退成辅助。还想保留大教师的话可以先蒸馏出一个中等模型再把中等模型的输出蒸馏给最终学生这条链路在工程上更稳。5.5 坑CE被蒸馏loss淹没分类学不会现象RKD loss 下降得很漂亮但学生验证精度稳定在随机水平附近像是完全没学过分类。原因beta和两个alpha配得太大RKD 的 loss 数值把 CE loss 完全压住梯度主要由关系监督贡献。关系监督只约束样本间的相对结构不约束样本和真实标签的对应关系学生确实学到了一个结构优美的特征空间但里面没有类别的概念。解决先看训练日志里三个 loss 的数值量级。出现这种情况立刻把beta降到 0.1 或者把beta固定为 0.5 但调低两个alpha确保 CE loss 的量级和蒸馏 loss 在同一个数量级。我习惯在前 5 个 epoch 只跑 CE loss 做预热等学生有基本的分类能力后再把 RKD loss 加进去效果比全程加最权重更稳。排查顺序再补充一点如果上面五条都排除了问题还在按这个顺序查。先打印三个特征张量的 shape确认是 [B, D] 而不是 [B, C, H, W]再确认教师真的在eval()模式下Dropout 和 BatchNorm 在 train 模式下会给教师特征引入随机噪声最后检查两个 loss 的输入参数有没有传反。t_feat和s_feat传反是那种最不想承认的错误训练看起来一切正常实际蒸馏的方向完全反了。6. 验证蒸馏效果关系矩阵可视化与三个辅助指标6.1 把特征关系画出来训练完之后最直接的验证方法是可视化特征关系矩阵。取验证集的一部分样本比如 100 个分别过教师和学生模型算出距离矩阵用热力图并排对比。如果两张图的块状结构基本一致说明学生学到了教师的样本间关系如果学生那边是一团模糊或噪声说明 RKD 没生效回去查特征层。import matplotlib.pyplot as plt import seaborn as sns # teacher_feat / student_feat 都是 [100, D] t_dist torch.cdist(teacher_feat, teacher_feat) s_dist torch.cdist(student_feat, student_feat) fig, axes plt.subplots(1, 2, figsize(10, 4)) sns.heatmap(t_dist.cpu().numpy(), axaxes[0], cbarFalse) sns.heatmap(s_dist.cpu().numpy(), axaxes[1], cbarFalse) axes[0].set_title(teacher distance) axes[1].set_title(student distance) plt.show()6.2 三个辅助指标精度之外我用三个辅助指标判断蒸馏质量。第一个是关系矩阵的 Spearman 相关系数把两个距离矩阵展平后算秩相关越高说明距离关系保留得越好。第二个是三元组一致性率随机采样一批三元组统计教师和学生对同一批三元组角度排序一致的比例。第三个是学生 GAP 特征在验证集上的 KNN 精度用几个邻居分类能反映特征空间是否真的压缩出了类间结构。这套验证流程是我第一次做 RKD 时踩出来的。那时候只盯着验证精度调了半天不知道关系有没有对齐把距离矩阵画出来才发现学生特征空间的结构是乱的类别之间的边界完全糊在一起。从那以后我每次蒸馏训练结束都强制走一遍精度 关系矩阵 三元组一致性率三件套关系矩阵看着不对劲就先不动权重的参数等特征结构对了再继续调。希望帮到你。本文还有配套的精品资源点击获取
返回列表