简介:面向深度学习与计算机视觉学习者,这份DEiT实战资源围绕Facebook提出的DeiT模型,展示如何在不依赖外部数据集的情况下,利用知识蒸馏策略完成ImageNet级别的高效训练,并落地到图像分类任务中。DeiT通过引入蒸馏令牌与教师模型交互,显著降低训练成本,四块GPU三天即可达到SOTA水平,适合希望快速上手Transformer图像分类、但缺乏超算环境的研究者与工程师。压缩包共2445个文件,主体为2437张可视化图片,涵盖训练曲线、混淆矩阵、样本预测等结果;另有6个Python脚本负责模型搭建、数据加载、训练与评估流程,1个JSON文件存储类别映射,1个TXT文档补充说明,整体约737MB,目录结构清晰,便于对照学习。目前已有869人学习该资源。下载后可获得从数据准备、脚本配置到模型训练、指标分析的一整套可参考实现,借助过程图片能快速验证蒸馏效果,迁移至自有分类数据集。
1. 直接上手 DEiT:别人三天训完 ImageNet,我们十分钟跑通推理
图像分类这个方向,2020 年之前基本是 CNN 的天下,Vision Transformer(ViT)虽然效果好,但训练极其挑剔,动辄几十个 epoch 加上超大 batch,普通实验室根本玩不转。DEiT(Data-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 针对这一点给出的答案很直接:不让模型从零硬学,而是让一个训练好的 CNN(RegNetY 系列)当老师,用蒸馏损失把知识“压”进学生 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, alpha=0.5, temperature=1.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(dim=1) # 学生蒸馏头与伪标签的交叉熵 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这里的temperature=1.0意味着 paper 里用的是硬标签蒸馏(hard-label distillation),不是经典的 soft distillation。为什么?因为硬标签蒸馏在 ImageNet 这种大规模分类任务上收敛更快、效果更稳定,且避免了对 teacher 输出概率分布的存储开销。alpha=0.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', encoding='utf-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', pretrained=True) model.eval() # 等价于 model.train(False),关闭 dropout 和 BN 的 batch 统计 # 5. 如果使用本地权重(比如你从分享链接里下载的 .pth 文件) # state_dict = torch.load('deit_small_patch16_224.pth', map_location='cpu') # 删除 'head.weight' 和 'head.bias' 前的依赖,因为 class.json 的类别数可能不是 1000 # model.load_state_dict(state_dict, strict=False)使用timm.create_model是最省事的方式,它会自动加载在 ImageNet 上预训练好的权重。如果你用的是资源包里的本地权重,注意strict=False参数:当你的class.json只有几十个类别时,模型的分类头维度不匹配,需要把最后一层替换掉。之后推理与 top-5 输出:
with torch.no_grad(): output = model(input_tensor) probabilities = torch.softmax(output, dim=1) top5_prob, top5_idx = torch.topk(probabilities, 5) # 打印 top-5 结果 for i in range(5): idx = top5_idx[0][i].item() print(f"Top {i+1}: {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), dim=1) top1 = probs.argmax(dim=1).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-4(batch=1024) | 线性缩放规则:lr = 5e-4 * batch / 1024 |
| 权重衰减 | 0.05 | 对 Transformer 偏大,但配合 AdamW 效果最好 |
| 训练轮数 | 300 epoch | DEiT 原文配置,小数据集可减到 100 |
| 学习率调度 | cosine 衰减 | 从峰值学习率余弦下降到 1e-5 |
| warmup | 5 epoch | 前 5 个 epoch 从零线性升到目标学习率 |
| batch size | 1024 | GPU 不够时可用 256/512,但要相应调整学习率 |
| 输入分辨率 | 224x224 | 与预训练权重一致,改 384 需微调 |
| RandAugment | 9/0.5 | 增强强度 9,概率 0.5 |
| MixUp | 0.8 | MixUp alpha 参数 |
| CutMix | 1.0 | CutMix alpha 参数 |
| EMA | 0.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 原文的实现方式,MixUp=0.8, CutMix=1.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', pretrained=True) 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(), lr=1e-4, weight_decay=0.05) scheduler = CosineAnnealingLR(optimizer, T_max=100, eta_min=1e-5) # 训练循环略:每个 batch 时,蒸馏分支与分类分支各算一个交叉熵,按 alpha=0.5 加权T_max=100表示学习率在 100 个 epoch 内从 1e-4 余弦衰减到 1e-5。在实际微调时,我一般会把 backbone 的 lr 设置成分类头的 0.1 倍,因为预训练权重已经学到了足够好的特征,全量微调反而会破坏之前的表征。
5. 避坑与排查:预训练权重、蒸馏分支与归一化三个高频翻车点
5.1 现象:加载预训练权重时报错,维度对不上或不存在的 key
这个问题出现的频率极高。timm.create_model('deit_small_patch16_224', pretrained=True)触发的默认权重是1000 类 ImageNet 版本。当你把模型的head和head_dist替换成自己数据集的类别数后,直接加载旧权重必然报维度不匹配。更隐蔽的情况是:加载权重时strict=True,只要有一个 key 的 shape 不同就整体加载失败。
原因与解决:这是分类头输出维度(1000)和新分类头维度(例如 10)不一致。解决方式有两个:一是加载权重前先把新分类头接上,然后strict=False加载,最后重新初始化分类头;二是加载权重后、替换分类头之前就先加载,再替换。注意前者更安全,因为你不会在替换后忘记重初始化分类头。
state_dict = torch.load('deit_small_patch16_224.pth', map_location='cpu') # 加载到模型之前,先从 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, strict=False)这段代码的逻辑是:把分类头相关的四个 key 从权重文件中剔除,再用strict=False加载。此时模型自身新初始化的分类头会保留随机值。如果不做这一步,直接strict=False也可以,但模型可能会静默地漏掉某些参数没有加载成功,导致推理结果全乱。我习惯把这一步变成固定的动作,确保内存中没有旧分类头残留。需要再初始化分类头以满足类别数目要求:
nn.init.trunc_normal_(model.head.weight, std=0.02) nn.init.constant_(model.head.bias, 0) nn.init.trunc_normal_(model.head_dist.weight, std=0.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 的输出并不代表模型真正“有信心”。第三,图片是 RGBA(4 通道)但没转成 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, dim=1) + torch.softmax(dist_out, dim=1)) / 2这里features[:, 0]对应 class token 位置的输出,features[:, 1]对应 distillation token 位置的输出。注意:两个分支在训练时被设计为互补关系,它们的预测分布不一定一致,但平均后通常能消掉一部分模型自身的随机误差。如果你的任务对稳定性要求高(比如工业质检场景),这个技巧几乎白捡精度。
另一个进阶用法是蒸馏 token 的“可解释性”价值——对比head与head_dist的预测差异,可以快速定位模型对哪些样本“不确定”。当两个分支给出不同类别时,这张图片大概率属于容易混淆的类别或分布外数据。这个信号可以作为人工复核的筛选条件。我一般会留下混淆样本,单独看它们的预处理和图像内容,确认是数据标注问题还是增强过度。
最后说一个习惯性建议:无论你的数据集多小,微调完成后都要先用测试集跑一遍双分支输出对比,如果head和head_dist差异过大(超过 5% 的样本预测不一致),先检查训练日志有没有过拟合,再检查数据增强强度是不是放得太开。从那以后,我每次部署 DEiT 分类模型前,都会强制走一遍双头 check + 归一化复核这两个动作,能省下后面大量的返工时间。希望这些细节对你快速落地这个项目有帮助。
本文还有配套的精品资源,点击获取