ARTICLE DETAIL

资讯详情

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

PyTorch实现对比表示蒸馏CRD:从知识蒸馏到特征学习的范式跃迁

PyTorch实现对比表示蒸馏CRD:从知识蒸馏到特征学习的范式跃迁 简介面向深度学习模型压缩与部署场景一套基于Pytorch实现对比表示蒸馏CRD算法的完整项目覆盖从教师网络到学生网络的知识迁移全过程适合有一定深度学习基础、希望掌握知识蒸馏实践的开发者。压缩包共39个文件包括35个Python源码、3个Shell脚本和1个Markdown说明文档整体仅55KBPython模块涵盖数据预处理、模型定义ResNet/MobileNet/ShuffleNet等、多种蒸馏损失实现FitNet、AT、SP、KD、CRD等以及pretrain、train_teacher、train_student等可执行训练脚本Shell脚本负责获取预训练权重与运行CIFAR蒸馏实验。项目还提供流程式README说明依赖安装、代码结构、关键步骤与运行指引帮助学习者快速复现实验、理解对比学习与知识蒸馏结合的思想。目前已有190人学习借助清晰的目录组织和完整源码可将该技术迁移至实际模型压缩与嵌入式部署场景通过本项目积累模型轻量化与性能保持的实践经验。1. 知识蒸馏走到 CRD 这一步换的是比较对象知识蒸馏Knowledge Distillation这几年早已不是“训练一个学生网络模仿教师输出”这么简单。早期 Hinton 提出的 KD 用软标签传递类别概率后来 FitNets 把中间特征图拉齐再到 Attention Transfer、SPKD 这类基于特征关系的方法本质上都在做同一件事找一个合适的“比较对象”让学生向教师看齐。可以这么说蒸馏效果的上限很大程度上取决于“拿什么比”。而对比表示蒸馏 CRDContrastive Representation Distillation换了一个比较维度——不再直接让学生特征逼近教师特征而是通过对比学习的方式让学生学会区分“与教师特征匹配的正样本”和“不匹配的负样本”从而把教师的高阶语义结构迁移给学生。这个思路源于 2020 年 CVPR 的 Contrastive Representation Distillation 论文实践表明在 CIFAR 和 ImageNet 系列任务上CRD 往往能比传统特征蒸馏高出几个点。本篇文章围绕“基于 PyTorch 实现 CRD”展开会先把知识蒸馏的范式变化讲清楚再剖析 CRD 的核心机制——对比学习的引入、正负样本对怎么构造、InfoNCE 损失如何在蒸馏场景下发挥作用然后落到一份可以直接运行的 PyTorch 代码上最后说透参数设置和踩坑点。2. 蒸馏范式演变从软标签到特征对比CRD 解决了什么2.1 Hinton KD 的本质与瓶颈经典 KD 的损失由两部分组成学生与硬标签的交叉熵加上学生与教师软标签的 KL 散度。教师网络的 softmax 输出经过温度系数 T 软化之后包含了类别间的相似结构信息——比如“猫”和“狗”的置信度分布比“猫”和“卡车”更接近这种暗知识是硬标签给不了的。# 经典 KD 损失核心伪代码风格 import torch import torch.nn.functional as F def kd_loss(student_logits, teacher_logits, labels, T4.0, alpha0.7): # 教师软标签软化除以温度 T再走 softmax teacher_soft F.softmax(teacher_logits / T, dim1) # 学生同样软化后计算 KL 散度 student_log_soft F.log_softmax(student_logits / T, dim1) kd F.kl_div(student_log_soft, teacher_soft, reductionbatchmean) * (T * T) # 硬标签交叉熵 ce F.cross_entropy(student_logits, labels) return alpha * kd (1 - alpha) * ce代码里T * T这个缩放是因为 KL 散度对温度敏感回传梯度时需要把温度影响补偿回来否则软标签的梯度会随 T 增大而稀释。alpha控制软标签损失的权重实践中一般取 0.7 到 0.9。这个方案的局限性在于——它只在输出层做比较对中间层语义的迁移无能为力。如果你在 ImageNet 这种大尺度分类任务上训练 ResNet 学生网络单纯 KD 往往只能小幅提升因为学生网络自己学到的特征表达和教师差着好几个抽象层级。2.2 特征蒸馏的方法与共同缺陷为了让学生网络学到中间层特征后续工作开始对齐特征图。FitNets 用 L2 损失直接匹配教师和学生的中间特征Attention Transfer 对齐注意力图SPKD 用样本间的相似度矩阵做约束。# FitNets 特征对齐的最小实现 def fitnets_loss(student_feat, teacher_feat): # 学生特征通常要加一个卷积投影层把通道数对齐到教师 # 然后做 L2 归一化防止数值尺度影响 student_feat F.normalize(student_feat, dim1) teacher_feat F.normalize(teacher_feat, dim1) return F.mse_loss(student_feat, teacher_feat)这类方法的共同问题是强制特征逐像素或逐通道对齐对学生的容量要求很高。学生网络本身参数少硬去拟合教师的特征分布很容易丢失自己的表达自由度——这就是“过度约束”问题。特征蒸馏在教师和学生结构差异较大时表现尤其不稳定比如教师是 ResNet-34、学生是 MobileNetV2这两者的特征空间分布差异极大L2 对齐就像让一个小学生模仿书法家的笔锋模仿不好反而连自己写字的手感都丢了。2.3 CRD 的核心思路互信息视角下的蒸馏CRD 换了一个根本性的思路——把蒸馏任务重构成一个对比学习任务。它不再要求学生的特征在数值上接近教师而是要求学生网络具备和教师网络一样的“辨别能力”给定一个样本的教师特征学生能否从候选集合中找出和它最匹配的样本特征。这个想法来自互信息Mutual Information最大化。教师特征和学生特征之间的互信息越大说明学生保留的关于教师的语义信息越多。但互信息在高维连续空间中难以直接计算CRD 巧妙地把互信息最大化转化为一个对比学习目标让学生学会区分“来自同一输入样本的教师特征”正样本对和“来自其他输入样本的教师特征”负样本对。更直白地说传统蒸馏是“我教你变成我”CRD 是“我考你认不认识我”。前者是回归问题后者是判别问题。判别任务往往比回归任务更容易学习尤其是在特征维度很高时——L2 回归在高维空间里会遭遇维度灾难而对比学习天然在高维空间里有优势。这就是 CRD 取名的由来Contrastive Representation Distillation对比表示蒸馏。3. CRD 核心机制拆解正负样本对、InfoNCE 损失与记忆库3.1 对比表示蒸馏的框架结构先给出 CRD 的整体结构。它由一个教师网络、一个学生网络、一个投影头Projection Head和一个记忆库Memory Bank组成。教师网络负责提取教师特征向量学生网络负责提取学生特征向量。但两者并不直接做特征对齐而是各自接一个投影头把特征映射到一个低维对比空间通常是 128 维然后在低维空间里做对比学习。为什么要有这个投影头因为原始特征维度动辄上千直接在这个空间做对比学习一方面计算量太大另一方面高维空间里样本距离趋于均匀对比学习的信号会被稀释。投影头本质上是把“用于对比的语义”和“用于分类的语义”做一个解耦。记忆库则用来保存每个训练样本的教师特征。因为负样本对需要从“其他样本”中采样如果每个 batch 只取当前 batch 内的负样本数量太少对比学习的难度不够。记忆库维护一个全局的样本特征队列能够提供足够多且足够“难”的负样本。从 PyTorch 的视角来看CRD 不需要特殊的网络结构只要求教师和学生都输出特征向量然后在特征向量的基础上追加一个可训练的线性投影层即可。3.2 InfoNCE 损失在蒸馏中的形态CRD 采用的损失函数是 InfoNCEInformation Noise-Contrastive Estimation的变体。原始的 InfoNCE 损失如下# InfoNCE 损失参考实现用于理解CRD 中会做修改 def info_nce_loss(query, positive_key, negative_keys, temperature0.07): # query: [N, D] 学生投影后的特征 # positive_key: [N, D] 匹配的教师投影特征 # negative_keys: [N, K, D] K 个负样本的教师投影特征 N query.shape[0] # 计算正样本对的相似度 logits pos_logits torch.sum(query * positive_key, dim1, keepdimTrue) / temperature # [N, 1] # 计算负样本对的相似度 logits neg_logits torch.bmm(negative_keys, query.unsqueeze(2)).squeeze(2) / temperature # [N, K] # 拼接正负 logits计算交叉熵 logits torch.cat([pos_logits, neg_logits], dim1) # [N, K1] labels torch.zeros(N, dtypetorch.long).to(query.device) return F.cross_entropy(logits, labels)这段代码的逻辑是对第 i 个样本学生投影特征 query_i 与对应的教师投影特征 positive_key_i 组成正样本对与其他样本的教师投影特征组成负样本对。损失目标是让正样本对的相似度余弦相似度除以温度系数后尽可能大负样本对的相似度尽可能小。这本质上是一个 (K1) 类的分类问题。CRD 和标准 InfoNCE 有一个关键区别标准 InfoNCE 中 query 和 key 来自同一张图像的不同增强视角而 CRD 中 query 来自学生网络、key 来自教师网络。所以 CRD 学习到的是“学生特征到教师特征的语义映射”而不是“特征到自身的增强不变性”。这个区别决定了 CRD 的判别边界是教师特征的分布而不是数据增强的分布。3.3 维度匹配与特征对齐策略实际工程里会遇到一个非常具体的问题教师特征维度与学生特征维度不一致。比如教师是 ResNet-34最后一个池化层输出 512 维学生是 ResNet-18输出 512 维这还算一致。但如果是 MobileNetV2 输出 1280 维、教师输出 2048 维就需要处理维度匹配。CRD 的常见做法是在教师和学生各接一个投影层MLP把特征投影到同一个低维空间。具体来说投影头通常是一个两层 MLP第一层线性变换到 512 维接 ReLU第二层线性变换到 128 维。代码实现如下import torch.nn as nn class ProjectionHead(nn.Module): def __init__(self, in_dim, hidden_dim512, out_dim128): super().__init__() self.fc1 nn.Linear(in_dim, hidden_dim) self.relu nn.ReLU(inplaceTrue) self.fc2 nn.Linear(hidden_dim, out_dim) def forward(self, x): x self.fc1(x) x self.relu(x) x self.fc2(x) # 对输出做 L2 归一化稳定训练 return F.normalize(x, dim1)注意最后一步对投影后的特征做了 L2 归一化。原因在于InfoNCE 损失中计算的是余弦相似度而余弦相似度只关心方向不关心模长。如果特征模长差异过大模型可以通过增大模长来“作弊”让损失值虚低但不具备语义区分性。L2 归一化把特征限制在单位超球面上模型只能通过学习有意义的特征方向来降低损失。如果不想引入额外的投影层另一个做法是直接在各层特征上加 AdaptiveAvgPool 把空间维度压成 1x1再用 1x1 卷积调整通道数。但这个做法在教师和学生网络结构差异大时效果不如投影头因为缺少非线性变换两个特征空间的分布无法被充分对齐。3.4 负样本采样与记忆库实现CRD 对负样本的数量和质量都有要求。早期实验表明负样本数量从 256 提升到 16384 时精度持续上升。但每轮迭代都重新计算那么多教师特征不现实所以 CRD 引入记忆库机制把之前若干轮训练中算出的教师特征存入队列当前 batch 的负样本从队列中随机采样。这里的实现有几个细节值得注意。第一记忆库的容量要远大于 batch size一般设置为整个训练集的大小。第二负样本只会从记忆库中取不会取当前 batch 内其他样本的教师特征——因为当前 batch 的特征还没被存入记忆库而且 batch 内样本类别接近作为负样本区分度不够。第三记忆库中的特征需要是“教师当前参数”计算出的特征而不是历史参数的结果。实际实现中教师网络是固定的不训练所以记忆库中的特征不会过期但如果教师网络也参与更新就需要做动量更新或者定期刷新。import random class MemoryBank: def __init__(self, capacity, feature_dim): self.capacity capacity # 随机初始化记忆库后续逐步替换 self.bank torch.randn(capacity, feature_dim) self.bank F.normalize(self.bank, dim1) self.ptr 0 def update(self, features, indices): # features: [N, D] 教师投影特征 # indices: [N] 样本对应的全局索引 for i, idx in enumerate(indices): self.bank[idx] features[i] def sample_negatives(self, query_indices, num_negatives): # 排除当前 batch 中的样本从记忆库中采样负样本索引 candidates [i for i in range(self.capacity) if i not in query_indices] neg_indices random.sample(candidates, num_negatives) return self.bank[neg_indices]记忆库的容量对应数据集的样本总数。在 CIFAR-100 上就是 50000在 ImageNet 上就是 128 万。显存足够大的情况下直接把记忆库放在 GPU 上是可行的如果显存吃紧则放在 CPU 上通过索引访问时用.cuda()转到 GPU。负样本数量num_negatives通常是 256 到 4096 之间。数量太少对比任务过于简单学到的特征区分度低数量太大训练时间和显存开销增长。CRD 原文实验里 4096 个负样本就能在 CIFAR-100 上取得不错的效果再往上提升有限。4. PyTorch 实现 CRD 的完整训练流程4.1 数据准备与超参数设置先用 CIFAR-100 来跑通整个流程。CIFAR-100 有 100 个类别、50000 张训练图、10000 张测试图图像尺寸 32x32适合做快速实验。超参数的设置如下参数名推荐值说明教师网络ResNet-34通常是用更大模型或预训练模型学生网络ResNet-18结构相对简单便于观察蒸馏效果优化器SGDmomentum 0.9权重衰减 5e-4与常规图像分类一致学习率0.05余弦退火批量大小 64 时的经验值训练轮数240CIFAR-100 上充分收敛需约 200 轮以上投影头输出维度128对比空间维度温度系数 tau0.07InfoNCE 损失的温度负样本数量 K4096从记忆库中采样的负样本数KD 损失权重0.8蒸馏损失和学生分类损失的平衡系数这里的tau是一个需要认真调节的参数。它控制相似度分布的陡峭程度。tau越小softmax 后的分布越尖锐模型对正负样本的区分要求越高tau越大分布越平缓训练信号越温和。0.07 是自监督学习领域一个常见的经验值。4.2 教师网络定义与预训练CRD 需要先有一个训练好的教师网络。教师网络可以用常规的交叉熵损失在 CIFAR-100 上训练 240 轮得到。这里不赘述完整的教师训练过程但有一个工程细节保存教师网络时最好同时保存state_dict和完整的模型结构信息后续加载时直接load_state_dict即可。# 加载预训练教师 teacher resnet34(num_classes100) teacher.load_state_dict(torch.load(teacher_cifar100.pth)) teacher.eval() for p in teacher.parameters(): p.requires_grad False这里有几个关键点。第一教师网络必须固定在 eval 模式因为 BatchNorm 在 train 和 eval 模式下的行为不同。如果教师网络错放在 train 模式BatchNorm 会持续用当前 batch 的统计量做归一化导致输出的特征波动较大记忆库中的特征也会变得不稳定。第二教师网络的参数必须冻结否则反向传播会把梯度传到教师网络中但这部分梯度不会更新因为 requires_gradFalse却会额外占用显存。为了节省显存可以进一步使用torch.no_grad()上下文或者把教师网络放进torch.inference_mode()。4.3 CRD 训练循环主体训练循环是 CRD 工程实现的核心。这里给出一个完整的 PyTorch 训练循环骨架for batch_idx, (images, labels, indices) in enumerate(train_loader): images, labels images.cuda(), labels.cuda() indices indices.cuda() # 教师特征提取不计算梯度 with torch.no_grad(): teacher_feat teacher(images) # [N, D_t] teacher_proj teacher_proj_head(teacher_feat) # [N, 128] # 学生特征提取正常计算梯度 student_logits, student_feat student(images) student_proj student_proj_head(student_feat) # [N, 128] # 从记忆库采样负样本 neg_features memory_bank.sample_negatives( indices.cpu().tolist(), num_negativesK ).cuda() # [K, 128] # 计算 CRD 损失 crd_loss crd_criterion(student_proj, teacher_proj, neg_features) # 分类损失 ce_loss F.cross_entropy(student_logits, labels) # 总损失 loss crd_loss * 0.8 ce_loss * 0.2 optimizer.zero_grad() loss.backward() optimizer.step() # 更新记忆库 memory_bank.update(teacher_proj.detach(), indices)这段代码的关键在于理解数据流。教师网络输出原始特征teacher_feat后先经过教师侧投影头得到对比空间的特征学生侧同理。crd_criterion接收学生投影特征、教师投影特征以及从记忆库中采样的负样本计算 InfoNCE 损失。ce_loss是学生网络的分类损失它不会对教师网络产生任何梯度。损失权重 0.8 对应 CRD 蒸馏损失的占比。为什么学生分类损失占比只有 0.2因为对比表示蒸馏提供的是特征级别的监督信号分类损失只是兜底防止学生网络完全忽略分类任务。如果分类损失占比过大学生网络会倾向于只优化分类边界投影头的语义特征会退化。如果分类损失占比过小学生网络的特征可能缺乏类别区分度。4.4 CRD 损失函数完整实现def crd_criterion(student_proj, teacher_proj, neg_features, tau0.07): student_proj: [N, D] 学生投影特征L2归一化后 teacher_proj: [N, D] 教师投影特征L2归一化后 neg_features: [K, D] 从记忆库采样的负样本特征L2归一化后 N student_proj.shape[0] K neg_features.shape[0] # 正样本对 logits学生对教师逐样本点积 pos_logits torch.sum(student_proj * teacher_proj, dim1, keepdimTrue) / tau # [N, 1] # 负样本对 logits学生对全部负样本点积 neg_logits torch.matmul(student_proj, neg_features.t()) / tau # [N, K] # 拼接正负 logits logits torch.cat([pos_logits, neg_logits], dim1) # [N, K1] labels torch.zeros(N, dtypetorch.long).cuda() return F.cross_entropy(logits, labels)student_proj和teacher_proj都经过了 L2 归一化所以torch.sum(student_proj * teacher_proj, dim1)就是余弦相似度。除以tau之后相似度的数值范围会被放大到约 ±14因为 tau0.07这个范围下 softmax 的梯度比较有区分性。如果tau过大所有相似度都趋近于 0交叉熵损失的梯度趋近于零模型几乎学不到东西如果tau过小只有最难的负样本能产生梯度训练不稳定。labels 是一个全零向量因为逻辑上第 0 列永远是正样本列第 1 到 K 列是负样本列。这个写法虽然简洁但有一个安全隐患如果student_proj * teacher_proj的点积结果特别小比如负的模型依然能正常学习但不同样本的损失值差异会很大需要通过监控 loss 的量级来判断训练是否健康。4.5 投影头与特征提取的代码结构把整套代码组装起来核心模块就是三个类教师网络 教师投影头、学生网络 学生投影头、记忆库。在 PyTorch 中可以用一个nn.Module把学生网络和投影头包装起来方便管理参数class StudentModel(nn.Module): def __init__(self, num_classes100, feat_dim512, proj_dim128): super().__init__() # ResNet-18 最后一个池化层输出 512 维特征 self.backbone resnet18(num_classesnum_classes) self.proj_head ProjectionHead(feat_dim, proj_dimproj_dim) def forward(self, x): # 提取特征图前的特征通过 hook 实现或直接返回 logits, feat self._forward_feat(x) proj self.proj_head(feat) return logits, proj实际写代码时需要注意resnet18默认forward返回的是分类 logits要获取倒数第二层特征需要把网络拆开或用 register_forward_hook。最简单的方式是修改 resnet 的forwarddef forward(self, x): x self.conv1(x) x self.bn1(x) x self.relu(x) x self.maxpool(x) x self.layer1(x) x self.layer2(x) x self.layer3(x) x self.layer4(x) feat self.avgpool(x) feat torch.flatten(feat, 1) logits self.fc(feat) return logits, feat这个改动不会影响原有的分类性能只是多返回了一个中间特征。要注意的是如果教师和学生都这样改教师网络前向时返回的feat是在torch.no_grad()下计算的学生网络的feat则正常走梯度链路。5. 调参、验证与常见训练陷阱验证 CRD 是否真正生效要关注三类指标。第一类是学生网络的 Top-1 准确率这是最终指标第二类是 CRD 损失本身的下降曲线第三类是学生投影特征与教师投影特征的余弦相似度分布。单独监控 CRD 损失没有意义它只代表对比任务是否被解决不代表分类性能是否提升。需要把损失曲线和准确率曲线放在同一张图里看如果 CRD 损失下降但准确率不动大概率是投影头把特征投影到了与分类无关的方向上。这也引出了训练中最大的一个陷阱投影头与分类头在竞争特征表达能力。学生网络的主干特征既要用于分类又要用于对比蒸馏。如果投影头的梯度信号过强主干特征会被拉向“区分教师特征”的方向而牺牲分类边界。缓解方式有两种一是把 CRD 损失权重从 0.8 降到 0.5 或 0.3 观察效果二是让投影头输入的不是主干最终特征而是主干某一中间层特征使得分类和对比使用不同层级的特征表示。另一个常见问题是 BatchNorm 统计量的扰动。学生网络的 BatchNorm 在训练中会不断更新 running_mean 和 running_var这会影响投影头的输入分布。如果训练过程中 CRD 损失出现周期性震荡优先检查是否是因为学习率过大导致特征分布剧烈变化。可以尝试给学生网络单独设置一个较小的学习率或者使用 AdamW 替代 SGD——虽然这不符合原文的设置但在调参阶段确实更容易稳定收敛。温度系数 tau 的敏感性比很多人预期的高。0.07 不是万能值当数据集换成 ImageNet 这种类间高度重叠的数据时可以尝试把 tau 调到 0.1 或 0.15。一个实用的判断方法是打印出每个 batch 中正样本 logits 的均值。如果均值长期低于 2说明 tau 太大模型没有充分区分正负样本如果均值高于 6说明 tau 太小模型对负样本完全无感。负样本数 K 对显存和性能的影响也需要按实际调整。K4096 时显存占用大约是 4096×128×4 字节约 2MB本身不大但教师网络前向传播和特征存储需要额外显存。显存不足时优先降低 K 而不是降低 batch size因为对比学习对 batch size 的敏感度更高。如果想要更好的性能可以尝试把记忆库替换成 MoCo 风格的动量队列——但这会引入额外的超参数动量系数与 CRD 原版的设计已经有出入需要谨慎评估。最后一个验证技巧将教师和学生投影后的特征用 t-SNE 可视化。如果 CRD 真正生效同类别样本的学生特征和教师特征在投影空间中应该形成相似的簇结构而不是完全打散。这个可视化不能直接量化模型性能但能帮你快速判断是否出现了特征崩塌——即所有样本的投影特征都收敛到同一个点。如果发生特征崩塌优先检查投影头是否缺少 L2 归一化或者温度系数是否过小。本文还有配套的精品资源点击获取
返回列表