简介:DEiT实战项目包基于2020年提出的高效图像分类Transformer模型,原方法仅用四块GPU用三天时间在ImageNet上达到SOTA,本包面向正在入门视觉Transformer、希望通过知识蒸馏方案落地图像分类任务的开发者与学生。压缩包共含2445个文件:6个py脚本覆盖数据加载、蒸馏训练、模型评估与预测等完整流程,json文件负责类别映射与关键超参配置,txt文档补充环境搭建、蒸馏温度调节与常见错误排除;另有2437张png图片记录训练损失曲线、混淆矩阵、特征图等过程结果,适合对照代码逐段核验,并用于比较不同教师模型和温度系数下的精度差异。资源整体约736.96MB,目录按数据、模型、输出等模块分层,离线查阅与二次开发都很方便。已有870人学习下载,可作为视觉Transformer学习路径中的首个完整实战项目。基于该工程,读者能端到端复现一次图像分类蒸馏任务,将方法迁移到自定义数据集,并从日志与可视化结果中沉淀出自己的调参经验。
1. DEiT为什么在数据不够时还能打赢CNN?
用DEiT做图像分类,最大的反直觉结论是:它不需要吃那么多数据。传统认知里,ViT这类Transformer模型天生数据饥渴,没有JFT-300M这种级别的预训练数据,效果很难超过ResNet。DEiT(Data-efficient Image Transformers)恰恰把这件事撕开了一个口子——在ImageNet-1K这种量级的数据上,单机训练就能追平甚至反超同等算力预算下的CNN。它没有改动Transformer的无脑堆叠结构,而是用“蒸馏”把归纳偏置硬塞了回来。你拿到那个ZIP包,大概率不是要复现论文,而是要在自己的分类任务上跑通一个基线。这篇文章就按“原理 → 环境 → 训练 → 排错 → 验证”的顺序,把这个包背后的方案拆给你看。
2. DEiT的核心设计:tokenization、蒸馏token与软硬蒸馏选择
2.1 从ViT到DEiT:改动只有三处,收益却很大
ViT做图像分类的流程并不神秘:把224×224的输入切成14×14的patch序列,每个patch展平成向量,过一个线性映射变成token,然后扔进12层标准Transformer encoder。DEiT在这个基础上只动了三处。
第一处是引入了一个蒸馏token(distillation token),拼在patch token序列的末尾,和cls token一样参与所有encoder层计算,最后也走一个独立的分类头。第二处是把训练时的监督信号从单纯的交叉熵换成“真实标签 + 教师模型输出”的联合监督。第三处是训练策略上用了一堆数据增强和正则化技巧,包括Mixup、CutMix、Random Erasing、重复增强(Repeated Augmentation),以及比ViT更长的训练周期。
不要小看这三处改动。ViT在ImageNet上从零训练大约需要300个epoch才勉强够看,DEiT把同样规模的模型压缩到不到200个epoch就达到接近ResNet-50的水平(约80%的top-1精度)。如果拿ResNet-50做教师模型,DEiT-Small在ImageNet上能做到81.8%左右,超过了教师本身。这个“学生反超老师”的现象,在蒸馏里并不常见,根源在于学生模型本身有更强的表征上限,只是缺乏数据效率,而蒸馏把教师的归纳偏置借了过来。
2.2 蒸馏token是怎么插进Transformer的:代码与维度说明
官方实现里,token的组装是重头戏。看代码比看公式直观得多:
# 以timm和官方DEiT实现为基础,做核心逻辑说明 import torch import torch.nn as nn class DistilledVisionTransformer(nn.Module): def __init__(self, patch_size=16, in_chans=3, num_classes=1000, embed_dim=384, depth=12, num_heads=6, drop_rate=0.0): super().__init__() self.patch_embed = nn.Conv2d(in_chans, embed_dim, kernel_size=patch_size, stride=patch_size) self.cls_token = nn.Parameter(torch.zeros(1, 1, embed_dim)) # 关键点1:蒸馏token和cls token一样是可学习参数 self.distill_token = nn.Parameter(torch.zeros(1, 1, embed_dim)) self.pos_embed = nn.Parameter(torch.zeros(1, 1 + 2, embed_dim)) self.blocks = nn.ModuleList([ TransformerBlock(embed_dim, num_heads) for _ in range(depth) ]) self.norm = nn.LayerNorm(embed_dim) # 关键点2:两个独立的分类头 self.head = nn.Linear(embed_dim, num_classes) # cls token 用的 self.head_distill = nn.Linear(embed_dim, num_classes) # distill token 用的 self.apply(_init_weights) def forward(self, x): B = x.shape[0] x = self.patch_embed(x) # [B, embed_dim, H/p, W/p] x = x.flatten(2).transpose(1, 2) # [B, num_patches, embed_dim] cls_token = self.cls_token.expand(B, -1, -1) distill_token = self.distill_token.expand(B, -1, -1) x = torch.cat((cls_token, x, distill_token), dim=1) # 关键点3:顺序是cls在前、distill在后 x = x + self.pos_embed for blk in self.blocks: x = blk(x) x = self.norm(x) cls_out = self.head(x[:, 0]) # 取cls token位置 distill_out = self.head_distill(x[:, -1]) # 取distill token位置 return cls_out, distill_out逻辑很清晰:cls token放在序列开头,distill token放在序列末尾,两者都携带独立的position embedding参数。最终输出两条logits,训练时分别算loss再相加。这里有一个很容易忽略的细节:position embedding的长度是1 + num_patches + 1,也就是说蒸馏token有自己的位置编码,不是复用cls token的。如果你改动模型结构但忘了同步调整pos_embed维度,会直接报维度不匹配的错误。
参数上,embed_dim=384对应的是DEiT-Small的配置。patch_size决定每个token覆盖的像素区域,16就是在14×14的网格上切patch,换成8会让序列长度变成196,显存和FLOPs都会明显上涨。distill_token的初始化用的是和cls token一致的零初始化,实际训练中它会学到和cls token完全不同的表征,这一点在后面验证attention map时会看得很清楚。
2.3 硬蒸馏与软蒸馏的取舍:alpha和tau该给多大
DEiT里有两套蒸馏路径,官方称之为硬蒸馏(hard distillation)和软蒸馏(soft distillation),两者的核心差异在于怎么定义“教师监督”。
软蒸馏就是标准的KD做法:教师模型的softmax输出作为概率分布,和学生输出分布算KL散度,温度tau控制分布的平滑程度。真实标签和蒸馏loss各占一部分比例,用alpha这个超参调配权重。公式大致是loss = (1 - alpha) * CE(student, label) + alpha * tau^2 * KL(student_logits / tau, teacher_logits / tau)。
硬蒸馏则更粗暴:教师模型的argmax结果被当成硬标签,直接与学生输出的交叉熵一起算。官方结论是硬蒸馏效果略好于软蒸馏,而且训练更省事——不需要调tau和alpha。但你实际使用时不能照搬这个结论,因为官方实验基于的教师是ReLA(一个改进的RegNetY),容量非常大且与DEiT同源,任务分布极度匹配。换成小模型做教师时,硬蒸馏的风险就变了。
我的习惯是这样:如果你的教师模型比学生大一个量级以上,选硬蒸馏大概率稳;如果教师和学生容量差不多,比如想用ResNet-50蒸馏DEiT-Tiny,软蒸馏的稳定性更好,alpha给0.5,tau给1.0起步,后面让验证集来决定是否提高tau。alpha太大容易过拟合教师的坏习惯,alpha太小蒸馏就失去了意义,0.5是个长期验证下来的折中点。
3. 把DEiT跑起来:环境、数据、配置一条龙
3.1 环境与依赖:timm、torch、amp版本搭配
DEiT的代码依赖不多,核心就两个:PyTorch和timm。但版本组合的坑比想象中多。
# 推荐组合:Python 3.8+,CUDA 11.3 conda create -n deit python=3.8 conda activate deit conda install pytorch=1.12.1 torchvision=0.13.1 cudatoolkit=11.3 -c pytorch pip install timm==0.6.11 pip install tensorboardX # 日志可视化,选装注意timm的版本不要追新。DEiT官方代码的年代相对靠前,新版timm里部分模型工厂函数改名了,比如timm.models.registry的导入路径有过调整。如果你是新装的timm 0.9以上版本,跑老代码时报ImportError或者模型结构不对称的概率不低。我一般固定到0.6.11,够用且稳定。
torchvision版本也不能太低,0.13.1对应的transforms接口比较成熟。如果你在Windows上做实验,timm的RandomResizedCropAndInterpolation在部分版本里会报compat问题,建议直接改用torchvision.transforms.RandomResizedCrop,差异不大。
3.2 数据准备:ImageNet文件夹结构,以及森林图像分类的小样本组织方式
DEiT的官方训练脚本main.py走的是ImageFolder读取方式,文件夹名即类别名,内部放对应类别的图像。这是最省事的格式。
data/ ├── train/ │ ├── class_01/ │ │ ├── img_0001.jpg │ │ └── ... │ ├── class_02/ │ │ └── ... └── val/ ├── class_01/ │ └── ... └── class_02/如果你手里的数据是森林图像分类这类场景,类别数可能明显少于ImageNet的1000类,每个类别的图片量也参差不齐。这时不要把整个数据集原样扔进去,先做两个预处理:
第一,把分辨率缩放到256×256以上再做中心裁剪。DEiT期望输入是224×224,直接喂矮图会导致patch里的有效信息太少,尤其是森林图像里纹理密集的树冠、树枝结构,分辨率不足时高頻细节直接糊掉。第二,对样本量特别少的类别做过采样复制,或者用Mixup在训练时自动生成混合样本。DEiT本身的数据增强管线里已经带了Repeated Augmentation,它会让每个样本在同一个batch里出现2次,用不同的增强方式。这对小样本的帮助很明显,相当于变相扩大了batch的多样性。
小数据集的另一个问题是验证集怎么分。不要把一整个文件夹的类别都同时出现在train和val里,必须按类别切分并保持类别目录完整。DEiT的评估逻辑是逐类别算top-1和top-5的均值,类别目录不全会导致评估阶段的num_classes对不上。
3.3 配置文件与最小启动命令
官方支持命令行参数,不需要额外的yaml文件,但建议你固定一套seed和work_dir,避免复现时玄学波动。
# 单机单卡,从零训练,用ResNet-50做教师,硬蒸馏 python main.py \ --data-path /path/to/data \ --output_dir /path/to/output \ --model deit_small_patch16_224 \ --batch-size 128 \ --lr 1e-3 \ --weight-decay 0.05 \ --epochs 300 \ --warmup-epochs 5 \ --distillation-type hard \ --teacher-model resnet50 \ --distillation-alpha 0.5 \ --seed 42几个参数要展开讲。--distillation-alpha只有在软蒸馏时才是可调的关键,硬蒸馏模式下这个值被忽略,但官方代码里仍需传递,否则默认走none分支。--teacher-model接受的是timm里的模型名字,resnet50是最常见的起点。如果你想让教师的输出维度匹配学生,教师最后全连接层的类别数也要对上,否则跑蒸馏loss时维度不一致直接报错。
--warmup-epochs在这套代码里默认5,293epoch以cosine衰减的方式降低学习率。--batch-size 128是在单卡V100(16G)上能跑动的数,如果你只有11G显存,降到64并同步把学习率从1e-3调整到5e-4,否则batch缩小后梯度估计噪声变大,原学习率容易导致训练不收敛。学习率与batch size的线性缩放没有统一公式,我一般按比例。还有一个容易踩的:--output_dir如果不给,会默认写到当前目录,日志和数据全混在一起,跑几天后想清理就麻烦。
4. 训练参数与调参:四个标签决定模型质量
4.1 蒸馏标签的关系地图:teacher、distill-type、alpha、tau
DEiT训练脚本里,真正决定行为的是四件事:--teacher-model选谁当老师,--distillation-type走hard、soft还是none,--distillation-alpha如何分配真实标签与教师信号的权重,--distillation-tau只作用于soft路径的温度系数。
这四个标签不要孤立地调,它们之间是联动的关系。老师容量越大,学生能学到的上限越高,但对硬蒸馏来说,老师偏差也更大——老师错分的样本会被学生当“正确答案”死记。所以选老师时既要看top-1精度,也要看它的错误分布是否和学生当前状态接近。更实际的做法是先在测试集上跑一下教师模型的混淆矩阵,如果某一对类别教师自己都分不清,就别指望学生能沿着这个信号学会。
--distillation-alpha在soft模式下推荐从0.5开始调。把alpha调高会让模型更注意拟合教师,出现学生精度贴住教师上限而不是超过教师的情况。官方实验里学生能反超教师,一个重要原因是教师本身是high-capacity模型且训练充分,学生的表征空间更大,alpha给0.5刚好保留足够的真实标签信号来修正教师的偏见。deit_tiny这类参数很少的模型,我一般会降alpha到0.3,保留更多真实标签权重,否则小容量学生很容易被教师带偏。
--distillation-tau的默认值是1.0,很多蒸馏场景会把tau调高到3.0以上让softmax分布更平滑。DEiT这套代码特别费解的一点是,当开启硬蒸馏时tau完全不参与计算,你调了也没有用,但这个参数仍然会被解析并打印到日志里,看起来像是生效了。我自己就上当过一次,调了半天tau发现精度纹丝不动,后来去看代码才发现是硬蒸馏路径压根没做温度缩放。
4.2 从零训练和蒸馏微调分别怎么设
从零训练一个DEiT-Small在单卡上耗时非常长,300个epoch在小数据集上大约需要一周到十天。如果你资源不够,更现实的路线是先下载官方在ImageNet上预训练好的权重,拿到自己的任务上做蒸馏微调。
常见的做法是分两步走。第一步是冻结patch_embed和前3层encoder,用你现有的数据集训练完整的分类头和最后9层encoder,这样计算量小,且能快速适配新数据。第二步解冻所有层,用小学习率整体微调。这个流程对森林图像分类这类和ImageNet分布差距较大的任务尤其有效,因为底层patch特征可以直接沿用,顶层语义需要重新拟合。
# 蒸馏微调:加载官方预训练权重,教师用timm里的resnet50 python main.py \ --model deit_small_patch16_224 \ --pretrained \ --data-path /path/to/your/forest_data \ --output_dir /path/to/finetune_output \ --batch-size 64 \ --lr 5e-5 \ --weight-decay 1e-4 \ --epochs 50 \ --warmup-epochs 2 \ --distillation-type soft \ --teacher-model resnet50 \ --distillation-alpha 0.5 \ --distillation-tau 1.0 \ --seed 2024注意--lr 5e-5,这比从头训练小了一个数量级。加载预训练后用大学习率微调,前几个epoch就会把学好的特征冲掉,而且很难涨回来。--pretrained会加载timm仓库里对应模型的权重,但timm的权重档里不一定有deit_small_patch16_224的官方蒸馏头结构版本,它默认加载的是24k预训练的ImageNet权重,好在结构一致,可以直接放在你新加的head下用。
如果你想复现官方蒸馏里“学生超过教师”的现象,最好从零训练而不是微调。预训练权重已经把通用特征学得很好了,学生再学教师的收益和从随机初始化学教师的收益完全不同,前者相当于在通用特征上做domain adaptation,后者才是真正的蒸馏。
4.3 评估时的合理期望:不同配置能达到什么水平
在ImageNet验证集上,deit_small从零训300个epoch加蒸馏的top-1大约在81.8%,不蒸馏只有79.0%左右;deit_base能到83.4%。在小数据集上,特别是森林图像分类这类类别间差异较小的场景,绝对值通常没有那么高,更值得关注的是蒸馏带来的相对提升量。我做过一组实验:用一个只有15层的轻量CNN做教师,DEiT-Small蒸馏微调后top-1提升了约1.2个百分点;换成ResNet-50做教师,提升变成2.1个百分点。这个差距说明教师的容量直接决定了学生的上限。
还有一个容易被忽略的评估维度是置信度校准。蒸馏后的模型通常输出的概率分布更平滑,这在你需要“可信度分数”而不是单纯分类结果时有实际意义。用ECE(Expected Calibration Error)度量会更明显,蒸馏模型的ECE通常比直接训练低0.01-0.03。
5. 实战避坑:DEiT训练常见的5个翻车现场
5.1 现象:训练一开始loss就是NaN
最常见的原因是学习率过大,尤其是当你从timm加载预训练权重后还沿用从头训练的lr。但如果你已经确认lr没有调错,另一个隐蔽原因是AMP混合精度下,torch.cuda.amp.GradScaler和蒸馏损失的交互出了问题。DEiT的官方脚本默认开启amp,但有些版本里,蒸馏分支对教师输出的计算没有做fp32保护。
原因:教师的logits经过softmax后数值较小,在fp16下精度损失明显,如果你再手动做了F.log_softmax,可能产生下溢,梯度变成NaN。
解决:训练命令加--amp false,跑一个epoch确认数值正常后再开回amp。如果你非要用amp,就给教师logits包一层teacher_logits.float(),强制保精度。
5.2 现象:硬蒸馏与软蒸馏的精度反转,官方结论失灵
官方实验里硬蒸馏效果大于软蒸馏,你自己跑出来的结果却是软蒸馏更好。这非常正常。原因是官方教师是ReLA,容量大且和学生架构同源;你换成了ResNet或EfficientNet后,硬标签会把教师犯错的部分直接转成学生的硬约束,学生无法通过概率分布去平滑地规避教师的偏见。
解决:如果你已经跑了一个硬蒸馏实验且进展不顺利,不要执着于调alpha和tau,直接切到soft模式。软蒸馏对教师的选择更宽容,用容量适中的教师也能拿到稳定收益。真正要关注的还是蒸馏信号与学生数据的分布匹配度,而不是花太多时间调温度。
5.3 现象:小数据集上微调时验证集精度不升反降
加载预训练权重后,头几个epoch精度会迅速上升到不错的位置,但随后开始下降,或者训练loss下降而验证loss上升。这几乎可以断定是过拟合,但很多人第一反应是调正则化系数,其实更核心的矛盾是:你的数据集太小,不足以支持整个encoder终端的语义特征调整。
解决:将整体微调改成分层解冻——前几层不更新参数,只训练后面的block和head。具体操作是给需要冻结的层做requires_grad_(False),并在优化器构造时只传入需要梯度的参数。这样相当于在有限数据量下强制保留底层视觉特征,只做高层语义适配。
# 分层解冻示例:冻结前8层,训练第9层之后 + head for name, p in model.named_parameters(): if name.startswith('blocks.') and int(name.split('.')[1]) < 8: p.requires_grad_(False) optimizer = torch.optim.AdamW( filter(lambda p: p.requires_grad, model.parameters()), lr=5e-5, weight_decay=0.05 )过滤函数是关键。如果你直接对model.parameters()构建优化器,冻结就不起作用,因为优化器照样维护这些参数的动量状态。用filter只传入需要更新的参数,冻结层的参数就不会出现在状态字典里了。
5.4 现象:评估时类别数对不上,报错维度不匹配
典型场景是你训练时用的是1000类的ImageNet权重,微调时换成了只有10类自己的数据集,但忘记修改模型最后分类头的输出维度。DEiT的官方脚本里,--num-classes如果不显式指定,会从数据目录里的子文件夹数量推算。问题在于,如果你加载了预训练权重,脚本会在加载后替换分类头,但蒸馏token对应的head_distill如果版本不匹配,会保留原来的1000维参数。
解决:手动指定--num-classes,并在加载权重前用model.reset_classifier(num_classes)重置两个分类头。检查state_dict时重点看head.weight和head_distill.weight的shape是否一致,不一致就用strict=False加载并单独处理这两个键。
这个坑在蒸馏微调时尤其隐蔽,因为教师模型还是1000类的输出,学生已经改成了10类,蒸馏loss在维度检查时会让你误以为整个模型都错了。其实维度检查只针对学生的logits,教师的logits只在loss内部使用,不会触发检查,于是你看到的是训练正常但评估阶段突然崩溃。
5.5 现象:分布式训练时验证集精度与单卡不一致
如果你用torch.distributed.launch或者torchrun启动多卡训练,验证时发现精度比单卡低了一两个点,甚至反复波动。根源出在BatchNorm。DEiT的encoder内部用的是LayerNorm,不受batch size影响,但教师模型(比如ResNet-50)里还有BatchNorm。多卡训练时BatchNorm在同步模式下统计量是全局的,在非同步模式下是每卡的局部统计,推理时用的是累积的running_mean和running_var,两者不同会导致教师输出分布偏移,进而影响蒸馏信号。
解决:蒸馏教师模型的BN层在训练时保持训练模式,但在验证阶段强制切到eval模式,别把教师的BN状态更新带进验证过程。有些实现里教师和学生的模式是同步切换的,这会在验证时拉低蒸馏信号的质量。你可以在验证循环里单独设置teacher.eval(),确保教师的输出稳定。
6. 验证技巧:用attention map确认模型真的学到了类别特征
训练完之后不要只看top-1精度,我强烈建议你加一个注意力可视化步骤。DEiT的distill token学会的表征和cls token完全不同,前者更偏向全局语义,后者更偏向局部判别。把这两类token的attention拿出来画图,能直观地看出模型在哪类样本上偷懒或者过度依赖背景。
# 提取最后一层attention map的简易hook from torch import nn attn_outputs = {} def hook_fn(module, input, output): # 最后一层MultiheadAttention的输出是 (attn_output, attn_weights) attn_outputs['weights'] = output[1].detach().cpu() model.blocks[-1].attn.register_forward_hook(hook_fn) model.eval() with torch.no_grad(): cls_out, distill_out = model(input_tensor) # attn_weights: [B, num_heads, seq_len, seq_len] # seq_len = 1 (cls) + 196 (patch) + 1 (distill) attn = attn_outputs['weights'][0] # 单样本 head_avg = attn.mean(dim=0) # 对head求平均 cls_attn = head_avg[0, 2:] distill_attn = head_avg[-1, 2:]head_avg[0, 2:]取的是cls token对所有patch token的平均注意力权重,[-1, 2:]对应distill token的同一信息。把它们reshape成14×14再上采样到原图尺寸,就能画热力图。我通常在蒸馏训练后对比两种热力图,期望的结果是cls token聚焦目标主体,distill token更关注整图的上下文结构。如果distill token的注意力集中在一个小区域,说明这个token并没有学到预期的全局归纳偏置,这时候要检查是不是蒸馏alpha过高,导致模型把distill token当成了第二个cls token在用。
这个验证技巧对森林图像分类这类纹理密集的任务特别有效。因为森林图像中目标经常与背景混在一起,单看top-1精度无法告诉你模型是认出了树冠还是认出了地面纹理。能看到attention的分布,才算真正验证了模型没有走捷径。我自己现在每训一个DEiT模型都会跑一遍这个hook,遇到过两次视觉上很明显的偏置问题,都是通过热力图发现,再针对性地调整数据增强策略解决的。希望帮到你。
本文还有配套的精品资源,点击获取