
简介这是一份基于DEiTData-efficient Image Transformers的图像分类实战资源包主要面向希望系统掌握Transformer蒸馏训练技巧的深度学习开发者与研究人员。资源围绕DEiT模型展开涵盖数据组织、训练脚本、标签映射与类别配置等关键内容既适合复现论文效果也适合迁移到自己的分类任务中。压缩包共2445个文件其中以2437张png格式的训练与验证图像为主配合6个Python脚本、1个json配置文件和1个txt说明文件整体体积约736.96MB目录结构清晰便于按模块调用和二次开发。该资源已有869人学习浏览说明其在实际使用中具备一定参考价值尤其能为入门Transformer的读者提供直观的工程样例。通过这份资源读者可以获得完整的DEiT图像分类项目框架理解蒸馏训练中教师网络与学生网络的配合方式并可直接基于现有脚本开展实验从数据准备到模型评估形成闭环节省从零搭建模型的时间。1. DEiT实战一个zip包带你把图像分类从CNN换到Transformer第一次在ImageNet上把DEiT训到和ResNet-50相当的精度我的第一反应是“早该换过来了”。DEiT全称Data-efficient Image Transformers本质是一个带蒸馏令牌的ViT变体专门解决ViT在中小数据集上“训练不动”的问题。这个zip包给你的是模型定义、蒸馏训练脚本和推理样例目标就是用DEiT跑通图像分类任务。适合手里只有几万张标注图的团队也适合想复现Transformer图像分类的入门工程师。训练过的人都知道普通ViT在小数据集上特别容易过拟合而DEiT用一整套蒸馏策略把这个坑填上了。下文按机制、训练、推理、避坑四个方向展开配置和代码都能直接抄读完你能自己改数据集、调参数、跑通整个流程。2. DEiT核心机制蒸馏令牌、软硬蒸馏和那组决定成败的训练超参2.1 蒸馏令牌一个“列向量”如何传递教师信号先把概念说透。ViT把一张图切成16x16的patch展平成序列送进Transformer编码器。为了让输出具备分类能力输入序列里会额外拼一个class token也就是一个可学习的向量。DEiT在此基础上又加了一个distillation token位置通常放在patch序列的末尾和class token一样参与所有block的自注意力计算。两个token的区别在输出端class token接一个全连接层作为分类头distillation token接另一个全连接层两个头互不干扰。为什么不让蒸馏信号直接叠加到class token上因为实验发现让两者各自经过独立的head注意力分布差异更大互补性更强。蒸馏token只负责“从教师那吸收知识”class token只负责“把真实标签学好”分工明确。教师的能力是怎么传过去的你在训练时把distillation token和patch序列一起过自注意力每层的token都会和patch交互教师模型的输出会被当作额外监督信号强行牵引distillation分支的特征。数据越少这个牵引越关键。作者在论文里验证过当训练集缩到ImageNet的1%时加不加蒸馏token精度差距能到几个百分点。2.2 三种蒸馏模式怎么选hard、soft、CNN蒸馏各吃多少数据DEiT论文里给了三种蒸馏路径按落地稳妥程度排序如下。hard蒸馏教师对每张图输出argmax标签把这个硬标签当成额外监督和真实标签一起算交叉熵。代码最简洁不用保存教师logits显存压力小是我默认的配置。复现DEiT只是为了感受效果选hard就够了。soft蒸馏用KL散度让学生的class token输出逼近教师的完整概率分布。信息量比hard大但对教师logits质量很敏感。教师置信度畸高时学生会把错误信念也学进注意力层。一般要配合温度缩放温度T取2到4我在森林图像分类这种类别高度重叠的任务上试过soft比hard高不到1个点但调温过程很费时间。CNN蒸馏DEiT作者实验发现教师用CNN比用Transformer效果更好因为CNN的归纳偏置和Transformer的建模方式互补蒸馏得到的注意力图更集中。官方选的是RegNetY-16GF但你不用完全复刻手头任何精度显著高于学生的CNN都能当教师。我自己用ResNet-50当教师效果也很稳。如何选数据在十万张以下、类别数不多默认hard加CNN教师数据量上到五十万再考虑soft蒸馏。教师选择有一个底线教师精度必须明显高于学生基线否则蒸馏就是练废。选之前先在验证集上测一下教师精度别让玄学代替评估。2.3 DEiT训练超参表照着抄能少走一半弯路DEiT不是靠某一个技巧赢的它把所有训练细节都拉满。我把和精度强相关的参数整理成一张表后面训练脚本会用到这些配置。参数取值说明输入分辨率224x224DEiT系列默认尺寸epochs300小数据集可缩到150但要配合更强增强batch size1024个人用户用梯度累计模拟这个量级optimizerAdamWlr5e-4weight_decay0.05学习率策略cosine decay前5个epoch做warmup从1e-6涨到目标lr数据增强RandAugment强度9概率0.5Mixup / CutMix0.8 / 1.0两者同时开启EMA0.99996推理时加载EMA权重而不是最后一步权重repeated augmentation开启同一张图在多个epoch被重复采样这条参数表里有几个容易翻车的细节。weight_decay是0.05不是ViT常用在Large模型上的0.3小模型用0.3会直接loss不降。RandAugment在timm里通过re_prob参数控制和torchvision的randaugment实现不是一套直接从torchvision搬过去会吃亏。EMA要全程维护一份影子权重最后加载的是影子权重这一步漏掉会让验证精度平白掉1到2个点。3. 用DEiT训练图像分类模型数据目录与最小训练脚本3.1 数据准备把jpg按类别放进train和val目录DEiT训练要求一个标准目录结构timm的create_dataset能自动读取不需要手写Dataset类。目录结构是硬约束建议手动创建别用符号链接绕路。# 建目录把图片按类别放好 mkdir -p dataset/train dataset/val # 以森林图像分类为例train下每个类别一个目录 # dataset/train/rainforest/ # dataset/train/desert/ # dataset/train/grassland/ # dataset/val/rainforest/ # dataset/val/desert/ # dataset/val/grassland/ # train和val下类别名必须完全一致逻辑说明timm的folder数据集把子目录名当作类别名类别顺序按字母序自动编码。train和val如果类别集合不一致验证时label数量对不上会直接报错这是新手最容易踩的坑之一。参数说明类别数尽量不要超过10000DEiT的head是线性层类别过多会带来显存和收敛速度双重压力。数据量不足时先保证每个类别至少有50张训练图验证集每类20张以上否则后面的评估结果没有统计意义。3.2 最小训练脚本一个带教师蒸馏的完整loop下面这个脚本是我整理过的最小可跑版本去掉了分布式、EMA、断点续训这些复杂逻辑只保留核心蒸馏流程。你只要把数据目录和类别数改成自己的就能跑。import torch import torch.nn as nn import timm from timm.data import create_dataset, create_loader, resolve_data_config from timm.optim import create_optimizer_v2 from timm.scheduler import create_scheduler # 教师模型固定权重学生用DEiT-Tiny teacher timm.create_model(resnet50, pretrainedTrue, num_classes2) teacher.set_grad_enabled(False) student timm.create_model(deit_tiny_patch16_224, pretrainedTrue, num_classes2, drop_rate0.0) config resolve_data_config({crop_pct: 0.875}, modelstudent) train_loader create_loader( create_dataset(folder, rootdataset/train), input_size(3, 224, 224), batch_size64, is_trainingTrue, re_prob0.25, mixup_alpha0.8, cutmix_alpha1.0, num_workers4, pin_memoryTrue) optimizer create_optimizer_v2(student, optadamw, lr1e-4, weight_decay0.05) scheduler create_scheduler(optimizer, schedcosine, epochs150, warmup_epochs5, warmup_lr1e-6) ce_loss nn.CrossEntropyLoss() for epoch in range(150): student.train() for images, targets in train_loader: images, targets images.cuda(), targets.cuda() with torch.no_grad(): tea_logits teacher(images) tea_labels tea_logits.argmax(dim1) # hard蒸馏标签 cls_feat, dist_feat student.forward_features(images) out_cls student.head(cls_feat) out_dist student.head_dist(dist_feat) loss 0.5 * ce_loss(out_cls, targets) \ 0.5 * ce_loss(out_dist, tea_labels) optimizer.zero_grad() loss.backward() optimizer.step() scheduler.step(epoch) torch.save(student.state_dict(), fdeit_tiny_epoch{epoch}.pth)逻辑说明这里用forward_features拿到两个token的嵌入向量分别过student的head和head_dist再各自算交叉熵最后加权合并。和直接用model(images)拿单一输出相比这样才真正实现了DEiT的蒸馏训练。head和head_dist是timm源码里针对DEiT定义的两个独立分类层这结构就是蒸馏令牌发挥作用的关键。参数说明resnet50的num_classes改成了2实际任务中改成你的类别数。learning rate我建议用1e-4而不是论文里的5e-4原因是单卡batch size只有64论文是1024批大小差着数量级lr也要跟着降否则梯度噪声太大。mixup和cutmix同时开启这是DEiT训练稳定的重要因素。如果你显存不够batch size可以再降到32但lr要同步降到5e-5左右。3.3 训练命令与日志解读看这三条曲线# 单卡训练 python train_deit.py # 多卡训练需要脚本里先补init_process_group相关逻辑 # torchrun --nproc_per_node4 train_deit.py逻辑说明脚本本身没有分布式初始化代码所以torchrun那行现在跑不了多卡要自己加init_process_group样板。单卡训练时建议先用小类别数跑2个epoch验证整体流程别一上来就用全量数据省得数据问题在训练中段才暴露。看日志时别只盯的accuracy一个指标。第一看train loss在头5到10轮有没有明显下降第二看验证错误率在50轮内是否开始收敛第三看收敛后的精度有没有超过直接用ResNet-50当教师的上限。三个指标里有两个不对直接停掉去查增强配置、lr和wd别硬跑300轮。我第一回跑的时候torchvision的preprocess和timm的normalize参数不一致导致推理阶段预测结果全是垃圾这种错logs的准确率是看不出来的要单独验证数据预处理路径。4. 利用训练好的DEiT做图像分类推理权重加载、预处理与结果可视化4.1 推理脚本加载权重和图像预处理一次到位训练完的pth文件不是直接load就完事下面这段脚本处理了三个容易出错的点timm权重对应的预处理、两个分类头的加载、模型推理模式。import torch import torch.nn.functional as F import timm from torchvision import transforms from PIL import Image classes [rainforest, desert, grassland] transform transforms.Compose([ transforms.Resize(256), transforms.CenterCrop(224), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]) ]) model timm.create_model(deit_tiny_patch16_224, pretrainedFalse, num_classeslen(classes)) model.load_state_dict(torch.load(best.pth, map_locationcpu), strictFalse) model.eval().cuda() img transform(Image.open(test.jpg).convert(RGB)).unsqueeze(0).cuda() with torch.no_grad(): cls_feat, dist_feat model.forward_features(img) out_cls model.head(cls_feat) out_dist model.head_dist(dist_feat) out (out_cls out_dist) / 2 prob F.softmax(out, dim1)[0]逻辑说明用strictFalse是因为你把head和head_dist的维度从1000改成了自己的类别数原权重里这两个key对不上不用strictFalse连加载都会报错更别提推理了。forward_features拿到两个token的嵌入后必须分别过对应的分类头再融合这个融合操作的合理性来自蒸馏令牌的设计初衷。参数说明Resize到256再CenterCrop到224是官方做法crop_pct约等于0.875。如果你的训练阶段的预处理不是这一套推理时也要完全保持一致任何不一致都会带来1到2个点的精度损失。两个头相加后取平均会让置信度分布比单头更平滑这是DEiT在推理阶段的实际经验不用纠结数学依据直接用就行。4.2 输出解析与Top-5置信度用softmax看原始logits模型输出的out_cls和out_dist都是未经softmax的logits。直接比大小会踩坑因为两个头的数值范围并不同必须先过softmax再做比较。top5 prob.topk(min(5, len(classes))) print(预测类别:, classes[top5.indices[0].item()]) print(置信度:, top5.values[0].item())逻辑说明topk返回最大的k个值和对应索引。代码里min(5, len(classes))防止类别数小于5时越界。prob是一维张量所以用[0]取第一个样本的结果。参数说明类别数大于10时Top-5才有参考价值类别只有2到3类时不如直接看置信度。如果多张图一起推理注意你用的是batch1的输入topk要做dim维度判断别直接索引[0]。4.3 用混淆矩阵评估自己的类别只看accuracy会骗自己accuracy只在类别均衡时才有意义。森林图像分类这种任务里模型很容易把草地误判成森林因为整张图的绿色比例接近但这种系统性错误用单一accuracy根本看不出来。from sklearn.metrics import confusion_matrix, classification_report import glob y_true, y_pred [], [] for path in glob.glob(test/*/*.jpg): true_label path.split(/)[-2] y_true.append(classes.index(true_label)) img transform(Image.open(path).convert(RGB)).unsqueeze(0).cuda() with torch.no_grad(): out model(img) y_pred.append(out.argmax(dim1).item()) print(classification_report(y_true, y_pred, target_namesclasses)) print(confusion_matrix(y_true, y_pred))逻辑说明用glob遍历测试目录路径倒数第二级是类别名把它转成类别索引。收集真值和预测值后用sklearn直接输出精确率、召回率和混淆矩阵。从矩阵里你能看到具体是哪个类别被混淆比看一个浮点数的accuracy有用得多。5. 避坑指南从zip解压到训练Loss不降的五个高频问题这章集中写我实战里反复撞上的问题每一条都是血泪经验按现象到原因到解决展开。5.1 zip压缩包解压报CRC错误先查伪加密和压缩工具现象下载的DEiT zip包在Windows上右键解压中途报“CRC校验失败”部分文件能解出来但图片已经在传输中损坏。Linux下用unzip解压同样报错。原因常见有两种。第一种是网络传输丢包压缩包本身已经损坏第二种是“zip伪加密”——压缩包的加密头标志被第三方打包工具改动过系统自带解压器识别出加密状态又没有密钥于是拒绝解压或解出垃圾文件。解决先重新下载一次再解压排除传输问题。如果还是报错换7-Zip解压7-Zip对ZipCrypto伪加密有兼容处理能无视那个错误的加密标志直接解出来。我遇到过一次最隐蔽的解压后图片能打开但训练精度一直偏低逐张检查才发现压缩包里混了几张半损坏图片被静默解出。所以解压后务必先用命令统计图片数量和size再抽样算md5别急着进训练流程。# Linux下用unzip解压后统计每个类别的图片数量 unzip deit_image_classification.zip -d deit_project/ find deit_project -name *.jpg | wc -l逻辑说明unzip是Linux下最常用的zip解压命令加-d指定输出目录。find统计jpg数量后和压缩包内清单对比数量不一致说明有文件解压失败。这张检查表能挡住绝大多数静默损坏。5.2 加载权重时size mismatchdistillation token带出两个head现象用timm的deit_tiny_patch16_224加载官方预训练权重后把num_classes改成自己的类别数加载时报KeyError或size mismatch。原因DEiT的state_dict里除了patch_embed、blocks还有head和head_dist两个分类头。改成2类之后这两个头的权重尺寸从1000被重置为2原权重对不上。解决把两个head的权重直接从预训练权重里剔除只保留特征提取部分。from collections import OrderedDict state torch.load(deit_tiny_patch16_224.pth, map_locationcpu) new_state OrderedDict() for key, value in state.items(): if key.startswith(head): continue new_state[key] value model.load_state_dict(new_state, strictFalse)逻辑说明跳过head和head_dist让两个分类头保留随机初始化状态。patch embedding和transformer block的权重是通用的视觉特征提取能力保留下来可以在自己的数据集上快速收敛只训两个head的话训练轮数可以比其他方案减少一半。5.3 训练loss不下降别急着调lr先看增强配置现象DEiT训练了30轮训练loss还在2.5附近震荡验证accuracy一直没有起色怎么调learning rate都没用。原因最常见的是weight_decay设成了0.3小模型被正则过头。另一个更隐蔽的是mixup和cutmix同时开启时loss曲线初始值偏高混合标签的交叉熵天然就是比普通标签大不熟悉的人会误以为模型没在学。解决先关掉mixup和cutmix只保留RandAugment跑5轮观察loss能不能降到接近基线水平。能降再逐步把mixup加回来。开mixup后loss比普通交叉熵高0.3到0.5是正常现象别在这个状态下去调lr。同时把weight_decay改回0.05lr降到1e-4八成问题能解决。我之前在一版配置里把lr从1e-4改成5e-4单卡上Loss好几个epoch不降改回去立刻好转。有这时间你还不如去睡觉第二天再看曲线。5.4 GPU显存OOM1024的batch size不是给个人用户准备的现象照着论文把batch size设成1024单张消费级显卡直接OOM小显存卡连64都稳不住。原因DEiT-Tiny虽然参数量不大但224x224输入切patch后序列长度为196加2个token12层Transformer的激活值比同规模CNN大得多。显存瓶颈在激活不在参数。解决减小batch size并开启梯度累计用累计步数补足有效batch size。accum_steps 4 # 有效batch size 64 * 4 256 for step, (images, targets) in enumerate(train_loader): images, targets images.cuda(), targets.cuda() with torch.no_grad(): tea_logits teacher(images) tea_labels tea_logits.argmax(dim1) cls_feat, dist_feat student.forward_features(images) out_cls student.head(cls_feat) out_dist student.head_dist(dist_feat) loss 0.5 * (ce_loss(out_cls, targets) ce_loss(out_dist, tea_labels)) loss loss / accum_steps loss.backward() if (step 1) % accum_steps 0: optimizer.step() optimizer.zero_grad()逻辑说明把loss除以accum_steps再backward梯度累计多个batch后再update一次数学上等价于batch size乘以accum_steps。DEiT全用LayerNorm没有BatchNorm所以这条策略不会造成BN统计量偏差。四步累计大概能把显存压力降一半是单卡训练最实用的招。5.5 验证精度低或波动大EMA权重和验证集缓存没跟上现象训练结束后加载最后一步权重做验证精度比训练中期的checkpoint还低换两次验证集评估的结果能差两个百分点。原因DEiT官方训练全程维护EMA权重这个影子权重对训练末尾的噪声更钝化所以性能更好。如果你只保存最后一步step权重精度变低是必然。验证集波动大则多半是验证loader里shuffle顺序和缓存导致样本重复评估。解决给模型包一层EMA每个step更新一次最后保存并加载EMA权重。from timm.utils import ModelEmaV2 ema_model ModelEmaV2(student, decay0.99996) # 在训练循环每个step末尾调用 # ema_model.update(student) # 训练结束时保存ema_model的state_dict ema_state ema_model.state_dict() torch.save(ema_state, best_ema.pth)逻辑说明decay接近1意味着EMA更新非常慢所以每个step都要调用一次update不是每隔几个epoch才同步一次。推理时加载的是ema_state且EMA模型自己也要调成eval模式drop_path和dropout这时候必须关闭否则验证结果忽高忽低。6. 进阶把DEiT微调迁移到自己的小数据与ONNX导出6.1 从零训练换微调先冻结特征抽取器目标数据集只有几千张图时不建议从随机初始化开始。正确做法是加载ImageNet预训练权重冻结patch_embed和前半段block只微调后半段block和两个head学习率设为预训练时的十分之一也就是5e-5。直接用5e-4的lr去动预训练权重一步更新就能把特征冲乱。我的习惯是先用冻结策略跑10个epoch观察loss曲线。如果下降速率正常再解冻全部层做一次全量微调每个epoch保存一个checkpoint。这样做的成本比从零训练低一半精度却通常更高。数据类别和预训练类别差异大时比如森林图像分类和ImageNet本身语义接近这个策略效果更好如果你做的是医学切片这类领域差异很大的任务冻结的层数要相应减少。6.2 导出ONNX时动态维度和大小的处理模型要部署到生产环境导出ONNX是常见做法。DEiT导出有两个坑forward返回值是tupleONNX导出器对tuple输出支持不稳定两个分类头的位置必须都在模型定义里。我一般在导出前先用一个包装module把两个head的融合结果包成单一输出。import torch model.eval() x torch.randn(1, 3, 224, 224).cuda() torch.onnx.export(model, x, deit.onnx, input_names[images], output_names[logits], dynamic_axes{images: {0: batch}, logits: {0: batch}})逻辑说明dynamic_axes把batch维度标成动态导出后的模型在onnxruntime里任意batch都能推理。导出的前提是你已经用包装module把两个head的相加逻辑包进去了否则日志里会报tuple输出不支持的错。导出后拿onnxruntime加载对同一张输入图和PyTorch输出对比误差小于1e-4才算成功。如果误差过大回查normalize参数和推理时的dropout开关。把ONNX跑通后这个模型就能脱离PyTorch环境在CPU上部署了。DEiT这个方向我在几个分类任务里用过最大的教训是它的精度提升不靠单点创新而是蒸馏令牌、EMA、强增强叠加出来的系统工程别一上来就按自己理解删掉某个模块。训练稳定性的优先级永远高于峰值精度先把流程跑稳再慢慢调。希望帮到你。本文还有配套的精品资源点击获取