DEiT图像分类实战:数据高效Transformer的训练与推理
简介面向深度学习与计算机视觉学习者这份DEiT实战资源围绕Facebook提出的DeiT模型展示如何在不依赖外部数据集的情况下利用知识蒸馏策略完成ImageNet级别的高效训练并落地到图像分类任务中。DeiT通过引入蒸馏令牌与教师模型交互显著降低训练成本四块GPU三天即可达到SOTA水平适合希望快速上手Transformer图像分类、但缺乏超算环境的研究者与工程师。压缩包共2445个文件主体为2437张可视化图片涵盖训练曲线、混淆矩阵、样本预测等结果另有6个Python脚本负责模型搭建、数据加载、训练与评估流程1个JSON文件存储类别映射1个TXT文档补充说明整体约737MB目录结构清晰便于对照学习。目前已有869人学习该资源。下载后可获得从数据准备、脚本配置到模型训练、指标分析的一整套可参考实现借助过程图片能快速验证蒸馏效果迁移至自有分类数据集。1. 直接上手 DEiT别人三天训完 ImageNet我们十分钟跑通推理图像分类这个方向2020 年之前基本是 CNN 的天下Vision TransformerViT虽然效果好但训练极其挑剔动辄几十个 epoch 加上超大 batch普通实验室根本玩不转。DEiTData-efficient Image Transformers就是冲着这个痛点来的——Facebook 在 2020 年提出的一篇 Transformer 模型只靠 4 块 GPU 训了三天没碰任何外部数据就在 ImageNet 上打到了 SOTA 级别。它做到了“让 Transformer 在中小规模数据上也能训练”而不是像 ViT 那样必须靠 JFT-300M 这种庞大数据集撑腰。对大多数做图像分类的从业者来说DEiT 是一个比 ViT 更现实的起点显存压力小、收敛快、蒸馏机制可解释。这份资源包含了一个完整的 DEiT 图像分类实战项目核心构成是class.json类别标签映射文件外加 8 张测试图片配套原始博客里有完整的训练与推理代码。也就是说你拿到手不是一份只能看的教程而是一个能直接跑通的分类验证环境。它适合谁两类人一是刚接触 Transformer 图像分类、想用一个轻量级模型快速验证效果的同学二是已经在用 CNN 做分类、想对比 DEiT 和 ResNet 系列在同样数据上谁更稳的工程师。接下来我从模型机制讲到推理落地再把手把手把训练参数和坑都过一遍。2. DEiT 的核心机制蒸馏是这样把精度“教”出来的2.1 为什么 ViT 难训练DEiT 却能低资源收敛要理解 DEiT 的厉害之处得先明白 ViT 为什么难训。ViT 把图像切成 16x16 的 patch拉平后拼上位置编码送进标准 Transformer encoder。这个架构本身没有引入图像领域的归纳偏置比如 CNN 的局部连接和权值共享所以它需要海量数据来“自己摸索”出空间结构。当训练数据只有 ImageNet-1K 这种百万级规模时ViT 从小数据集上学到的特征泛化能力不如同等规模的 CNN表现甚至不如 ResNet。DEiT 针对这一点给出的答案很直接不让模型从零硬学而是让一个训练好的 CNNRegNetY 系列当老师用蒸馏损失把知识“压”进学生 Transformer。同时配合了一系列数据增强策略——RandAugment、MixUp、CutMix、EMA指数移动平均把训练难度降下来。整个训练流程使用 AdamW 优化器、cosine 学习率调度batch size 1024输入分辨率 224x224。这个配置组合在单机 4 卡 V100 上三天能跑完 300 epoch。这两张图你可以对照资源里的5e4d1ee0d.png和77291b3ad.png看左侧是训练 loss 曲线右侧是验证集 top-1 精度随 epoch 的变化。需要注意DEiT 的蒸馏不是简单地把两个 loss 加起来它多了一个“蒸馏 token”这个细节很多人第一次看会忽略。2.2 蒸馏 token 与 class token 的并行机制ViT 在输入序列前会加一个特殊的 class token对应输出的就是分类向量。DEiT 在此基础上又加了一个 distillation token它同样参与 attention 计算但输出只用来计算蒸馏损失不参与最终的类别预测。两个 token 是独立的不会互相干扰。训练时的损失由三部分构成真实标签的交叉熵损失作用于 class token 输出、蒸馏损失作用于 distillation token 输出、以及两者在损失函数中的权重配比。DEiT 用的是 hard distillation公式是# 蒸馏损失用 teacher 的硬预测标签argmax 结果作为伪标签 import torch.nn.functional as F def distillation_loss(student_logits, teacher_logits, labels, alpha0.5, temperature1.0): student_logits: 蒸馏 token 输出shape [batch, num_classes] teacher_logits: teacher 模型输出shape [batch, num_classes] labels: 真实标签 alpha: 蒸馏损失权重DEiT 论文默认 0.5 temperature: 蒸馏温度DEiT 默认 1.0 # 硬标签蒸馏取 teacher 预测的 argmax 作为伪标签 with torch.no_grad(): teacher_pred teacher_logits.argmax(dim1) # 学生蒸馏头与伪标签的交叉熵 distill_loss F.cross_entropy(student_logits, teacher_pred) # 真实标签交叉熵作用于 class token 输出 cls_loss F.cross_entropy(student_logits, labels) return (1 - alpha) * cls_loss alpha * distill_loss这里的temperature1.0意味着 paper 里用的是硬标签蒸馏hard-label distillation不是经典的 soft distillation。为什么因为硬标签蒸馏在 ImageNet 这种大规模分类任务上收敛更快、效果更稳定且避免了对 teacher 输出概率分布的存储开销。alpha0.5是默认权重实际调参中我试过 0.3~0.7 的范围差异不算大但 alpha 太低相当于放弃了蒸馏信号精度会掉 0.5~1 个点。2.3 为什么 teacher 选 CNN 而不是更强的 Transformer这是 DEiT 最反直觉的一个设计选择训练学生 Transformer 的老师偏偏选了一个 CNN 架构RegNetY-16GF。直觉上用一个更大的 Transformer 当老师不是更一致吗但论文实验显示CNN teacher 蒸馏出来的 Transformer 学生在 ImageNet 上取得了比 Transformer teacher 更好的精度。原因在于CNN 的归纳偏置局部性、平移等变性可以被 Transformer 学生“吸收”而 Transformer teacher 的知识结构跟学生太像反而没有提供额外的互补信息。这也意味着 DEiT 的蒸馏本质是把 CNN 的先验知识迁移到 Transformer 上补足后者数据效率不足的短板。你复现这个项目时如果手头没有 RegNetY 的权重可以暂时用 ResNet-50 替代但精度会略有下降约 0.3~0.5 个点因为 RegNetY 本身更强。3. 把推理跑起来预处理、权重加载与单张图片分类3.1 项目文件结构与代码骨架解压资源包后你会看到class.json和 8 张 png 图片。class.json是类别索引到类别名的映射类似{0: cat, 1: dog, ...}。8 张 png 是你的测试素材不需要额外下载数据。我把推理的关键步骤拆成三部分先看完整流程再逐个解释。import json import torch import torchvision.transforms as transforms from PIL import Image # 1. 加载类别映射 with open(class.json, r, encodingutf-8) as f: class_map json.load(f) # class_map 的格式为 {0: 类别A, 1: 类别B, ...} # 注意如果是从官网下载的权重class_map 顺序可能与 ImageNet 原始类别顺序一致 # 2. 定义预处理流程必须与训练时保持一致 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]) ]) # 3. 加载测试图片并预处理增加 batch 维度 image Image.open(0367e0199.png).convert(RGB) input_tensor transform(image).unsqueeze(0) # [1, 3, 224, 224]预处理这一步有三个关键点Resize(256)是为了保证长边大于等于 224这样CenterCrop(224)裁出来的区域包含足够的主体信息不会因为原图比例不对导致目标被裁掉一半Normalize的 mean 和 std 必须用 ImageNet 的统计值因为这个模型的预训练权重就是在 ImageNet 上学的换个统计值等于把输入分布整体偏移了精度直接崩掉。过了预处理之后需要加载模型和权重。注意 DEiT 的权重结构与原生 ViT 有一个关键差异——多了蒸馏分支。import torch from timm import create_model # 4. 创建 DEiT 模型以 deit_small_patch16_224 为例 model create_model(deit_small_patch16_224, pretrainedTrue) model.eval() # 等价于 model.train(False)关闭 dropout 和 BN 的 batch 统计 # 5. 如果使用本地权重比如你从分享链接里下载的 .pth 文件 # state_dict torch.load(deit_small_patch16_224.pth, map_locationcpu) # 删除 head.weight 和 head.bias 前的依赖因为 class.json 的类别数可能不是 1000 # model.load_state_dict(state_dict, strictFalse)使用timm.create_model是最省事的方式它会自动加载在 ImageNet 上预训练好的权重。如果你用的是资源包里的本地权重注意strictFalse参数当你的class.json只有几十个类别时模型的分类头维度不匹配需要把最后一层替换掉。之后推理与 top-5 输出with torch.no_grad(): output model(input_tensor) probabilities torch.softmax(output, dim1) top5_prob, top5_idx torch.topk(probabilities, 5) # 打印 top-5 结果 for i in range(5): idx top5_idx[0][i].item() print(fTop {i1}: {class_map[str(idx)]} ({top5_prob[0][i].item() * 100:.2f}%))这段代码里torch.no_grad()是必须的它关闭了 autograd 的梯度追踪推理模式下能减少显存占用并加速计算。topk直接取概率最高的前 5 个索引再查class_map。如果你的class_map的 key 是字符串而idx是整数记得用str(idx)转换这个类型不匹配的问题我见过不止一个人踩过。3.2 8 张测试图怎么快速批量跑完不要一张张手动跑写个循环批量处理import os from pathlib import Path image_dir Path(.) for img_path in sorted(image_dir.glob(*.png)): img Image.open(img_path).convert(RGB) tensor transform(img).unsqueeze(0) with torch.no_grad(): probs torch.softmax(model(tensor), dim1) top1 probs.argmax(dim1).item() print(f{img_path.name}: 预测为 {class_map[str(top1)]}置信度 {probs[0][top1].item() * 100:.2f}%)这里注意sorted()是为了让输出顺序固定。全局匹配*.png会把文件夹里所有 png 都跑一遍实际使用时你的图片放在这个目录下就行。我在本地用 CPU 跑deit_small_patch16_224单张大约 0.5 秒GPU 上几乎瞬时。4. 训练细节拆解优化器、损失函数与数据增强参数对照4.1 核心训练超参数一览与含义如果你不只是想跑推理还想在自己的数据集上微调或从零训练下面的参数表是从 DEiT 官方配置中提炼出来的核心项。这些参数在资源包的博客原文里有更详细的说明这里我用表格做一个速查。参数值作用优化器AdamW相比 Adam 增加权重衰减解耦Transformer 上收敛更稳基础学习率5e-4batch1024线性缩放规则lr 5e-4 * batch / 1024权重衰减0.05对 Transformer 偏大但配合 AdamW 效果最好训练轮数300 epochDEiT 原文配置小数据集可减到 100学习率调度cosine 衰减从峰值学习率余弦下降到 1e-5warmup5 epoch前 5 个 epoch 从零线性升到目标学习率batch size1024GPU 不够时可用 256/512但要相应调整学习率输入分辨率224x224与预训练权重一致改 384 需微调RandAugment9/0.5增强强度 9概率 0.5MixUp0.8MixUp alpha 参数CutMix1.0CutMix alpha 参数EMA0.9999指数移动平均衰减系数4.2 参数调整原则小数据集与大显存的不同玩法训练超参不是死的我根据自己的复现经验总结了几条实用原则学习率必须跟着 batch size 走。如果你只有单卡 12GB 显存batch size 只能开到 128那学习率就从 5e-4 按比例缩小到大约 1e-4 附近否则会发散。判断标准是第一个 epoch 的 loss 不应该比随机初始化时高太多——如果 loss 不降反升立即把学习率除以 10。RandAugment 的参数含义第一个数字 9 是全局增强强度数值越大增强越剧烈第二个 0.5 是每次应用增强的概率。小数据集上可以适当调高强度到 10~11但不要超过 15否则图像失真严重模型学到的是噪声不是特征。MixUp 和 CutMix 同时启用时每张图有 50% 的概率应用 MixUp、50% 概率应用 CutMix二者互斥。这是 DEiT 原文的实现方式MixUp0.8, CutMix1.0中的数值是 beta 分布的 alpha 参数值越大生成的混合图像越接近原始图。4.3 一个完整的微调代码骨架下面是一段可以直接改改就用的微调脚本核心逻辑import torch from torch import nn from timm import create_model from torch.optim import AdamW from torch.optim.lr_scheduler import CosineAnnealingLR # 用预训练的 DEiT替换分类头 num_classes len(class_map) model create_model(deit_small_patch16_224, pretrainedTrue) model.head nn.Linear(model.head.in_features, num_classes) # 蒸馏分支同样需要替换成对应类别数 model.head_dist nn.Linear(model.head_dist.in_features, num_classes) optimizer AdamW(model.parameters(), lr1e-4, weight_decay0.05) scheduler CosineAnnealingLR(optimizer, T_max100, eta_min1e-5) # 训练循环略每个 batch 时蒸馏分支与分类分支各算一个交叉熵按 alpha0.5 加权T_max100表示学习率在 100 个 epoch 内从 1e-4 余弦衰减到 1e-5。在实际微调时我一般会把 backbone 的 lr 设置成分类头的 0.1 倍因为预训练权重已经学到了足够好的特征全量微调反而会破坏之前的表征。5. 避坑与排查预训练权重、蒸馏分支与归一化三个高频翻车点5.1 现象加载预训练权重时报错维度对不上或不存在的 key这个问题出现的频率极高。timm.create_model(deit_small_patch16_224, pretrainedTrue)触发的默认权重是1000 类 ImageNet 版本。当你把模型的head和head_dist替换成自己数据集的类别数后直接加载旧权重必然报维度不匹配。更隐蔽的情况是加载权重时strictTrue只要有一个 key 的 shape 不同就整体加载失败。原因与解决这是分类头输出维度1000和新分类头维度例如 10不一致。解决方式有两个一是加载权重前先把新分类头接上然后strictFalse加载最后重新初始化分类头二是加载权重后、替换分类头之前就先加载再替换。注意前者更安全因为你不会在替换后忘记重初始化分类头。state_dict torch.load(deit_small_patch16_224.pth, map_locationcpu) # 加载到模型之前先从 state_dict 中删除分类头参数避免 key 不匹配 state_dict.pop(head.weight, None) state_dict.pop(head.bias, None) state_dict.pop(head_dist.weight, None) state_dict.pop(head_dist.bias, None) model.load_state_dict(state_dict, strictFalse)这段代码的逻辑是把分类头相关的四个 key 从权重文件中剔除再用strictFalse加载。此时模型自身新初始化的分类头会保留随机值。如果不做这一步直接strictFalse也可以但模型可能会静默地漏掉某些参数没有加载成功导致推理结果全乱。我习惯把这一步变成固定的动作确保内存中没有旧分类头残留。需要再初始化分类头以满足类别数目要求nn.init.trunc_normal_(model.head.weight, std0.02) nn.init.constant_(model.head.bias, 0) nn.init.trunc_normal_(model.head_dist.weight, std0.02) nn.init.constant_(model.head_dist.bias, 0)5.2 现象推理结果置信度极高但预测类别全错这个我有过惨痛教训。现象是8 张测试图每一张模型都给出了 99% 以上的置信度但预测的类别跟图片内容完全不搭边。此时模型并没有坏问题几乎出在预处理上。三处最常见的错误按频率排序第一忘记CenterCrop直接Resize(224)。原图被压缩变形比例彻底失真。DEiT 训练时用的是Resize(256) CenterCrop(224)推理必须一致否则模型相当于看了一堆畸形图。第二Normalize 的 mean/std 值写反了。有人会把 std 写成 0.5有人把 mean 和 std 的顺序搞混。一旦 Normalize 错的输入分布整体偏移模型输出置信度同样会虚高因为 softmax 的输出并不代表模型真正“有信心”。第三图片是 RGBA4 通道但没转成 RGB。虽然Image.open(...).convert(RGB)已经处理了但如果用cv2.imread读取图片默认是 BGR 顺序通道反转后颜色全偏。5.3 现象模型参数量巨大训练时显存溢出DEiT-Small 有约 2200 万参数理论显存占用不大但实际训练时 12GB 显卡仍然可能撑不住。原因在于 NVIDIA 的 amp自动混合精度没有被启用。PyTorch 1.6 提供了原生的混合精度训练scaler torch.cuda.amp.GradScaler() with torch.cuda.amp.autocast(): output model(input_tensor) loss criterion(output, labels) scaler.scale(loss).backward() scaler.step(optimizer) scaler.update()启用混合精度后显存大约能降到原来的 60%~70%训练速度提升 1.5~2 倍。如果仍然溢出把输入分辨率降到 160但会损失一部分精度。这里要提示如果用了 EMA 更新模型权重EMA 的参数更新要保持全精度混合精度只作用于前向与反向传播。6. 进阶用 DEiT 的蒸馏 token 做模型集成与鲁棒性检查跳过那些“模型能跑就行”的层面DEiT 最容易被忽视的资产是它自带的双头结构——head和head_dist。推理时这两个头分别产生独立的预测分布。虽然论文里只用head作为最终输出但把两个头的概率平均之后相当于做了一次“免费”的模型集成。我的实测经验在 ImageNet 验证集上两个头平均后的 top-1 精度比单独用head高 0.2~0.3 个点且对于容易混淆的相似类别比如不同品种的狗稳定性明显更好。实现起来就一行代码with torch.no_grad(): output model(input_tensor) # model 内部会返回 head 和 head_dist 两个分支的结果 # 如果是 timm 的默认实现需要手动调用模型的 forward_features 再分别过两个 head # 但更简单的方式是直接用 model 的返回默认是 head 的输出严格来说timm 的deit_small_patch16_224在推理模式下默认只返回 class token 的输出蒸馏分支被隐藏了。如果要用双头集成就得显式操作# 获取两个分支的 logits features model.forward_features(input_tensor) head_out model.head(features[:, 0]) # class token 分支 dist_out model.head_dist(features[:, 1]) # distillation token 分支 ensemble_prob (torch.softmax(head_out, dim1) torch.softmax(dist_out, dim1)) / 2这里features[:, 0]对应 class token 位置的输出features[:, 1]对应 distillation token 位置的输出。注意两个分支在训练时被设计为互补关系它们的预测分布不一定一致但平均后通常能消掉一部分模型自身的随机误差。如果你的任务对稳定性要求高比如工业质检场景这个技巧几乎白捡精度。另一个进阶用法是蒸馏 token 的“可解释性”价值——对比head与head_dist的预测差异可以快速定位模型对哪些样本“不确定”。当两个分支给出不同类别时这张图片大概率属于容易混淆的类别或分布外数据。这个信号可以作为人工复核的筛选条件。我一般会留下混淆样本单独看它们的预处理和图像内容确认是数据标注问题还是增强过度。最后说一个习惯性建议无论你的数据集多小微调完成后都要先用测试集跑一遍双分支输出对比如果head和head_dist差异过大超过 5% 的样本预测不一致先检查训练日志有没有过拟合再检查数据增强强度是不是放得太开。从那以后我每次部署 DEiT 分类模型前都会强制走一遍双头 check 归一化复核这两个动作能省下后面大量的返工时间。希望这些细节对你快速落地这个项目有帮助。本文还有配套的精品资源点击获取