RKD知识蒸馏实战:CoatNet教ResNet学关系 简介RKD知识蒸馏实战资源包面向具备基础深度学习训练经验、想在分类任务中通过特征蒸馏压缩模型的开发者。该方案使用CoatNet作为教师网络对ResNet进行蒸馏与常规Logits蒸馏不同RKD针对展平层特征开展蒸馏将损失拆解为二阶距离损失Distance-wise Loss与三阶角度损失Angle-wise Loss可广泛用于图像分类模型的轻量化部署。压缩包共2000余个文件以2406张png图像作为训练与验证数据另有7个Python脚本和1个pyc文件完整代码加数据压缩后约930.94MB。目前已有623人学习下载。通过这份资源开发者可以深入理解RKD损失计算、教师与学生网络在展平层的特征对齐方式并能直接基于脚本调整教师网络、学生网络与蒸馏权重适合中高级学习者作为特征蒸馏实战参考。1. RKD知识蒸馏实战当CoatNet把“关系”教给ResNet在知识蒸馏里大部分方法都在教学生“抄答案”——让学生模型的logits去逼近教师的logits或者让中间层特征图做逐像素对齐。但有一个场景这种“抄答案”会失效当学生模型的容量远小于教师时特征空间的结构差异太大硬对齐反而把学生带偏。RKDRelational Knowledge Distillation走的是另一条路不直接对齐单个样本的特征而是对齐样本与样本之间的关系包括两两之间的距离分布Distance-wise Loss和三三之间的角度分布Angle-wise Loss。这份资源用CoatNet当教师、ResNet当学生对展平层特征做RKD蒸馏。适合正在做模型压缩、换轻量骨干网络时精度掉点、以及想理解关系蒸馏为什么比特征对齐更稳的从业者。看完这篇笔记你可以直接照着手头的分类项目改出一套可跑的蒸馏脚本。2. RKD的核心设计为什么关系比特征更值得学2.1 展平层特征的“关系场”视角传统蒸馏方法比如FitNet要求学生中间层的feature map和教师逐像素逼近这隐含了一个假设师生的特征空间是对齐的。但实际中CoatNet这类Transformer风格的模型和ResNet这类纯卷积模型特征分布差异极大硬对齐会让ResNet被迫去拟合CoatNet的注意力分布反而丢失自己的归纳偏置。RKD的处理方式是把特征图展平成一维向量然后只在向量之间的几何关系上做约束。CoatNet输出的展平特征保留的是全局空间信息ResNet输出的展平特征则更偏局部纹理。两者关系结构做蒸馏本质上是在教ResNet“模仿CoatNet看待样本间相似性与差异性的方式”而不是模仿具体特征值。这就绕开了特征空间不对齐的问题。这份资源里教师CoatNet最后输出的展平特征一般维度在1000左右取决于具体变体学生ResNet在池化层后取展平特征再接一个小的Embedding层映射到相同维度。常见做法是Embedding层用两层的MLP中间带ReLU输出维度对齐到CoatNet的特征维。2.2 Distance-wise Loss软化的距离分布距离损失的核心不是让样本对距离数值相等而是让距离的分布相似。计算方式是在一个batch内对每个样本算出它与其他所有样本的特征欧氏距离得到一组距离值然后除以温度参数做softmax变成一个概率分布再用KL散度去约束学生这组分布逼近教师的分布。距离损失的计算代码如下import torch import torch.nn.functional as F def distance_wise_loss(student_feat, teacher_feat, temperature1.0): # student_feat / teacher_feat: [batch, dim] # 计算两两之间的欧氏距离矩阵 s_dist torch.cdist(student_feat, student_feat, p2) # [batch, batch] t_dist torch.cdist(teacher_feat, teacher_feat, p2) # distance-wise的softmax负距离越大相似度越高 s_logits -s_dist / temperature t_logits -t_dist / temperature # 按行做softmax得到概率分布 s_prob F.log_softmax(s_logits, dim-1) t_prob F.softmax(t_logits, dim-1) # 每行计算KL散度后取均值 loss F.kl_div(s_prob, t_prob, reductionbatchmean) return loss这里有一个关键参数temperature。温度越大softmax后的分布越平缓梯度越平滑温度越小越接近one-hot梯度容易爆炸。资源里默认取1.0但这个值需要根据数据集调我一般会先在0.5到2.0之间试几轮。另外要注意torch.cdist在计算时不包括样本自身的距离对角线为0softmax后对角线这一项会参与概率分布但因为所有对角线都被归一化影响不大。如果你的batch内存在重复样本对角线项会干扰分布这种情况可以手动mask掉对角线再归一化。2.3 Angle-wise Loss三元组夹角约束角度损失比距离损失高一阶它关心的是三个样本组成的几何角度。算术上就是取三个样本的特征向量a、b、c计算向量b-a和c-a的夹角余弦。这个角度反映的是特征空间中样本间的“方位关系”比距离更稳定——因为距离会受特征尺度影响但角度天然对缩放不变。角度损失计算如下def angle_wise_loss(student_feat, teacher_feat, temperature1.0): batch student_feat.size(0) # 生成所有可能的三元组索引 idx torch.arange(batch) triples [] for i in idx: for j in idx: if i j: continue for k in idx: if k i or k j: continue triples.append((i, j, k)) # 这里只保留ij对称性的三元组避免重复计算 s_angles, t_angles [], [] for (i, j, k) in triples: v1_s student_feat[j] - student_feat[i] v2_s student_feat[k] - student_feat[i] cos_s F.cosine_similarity(v1_s, v2_s, dim0) s_angles.append(cos_s) v1_t teacher_feat[j] - teacher_feat[i] v2_t teacher_feat[k] - teacher_feat[i] cos_t F.cosine_similarity(v1_t, v2_t, dim0) t_angles.append(cos_t) s_angles torch.stack(s_angles) t_angles torch.stack(t_angles) s_logits s_angles / temperature t_logits t_angles / temperature s_prob F.log_softmax(s_logits, dim-1) t_prob F.softmax(t_logits, dim-1) loss F.kl_div(s_prob, t_prob, reductionbatchmean) return loss这个朴素实现在batch64时会生成约25万个三元组Python循环直接跑会非常慢。资源里的做法是用矩阵运算一次性算出所有角度先对每个样本计算与其他样本的向量差然后做批量点积和范数运算得到一个[batch, batch, batch]的夹角张量。具体优化代码在项目里已经写好了这里不展开。实际使用时还有个重要参数三元组的采样策略。全量三元组在batch较大时计算开销太高常见做法是限制角度的候选范围比如只取距离当前锚点最近的K个样本组成三元组。资源默认K4即每个锚点取最近的4个邻居参与角度计算这样计算量从O(n^3)降到O(n*K^2)训练速度快很多。3. 环境与数据准备把CoatNet教师先“寄存”下来3.1 加载CoatNet预训练权重并冻结RKD的第一步是把教师模型完整跑一次前向拿到展平特征。CoatNet是结合卷积与自注意力的混合结构预训练权重一般从timm库加载。注意教师模型和学生的数据预处理必须完全一致否则特征分布会错位。常见做法是加载CoatNet后立刻冻结所有参数并切到eval模式同时开启torch.no_grad()。这里有个容易被忽视的点CoatNet里的BatchNorm层在eval模式下用的是running statistics如果你的数据分布和预训练数据集差异较大running stats可能不准。更稳妥的做法是让CoatNet在训练模式下前向保持BatchNorm的batch统计但不回传梯度。资源里的教师前向是这样处理的import torch import timm device cuda:0 # 加载CoatNet教师模型 teacher timm.create_model(coatnet_0_rw_224, pretrainedTrue).to(device) teacher.eval() # 冻结全部参数 for p in teacher.parameters(): p.requires_grad False def get_teacher_feat(x): # 关键点不用torch.no_grad()包裹 # 让BatchNorm使用当前batch的统计量特征更贴近实际分布 with torch.cuda.amp.autocast(): # 取倒数第二层输出作为展平特征 feat teacher.forward_features(x) # 返回 [batch, dim] return feat为什么不用torch.no_grad()因为no_grad下BatchNorm的running stats不会更新如果训练数据分布与ImageNet差异较大教师特征会有轻微漂移。资源里的做法是让教师做前向但冻结参数这样BatchNorm能利用当前batch的统计量同时梯度并不会传到教师网络。代价是会多占一点显存不过相对收益是值得的。3.2 数据预处理与batch_size选择ResNet和CoatNet的预处理基本一致ImageNet的标准化mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]Resize到224x224。但要注意数据增强策略CoatNet训的时候用了RandAugment之类的强增强ResNet如果只用弱增强学生学到的关系分布会偏简单反之如果学生也用强增强训练不稳定。资源里的做法是教师和学生共用同一套增强策略即每个batch在增强后同时喂给两个模型保证它们看到的样本对完全一样。这个细节对RKD至关重要——如果教师和学生看到的是不同增强版本的同一样本距离和角度的对应关系就被打乱了。所以线上增强必须放在数据加载器里统一完成后分叉不要在模型内部各自做随机增强。batch_size上RKD比普通蒸馏更挑。因为距离矩阵是batch内样本两两计算的batch太小、样本类别单一距离分布就会失真。常见做法是最少64起步。如果你的显存不够可以试试梯度累积或者干脆减少三元组邻居数K而不是减小batch。3.3 特征展平与维度对齐的Embedding层CoatNet的forward_features输出的是一个一维特征向量例如[batch, 1024]ResNet的卷积特征则是四维的[batch, C, H, W]需要先经过全局池化成[batch, C]再接Embedding层把C映射到1024。这里有两个注意点第一Embedding层不要做得太深。两层MLP就够输出维度取教师的特征维度中间隐藏层取教师维度的两倍即可。太深的Embedding层的表达能力太强学生会把教师特征“背下来”导致蒸馏完成后去掉Embedding层ResNet主干学到的关系被压缩在Embedding里。第二Embedding层的初始化方式。常见做法是用Xavier初始化并把最后一层的偏置置零。如果初始化不当蒸馏初期学生特征分布偏移教师太远KL散度可能一开始就饱和到0导致梯度消失。资源里给出了一个简单的初始化函数class EmbeddingMLP(nn.Module): def __init__(self, in_dim, out_dim, hid_dimNone): super().__init__() hid_dim hid_dim or out_dim * 2 self.mlp nn.Sequential( nn.Linear(in_dim, hid_dim), nn.ReLU(inplaceTrue), nn.Linear(hid_dim, out_dim) ) self._init_weights() def _init_weights(self): for m in self.mlp: if isinstance(m, nn.Linear): nn.init.xavier_uniform_(m.weight) nn.init.zeros_(m.bias) # 假设ResNet倒数第二层输出2048维 student_feat_dim 2048 teacher_dim 1024 embedding EmbeddingMLP(student_feat_dim, teacher_dim).to(device)这里in_dim要看你用的ResNet具体版本ResNet18是512ResNet34是512ResNet50是2048。确定方法很简单把ResNet去掉最后的全连接层池化后打印shape即可。4. 训练主循环分类Loss与关系Loss如何协同4.1 训练框架总览整个训练主循环并不复杂但三个loss的配比要仔细调。分类loss保证学生能完成基本任务distance_loss和angle_loss负责迁移CoatNet的“关系观感”。通常总loss是total_loss ce_loss alpha * distance_loss beta * angle_lossalpha和beta这两个权重很关键。资源里的默认值是alpha2.0和beta1.0即距离损失权重是角度损失的两倍。这个比例有一定的道理距离损失是二阶关系角度损失是三阶关系高阶关系更难学所以权重适当降低。但具体值必须结合你的数据集调整我的经验是从教师和学生特征分布差异来估计——若两者特征分布差距大权重就要偏大差距小反而要降低否则容易盖过分类loss。4.2 教师前向缓存策略蒸馏训练有一个常见的性能问题教师模型每个epoch都要重复前向计算。如果你的数据集几千张图片还好但如果是几万甚至几十万张教师的重复前向会浪费大量时间。资源里做了一个简单缓存——把所有训练样本的教师特征在训练第一遍前提前算好存成npy文件或内存字典后续训练直接从缓存取。from tqdm import tqdm def precompute_teacher_feats(teacher, data_loader, cache_path): teacher_feats {} teacher.eval() with torch.no_grad(): for images, targets in tqdm(data_loader): images images.cuda() feats teacher.forward_features(images).cpu().numpy() for idx, tgt in enumerate(targets): # 用图片路径或唯一索引做key teacher_feats[tgt.item()] feats[idx] # 保存为npy np.save(cache_path, teacher_feats) return teacher_feats每次数据加载器返回的batch里把当前图片对应的教师特征一起返回。这需要你的数据加载器能返回样本的唯一标识——常见做法是Dataset里保存图片路径用路径当key。缓存完成后训练时教师模型可以完全不参与前向显存压力骤降训练速度也能提升两倍上下。注意这个缓存的坑数据增强。如果训练时用到了RandomCrop、随机翻转这类在线增强同样的图片每次增强后的版本都不同但教师特征只缓存了一份这会造成师生输入对不齐。解决方式有两个要么缓存教师对“原始图”的特征学生也用原始图做弱增强蒸馏另外加一条强增强分支做分类要么干脆不缓存在线同时跑教师前向。资源里的做法是折中先缓存一份教师特征然后学生数据增强强度降到只有随机翻转和中心裁剪保证学生看到的和教师缓存对应的原始图在语义和空间结构上基本一致。这样做虽然学生增强变弱了但关系蒸馏的收益补偿了分类精度。4.3 完整训练循环主体下面的代码是在资源代码基础上整理出来的核心训练循环包含了三loss的计算逻辑import torch import torch.nn as nn import torch.optim as optim import torch.nn.functional as F from torch.cuda.amp import GradScaler, autocast def train_one_epoch(student, embedding, teacher, train_loader, optimizer, scheduler, scaler, alpha, beta, temperature): student.train() embedding.train() running_ce, running_dw, running_aw 0.0, 0.0, 0.0 ce_criterion nn.CrossEntropyLoss() for images, labels in train_loader: images images.cuda() labels labels.cuda() with autocast(): # 学生前向ResNet主干 分类头 logits, feat student(images) # feat: [batch, embed_dim] # 教师前向或在缓存中检索 with torch.no_grad(): t_feat teacher.forward_features(images) # [batch, teacher_dim] # 学生特征嵌入对齐 s_feat embedding(feat) # [batch, teacher_dim] # 三个loss ce_loss ce_criterion(logits, labels) dw_loss distance_wise_loss(s_feat, t_feat, temperature) aw_loss angle_wise_loss(s_feat, t_feat, temperature) total_loss ce_loss alpha * dw_loss beta * aw_loss scaler.scale(total_loss).backward() scaler.step(optimizer) scaler.update() optimizer.zero_grad() running_ce ce_loss.item() * images.size(0) running_dw dw_loss.item() * images.size(0) running_aw aw_loss.item() * images.size(0) scheduler.step() avg_ce running_ce / len(train_loader.dataset) avg_dw running_dw / len(train_loader.dataset) avg_aw running_aw / len(train_loader.dataset) print(fCE: {avg_ce:.4f} DW: {avg_dw:.4f} AW: {avg_aw:.4f}) return avg_ce, avg_dw, avg_aw代码里的混合精度训练AMP需要特别说明RKD的KL散度对数值精度不如分类loss敏感但torch.cdist在FP16下可能出现溢出导致NaN。我在实际跑的时候遇到过几次后来所采取的做法是距离和角度loss计算部分强制回到FP32with autocast(): ce_loss ce_criterion(logits, labels) # 特征转为FP32计算关系损失避免FP16下的溢出问题 s_feat_f32 s_feat.float() t_feat_f32 t_feat.float() dw_loss distance_wise_loss(s_feat_f32, t_feat_f32, temperature) aw_loss angle_wise_loss(s_feat_f32, t_feat_f32, temperature)这样做的代价是增加少量显存和计算量但换来的是稳定。AMP在蒸馏任务里并不是必需品如果你的显存充足整个训练可以直接关掉AMP减少一个变量。4.4 优化器与学习率策略优化器方面资源用的是AdamW配合余弦退火调度器。这种组合的好处是AdamW对Embedding层这类随机初始化的小模块更友好而余弦退火能让温度上限在后期慢慢降下来帮助蒸馏后期精细拟合。学习率设置上ResNet主干用线性缩放规则——基础学习率0.001乘以batch_size/256Embedding层则单独设一个较大的学习率0.005左右因为它是从头开始训的需要更快收敛。常见的踩坑是Embedding层学习率设得和主干一样导致蒸馏前期关系loss收敛极慢整个训练过程被Embedding层拖慢。对Embedding层的参数单独做参数分组是保证训练效率的重要一步optimizer optim.AdamW([ {params: student.parameters(), lr: 0.001 * (batch_size / 256)}, {params: embedding.parameters(), lr: 0.005}, {params: student.fc.parameters(), lr: 0.001 * (batch_size / 256)} ]) scheduler optim.lr_scheduler.CosineAnnealingLR(optimizer, T_maxepochs)T_max设为总训练轮数让学习率在最后一个epoch降到接近0。如果你用的数据集比较小训练轮数在50以内可以适当把base_lr调低20%左右因为RKD的关系信号在小数据上容易过拟合不仅没帮你提升泛化反而可能记住训练集的关系结构。5. 避坑排查RKD实战中的五个典型翻车现场这一章的坑全是我实际跑RKD过程中碰到的每条都是先给现象再给原因和解决方案照着排查能省很多时间。5.1 现象距离损失一开始就是NaN这个问题大概率出在FP16混合精度上。前面提过torch.cdist在FP16下计算欧氏距离时如果特征范数较大或batch内维度较高中间平方累加过程可能溢出。排查办法是直接把distance_loss的输入转成FP32。如果已经转了仍然NaN再看温度参数是不是设成了0或者负值——softmax除以0才会产生NaN。此外还有一个不起眼的原因特征里有NaN。CoatNet在加载预训练权重后如果输入数据包含异常像素值比如归一化没做好出现Inf特征就会从源头带入NaN。建议在训练循环第一轮就打印教师特征的torch.isnan().sum()确认源是干净的网络。5.2 现象distance_loss和angle_loss降了但分类accuracy一直不动这说明分类loss被两个关系loss压制得太厉害学生把所有精力都放在模仿关系上忽略了真正的分类任务。原因通常是权重alpha和beta设置过大。我查过资源默认训练日志alpha2.0、beta1.0适合中大规模数据集如果你的数据集只有几千张最好把alpha降到0.5以下。另一个可能原因是Embedding层太强学生主干已经没有梯度了。检查方法打印student的各层梯度范数如果主干部分的梯度接近0说明Embedding把损失全吸收了。解决方式是减少Embedding层的隐藏层宽度或者给Embedding层加一个小的L2正则限制它的表示能力。5.3 现象教师模型显存占用巨大batch_size被迫调小这个问题常发生在没有做特征缓存、教师也在线与学生同步前向的场景。CoatNet是Transformer结构自注意力在长序列上的显存开销高。除了提前做教师特征缓存见4.2节还有一种边训练边蒸馏的低显存方案教师forward的batch不一定要与学生一样大可以每隔两步对同一批数据算一次教师特征放到待处理队列里学生更新时从这个队列取历史特征算loss。但这么做会让师生特征不对齐实际效果略差。我一般优先做缓存缓存数据集特征总大小通常不会太大5万张图片的教师特征大约占用200MB内存完全可接受。5.4 现象师生距离分布差异巨大RKD_loss不下降这很可能是教师没有彻底冻结。如果教师的参数还在更新随着训练进行教师的特征分布漂移学生永远追不上一个移动的目标。检查方式训练循环里打印sum(p.requires_grad for p in teacher.parameters())确认是0。还有一个细节如果教师是用model.eval()和no_grad()跑的但教师内部结构里有Dropout层eval模式才正确的如果忘了evalDropout还会做随机失活同一个样本两次前向得到不同特征距离矩阵本身就带噪声了。5.5 现象Angle-wise Loss计算时维度爆炸朴素的角度损失实现会把三元组全量展开比如batch64时生成约25万个三元组每个三元组做一次向量差和夹角余弦显存和计算量会直接打满。解决方式依赖前面说的矩阵化批量计算项目源码里用的是torch.bmm和torch.norm的批量版本原理是把向量差算出来后一次性算所有点积和范数的乘积。如果你是从零开始写建议先用小batch比如16验证形状正确后再放大batch。批量计算角度损失的关键代码参考def angle_loss_batch(s_feat, t_feat, temperature1.0, k4): # s_feat: [batch, dim] batch s_feat.size(0) # 计算样本间向量差: [batch, batch, dim] s_diff s_feat.unsqueeze(1) - s_feat.unsqueeze(0) t_diff t_feat.unsqueeze(1) - t_feat.unsqueeze(0) # 取每个锚点的K个近邻避免全量三元组 # 返回角度向量做softmax与KL散度 ...这里的k就是每个锚点选近邻的数量。k越大约束越强但计算量增长越快。资源默认k4实际我在ImageNet子集上试过k8能稳定提升约0.3个点再往上就没有明显收益了。6. 进阶技巧用关系一致性指标验证蒸馏质量训练结束之后不能只盯着精度数字看。RKD真正想迁移的是CoatNet看待样本间关系的方式那么检验蒸馏效果最直接的办法是比较师生特征空间的“关系结构”是否接近。我习惯用一个轻量指标——关系一致性Relational ConsistencyRC。方法很简单取验证集的一部分特征分别计算师生的距离矩阵然后对齐行做置信度加权后的余弦相似度。余弦相似度越接近1说明距离结构越一致蒸馏越成功。def relational_consistency(s_feats, t_feats): # s_feats, t_feats: [num_samples, dim] numpy s_dist np.linalg.norm(s_feats[:, None, :] - s_feats[None, :, :], axis-1) t_dist np.linalg.norm(t_feats[:, None, :] - t_feats[None, :, :], axis-1) # 按行中心化 s_centered s_dist - s_dist.mean(axis-1, keepdimsTrue) t_centered t_dist - t_dist.mean(axis-1, keepdimsTrue) # 行间余弦相似度取均值 eps 1e-8 scores np.sum(s_centered * t_centered, axis1) / ( np.linalg.norm(s_centered, axis1) * np.linalg.norm(t_centered, axis1) eps ) return float(np.mean(scores))RC值在0.8以上通常蒸馏质量已经不错如果低于0.6说明学生的距离关系和教师还有较大差距此时优先调整温度参数和距离loss权重。配合这个指标我还会同时把学生的特征降维到二维做t-SNE可视化和教师特征并排对比。如果两组散点图的聚类形状大致相似说明关系结构确实已经迁移过去了。这一眼扫过去的数据形态比任何精度数字都直观。另外一个小技巧RKD训练完成后可以把Embedding层丢弃直接用ResNet主干的特征接一个线性分类头微调几个epoch。因为Embedding层只是对齐用的脚手架真正学到的关系知识在主干里。丢弃之后重新微调分类头有时精度还会小幅回升——这相当于先把关系结构蒸馏进主干再让主干重新适应分类目标像是给了ResNet一次“带着CoatNet视角重新学习”的机会。从那以后我跑RKD都强制走一遍关掉Embedding微调分类头的流程验证一下主干本身的关系表达能力。这步花不了多长时间但能确认真实收益来自主干而不是适配器。RKD这套思路适合那些追求更稳的蒸馏迁移方式的场景希望帮到你。本文还有配套的精品资源点击获取