CRD对比表示蒸馏Pytorch工程解析:从原理到CIFAR-100实战
简介基于PyTorch实现的对比表示蒸馏CRD算法实战项目面向深度学习模型压缩和加速需求适合想掌握知识蒸馏原理、教师网络与学生网络知识迁移的开发者与研究者。压缩包共39个文件包含35个Python脚本、3个Shell脚本和1个Markdown文档整体约55KB代码覆盖数据预处理、模型定义、CRD对比记忆与损失计算、蒸馏训练与评估流程Shell脚本可辅助下载预训练教师模型并一键运行CIFAR蒸馏实验。项目不仅实现CRD还内置KD、FitNet、AT、SP、FSP、CC、NST、VID、RKD等多种蒸馏算法并提供ResNet、VGG、MobileNet、ShuffleNet、WideResNet等骨干网络便于横向比较与二次扩展不同实验脚本、模型仓库与算法模块相互分离目录逻辑清晰可灵活替换骨干网络、蒸馏损失和数据集配置。已有190人学习配合说明文档和运行脚本适合希望快速上手知识蒸馏与模型轻量化实战的读者。1. CRD 对比表示蒸馏的 Pytorch 工程为什么值得拆知识蒸馏里最常遇到的一个场景是你已经训练好一个 WRN-40-2想在移动端换成一个 ResNet-8。传统 KD 让学生学教师的软标签确实能涨点但学生拿到的只是最后输出的概率教师中间层对视觉结构的理解几乎没传过去。CRD 把蒸馏后移用对比学习拉近师生表示空间同时让负样本对互相推开学生不再只是模仿结果而是学教师的特征判别力。这个项目给出了一整套基于 Pytorch 的 CRD 实现里面有完整的 distiller_zoo、crd/memory.py、crd/criterion.py 和训练脚本适合想理解对比学习如何落到模型压缩以及如何在 CIFAR-100 上跑通完整知识蒸馏流程的开发者。重点是源码把数据准备、教师训练、蒸馏训练、方法对照全串起来了不用自己再东拼西凑。2. CRD 对比表示蒸馏的核心实现memory.py、criterion.py 与 loops.py 的配合2.1 先想清楚对比表示蒸馏的损失在解决 logits 蒸馏的什么问题用 Pytorch 实现知识蒸馏最简单的 KD 损失是# 教师 logits 与学生 logits 的 KL 散度配合温度 T loss_kd F.kl_div( F.log_softmax(student_logits / T, dim1), F.softmax(teacher_logits / T, dim1), reductionbatchmean, ) * T * T这里温度 T 把硬标签平滑成软标签学生学到的是教师对类别关系的判断但梯度几乎集中在最后一层分类器卷积特征提取器的更新信号很弱。CRD 想传的是教师中间表示本身。它在特征层拉近同一个样本的师生表示同时要求该表示不要和大量负样本特征混在一起本质是 InfoNCE 风格的对比损失。这个项目里的 crd/criterion.py 就是承接这件事的模块。不使用直接 MSE 对齐的原因也很实际直接对齐师生特征会强迫学生每个维度都复刻教师忽略了特征内部的相关性对比损失只约束相对相似度负样本越多、约束越强而且对高维视觉表示更稳。工程上把 memory bank 和 loss 分开放在 crd/memory.py 与 crd/criterion.py 两个文件里说明设计者希望样本索引维护和损失计算解耦。后面训练脚本 train_student.py 通过这两个模块组合出完整的 CRD 损失逻辑很干净。2.2 memory_bank 的目录索引与特征更新逻辑对比损失需要大量负样本常规 batch 只有 64 或 128 张图远不够用。CRD 的做法是维护一个 memory bank保存每个训练样本最近的表示特征。打开 crd/memory.py核心数据结构大概是这样class ContrastiveMemoryBank: def __init__(self, size, feat_dim, use_momentum): self.features torch.randn(size, feat_dim).div(math.sqrt(feat_dim)) self.targets torch.zeros(size, dtypetorch.long) self.ptr 0 self.use_momentum use_momentum torch.no_grad() def update(self, features, targets, idx): # 用当前 batch 的特征覆盖对应样本的历史特征 for i, k in enumerate(idx): self.features[k] features[i].detach() self.targets[k] targets[i]逻辑并不复杂features的每一行对应一个训练样本的表示idx是样本在数据集里的索引。计算某个 anchor 时负样本直接从self.features里随机抽取。use_momentum参数决定更新方式纯覆盖更新速度快但训练初期特征变化大可能出现 memory bank 与当前模型特征不一致的问题。实践里更稳的做法是先让教师网络把整份训练集的特征过一次填好初始化 bank训练中再逐步覆盖。这个项目根目录下的 pretrain.py 承担的正是这种初始化而不是去训练教师模型。2.3 criterion.py 里 InfoNCE 风格的对比损失是怎么计算的memory bank 准备好之后损失函数要回答一个关键问题怎么判断学生特征和教师特征属于同一个样本。项目路线是在统一 embedding 空间里用内积求相似度然后做 cross entropy。伪代码如下class ContrastiveDistillLoss(nn.Module): def __init__(self, nce_k16384, nce_t0.07, feat_dim128): super().__init__() self.nce_k nce_k # 负样本数量 self.nce_t nce_t # 温度系数控制分布锐度 def forward(self, student_f, teacher_f, idx): # student_f: [B, feat_dim]teacher_f 同形状 # 正样本对teacher_f 与 student_f 的逐样本相似度 l_pos (student_f * teacher_f).sum(dim1, keepdimTrue) / self.nce_t # 负样本从 memory bank 中随机取 batch student_f.shape[0] neg_inds torch.randint(0, memory_bank.size, (batch, self.nce_k)) neg_feats memory_bank.features[neg_inds] l_neg torch.bmm(neg_feats, student_f.unsqueeze(2)).squeeze(2) / self.nce_t # 正负样本拼接后过 softmax取 log 概率 logits torch.cat([l_pos, l_neg], dim1) # [B, 1 K] labels torch.zeros(batch, dtypetorch.long) loss F.cross_entropy(logits, labels) return loss这里有两个容易踩的细节。第一teacher_f和student_f必须经过 L2 归一化因为内积上限由模长决定不归一化会让模型靠尺度作弊第二负样本索引neg_inds最好每个 anchor 独立采样避免同一 batch 内多个样本共享同一批负样本造成相关性。criterion.py 正式实现里还会对教师特征做投影头让教师特征维度和学生一致同时削弱教师特征方差过大的问题。nce_k越大越接近理想负样本分布但显存和时间开销也明显上升后面消融部分会展开。2.4 loops.py 中 train_step 如何把 CRD 损失拼进总 loss工程里真正跑训练的是 loops.py 里的 train_step。它不会和某一个具体 loss 绑定而是读取命令行指定的--distill参数从 distiller_zoo 里选对应的蒸馏器。以 CRD 为例一次迭代会同时算出分类损失、logits 蒸馏损失和对比损失loss_total loss_cls alpha * loss_kd beta * loss_crd放在循环里看更清楚for data, target, index in train_loader: data, target data.cuda(), target.cuda() feat_s, logit_s model(data) with torch.no_grad(): feat_t, logit_t model_t(data) feat_t [f.detach() for f in feat_t] loss_cls F.cross_entropy(logit_s, target) loss_div distiller.logits_loss(logit_s, logit_t) loss_crd distiller.feature_loss(feat_s[-1], feat_t[-1], index) loss loss_cls args.alpha * loss_div args.beta * loss_crd optimizer.zero_grad() loss.backward() optimizer.step()这里的index很关键CRD 的对比损失必须知道当前 batch 样本在 memory bank 里的位置才能把正样本对取出来。如果你的自定义数据 loader 不返回索引CRD 会直接报错或学不到东西。同时注意 teacher 特征整体被torch.no_grad()包住教师只负责提供目标不参与梯度回传这也是蒸馏和联合训练的一个重要区别。3. 在 CIFAR-100 上把 CRD 蒸馏跑起来预训练教师与 train_student.py 的完整流程3.1 环境准备从 pytorch 安装到数据自动下载先确认 Pytorch 版本。CRD 的代码大多基于 nn.Module、torch.cat 这些基础功能1.8 以上的稳定版本都能跑。推荐直接按 Pytorch 官网提供的命令安装 CUDA 版本pip install torch torchvision --index-url https://download.pytorch.org/whl/cu118项目自带的数据处理模块会在第一次运行时下载 CIFAR-100。网速不理想时也可以提前把数据解压到对应数据目录代码检测到已解压的文件就会直接复用。需要注意的不是 torch 版本而是 GPU 显存。教师用 WRN-40-2、学生用 ResNet-8 时CRD 默认的 16384 个负样本会占掉额外几百 MB 显存第五章会讲怎么降档。3.2 训练教师网络train_teacher.py 的参数写法蒸馏的前提是有一个可靠的教师。这个项目提供 train_teacher.py支持 cifar100 和 imagenet 两种数据集。最常用的一条命令是python train_teacher.py --dataset cifar100 --model wrn_40_2 --batch_size 64 --lr 0.05 --epochs 240 --trial 1--model可以换成 resnet56、vgg13 等--trial用来区分同一组参数下的重复实验checkpoint 会写到教师模型保存目录。为了公平对比CRD 实验里所有学生超参数一般保持固定只换蒸馏方法所以教师训练不需要反复调参。如果不想花两三个小时重新训练可以直接用 scripts/fetch_pretrained_teachers.sh 拉取预训练好的权重脚本会放到约定目录学生训练时用--path_t指向。3.3 用 fetch_pretrained_teachers.sh 或已有 checkpoint 快速进入蒸馏脚本本质是下载 Teacher checkpoint 的循环wget -P save/teacher_models/ https://example.com/pretrained/wrn_40_2.pth执行前先看 README 里写的具体下载地址确认目录结构是否符合项目预期。如果你有其他项目训练好的教师只需要把权重转成torch.load能直接读的 state_dict并且和学生模型使用同一套预处理流程即可。常见的误区是直接用不同分辨率或不同归一化参数训练出的教师来蒸馏后面的学生对不齐特征分布效果会明显下降。3.4 核心蒸馏命令run_cifar_distill.sh 中 CRD 参数的解释run_cifar_distill.sh 把完整的蒸馏命令包在了一个脚本里。以 CRD 为例核心命令可以缩成这样python train_student.py \ --path_t ./save/teacher_models/wrn_40_2_best.pth \ --distill crd \ --model resnet8 \ --dataset cifar100 \ --batch_size 64 \ --lr 0.05 \ --epochs 240 \ --crd_nce_k 16384 \ --crd_nce_t 0.07 \ --crd_feat_dim 128 \ --alpha 0.0 --beta 0.8--distill crd会在 distiller_zoo 中注册 CRD 蒸馏器--path_t指定教师权重路径。--crd_nce_k是负样本数量默认 16384接近 CIFAR-100 训练集一半的规模--crd_nce_t是温度调小会让 softmax 更尖锐正负样本差距更大--crd_feat_dim决定投影后的表示维度。alpha是 logits 蒸馏权重beta是 CRD 对比损失权重默认命令里alpha0.0说明这个配置想观察纯 CRD 的效果不加 KD 干扰。动手时可以先用这个基线之后再组合 alpha 和 beta。3.4.1 蒸馏过程中观察什么训练日志会周期性打印 loss_cls、loss_kd、loss_crd。如果 loss_crd 在下降说明学生表示的区分能力在变好同时看验证集 acc 曲线有没有稳步上升。CIFAR-100 上 ResNet-8 直接从随机初始化训练top-1 大概在 60% 出头CRD 蒸馏后通常会更高一些具体多少取决于教师强弱、负样本数和随机种子。要验证蒸馏是否生效最直接的办法是再跑一个不使用教师权重的普通训练然后把两条 val acc 画在一起对比。4. 在 distiller_zoo 里对照 KD、FitNet、AT、SP 等十种蒸馏方法定位 CRD 的收益4.1 distiller_zoo 的接口约定每个 loss 文件都长什么样distiller_zoo 文件夹里集中了 FitNet.py、AT.py、SP.py、NST.py、VID.py、RKD.py、CC.py 等十余个蒸馏实现。它们都归到同一个蒸馏器基类下对外暴露统一的获取 logits 差值和特征差值的接口。以 FitNet 为例它的核心逻辑是选教师和学生某一层特征然后算 MSE# FitNet 简化实现 def feature_loss(self, f_s, f_t): f_s self.embed_s(f_s) # 学生特征投影到教师通道数 f_t f_t.detach() return F.mse_loss(f_s, f_t)每次切换方法只需要改训练命令里的--distill数据加载、日志、checkpoint 都不动。这么做的好处是比较 CRD 和 SP、AT 的差异时能保证除了 loss 以外的变量完全一致。自己搭实验最容易翻车的地方就是每个方法单独写一套训练循环最后结果差异来自优化器迭代次数或数据增强而不是蒸馏方法本身。4.2 切换蒸馏方法从 --distill kd 到 --distill crd在同一个工程里做横向对比命令格式固定python train_student.py --path_t teacher --distill kd --model resnet8 --dataset cifar100 --trial 1 python train_student.py --path_t teacher --distill at --model resnet8 --dataset cifar100 --trial 1 python train_student.py --path_t teacher --distill crd --model resnet8 --dataset cifar100 --trial 1不同方法有各自的超参数需要额外关注。AT 的损失权重通常要调到 1000 级别才能和 KD 的 0.1 量级匹配SP 需要设置相似性尺度VID 有可学习的方差项。跑之前先看 README 或脚本里的参考值。CRD 在参数数量上相对友好主要就是--beta、--crd_nce_k、--crd_nce_t三个从 0.8、16384、0.07 起步基本不会出大问题。4.3 做一组可控的 CRD 消融负样本数、温度系数、中间特征尺度消融实验是理解 CRD 最可靠的手段。固定教师和学生、固定 240 epochs只动一个变量实验组crd_nce_kcrd_nce_tbetaval acc基线163840.070.8待记录少负样本40960.070.8待记录低温度163840.020.8待记录高温度163840.20.8待记录加 KD 混合163840.070.8 alpha0.9待记录为了控制随机性同一组至少跑两个不同 seedtrial 取 1、2、3最后报告均值。这样能看出项目在特定数据增强、优化器下的灵敏度。如果 K 降到 4096 后准确率几乎不变说明当前任务不需要那么多负样本可以从显存优化角度继续压缩。4.3.1 失败时看什么如果 loss_crd 不降先确认--crd_feat_dim和模型实际输出维度一致再检查 memory bank 的更新频率。一个常见坑是 teacher 特征在 CRD 里默认取feat_t[-1]也就是最后一个卷积 block 的输出而某些模型在 forward 最后多做了 GAP 和 reshape导致维度对不上这时需要去看 models 目录下注册的 feat 获取方式。4.4 读日志评估蒸馏是否生效训练日志里至少要看三个值交叉熵损失、蒸馏损失、验证集 acc。如果交叉熵已经收敛但蒸馏损失还在缓慢下降说明学生还能继续从教师身上拿信息。另一个技巧是每固定 epoch 保存一次学生 checkpoint单独在验证集上测试避免只看训练 loss 的波动。scripts 里给的 run_cifar_vanilla.sh 就是为了跑一个“无蒸馏基线”训练超参和学生完全一致只是去掉教师路径。和这个基线比CRD 的收益才可信。5. 部署前压缩 CRD 工程的显存与特征取层技巧5.1 显存不够先动这三个参数batch_size 64、nce_k16384时对比损失里torch.bmm的中间矩阵大小是 batch × nce_k × feat_dim显存随 K 线性增长。显存有限时按顺序降先降--crd_nce_k到 4096记录 acc 损失再降--crd_feat_dim到 64这会减小投影头参数和矩阵规模最后才动 batch_size因为 CRD 正样本对数与 batch 有关batch 太小会让单步梯度方差变大。如果还要再省可以在 train_student.py 里对 student feature 进入投影头前保持正常回传teacher feature 始终 no_grad避免多余中间变量。5.2 非标准学生模型的特征出口怎么改压缩到非标准结构时最容易出 bugfeat_t[-1]拿到的可能不是预期特征。改法是在模型定义文件的 forward 里增加返回值例如return logits, feat再在蒸馏器里按名字取。这个项目 models 目录下的 resnet.py、wrn.py、mobilenetv2.py 都遵守类似约定新增 ShuffleNet 时只要在 return 前把最后一个 block 的 feature map append 到 list就能无缝接进 CRD。改完先用一个小 batch 打印特征维度确认能通过所有蒸馏器接口。5.3 验证蒸馏效果最稳定的一个技巧先把 240 epochs 缩到 30 epochs 做冒烟测试python train_student.py --path_t teacher --distill crd --model resnet8 --dataset cifar100 --epochs 30 --trial smoke如果 30 epoch 的曲线里 CRD 已经明显高于 vanilla 基线说明配置正确如果和 vanilla 持平再检查负样本采样是否把正样本对排除干净以及 criterion.py 里是否错误地从 memory bank 中取到了当前样本。把冒烟测试固化成一条短命令比每次盯着完整训练更有效。本文还有配套的精品资源点击获取