news 2026/10/4 18:08:41

DEiT实战指南:中小数据集上高效训练图像分类Transformer

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
DEiT实战指南:中小数据集上高效训练图像分类Transformer

简介:DEiT 是由 Facebook 在 2020 年提出的高效图像分类 Transformer 模型,通过知识蒸馏与训练策略改进,消除了 Transformer 难以训练的痛点,在仅使用 ImageNet 数据、4 块 GPU 训练三天的条件下就达到了 SOTA 水平。该压缩包围绕 DEiT 实战展开,面向有一定深度学习基础、希望将 Transformer 应用于图像分类任务并复现实验的开发者与研究者。压缩包共 2445 个文件,整体约 736.96 MB;其中约 2437 张 png 图片是训练过程可视化与结果图表,便于直观对比不同配置,另有 6 个 Python 脚本负责数据准备、模型构建、训练与推理,1 个 JSON 文件保存类别映射及实验配置,1 个 TXT 文件提供说明。已有 870 人浏览学习。通过这份实战资源,可以系统掌握 DEiT 的知识蒸馏训练思路和图像分类完整流程,从数据组织、参数配置到模型评估都能快速上手,并结合可视化图表分析训练动态,减少自行复现的时间成本。

1. DEiT是什么:中小数据集也能训Transformer的图像分类方案

如果你手里只有几千张标注图片,却想用最新的图像分类模型思路试一试Transformer,通常会先被两个事实劝退:一是ViT从零训练在小数据上几乎必过拟合,二是想微调ImageNet预训练权重,结构、蒸馏头和训练参数全是未知数。DEiT(Data-efficient Image Transformers)就是冲这个场景来的,它用知识蒸馏加一个额外的蒸馏token,让Transformer在中小数据集上也能训出接近CNN的精度。解压这个zip工程包后,你拿到的是一套把DEiT用于图像分类的最小工程:数据准备、训练、验证、推理都在里面。适合谁?手里有几千到几万张图的算法工程师、做边缘设备分类方案的人,以及想从CNN切到Transformer但怕翻车的人。这篇笔记按你最基本的诉求来写:这是什么、怎么跑通、参数怎么调、坑在哪。

2. DEiT的两个关键机制:蒸馏token与训练策略

2.1 distillation token:CNN老师怎么教Transformer学生

DEiT在结构上最直观的变化,是在patch embedding之后、输入Transformer encoder之前,除了标准的class token(就是ViT里那个[CLS])之外,再多拼接一个蒸馏token(distillation token)。这个token和class token一起走完整个Transformer encoder,但最后各连各的分类头:class token的head学习真实标签,distillation token的head学习“老师”给的标签。

老师是谁?DEiT原论文用的是RegNet这一类CNN,而不是另一个Transformer。这有一个实际好处:CNN和Transformer的结构差异大,蒸馏出来的信息互补性更强。学生模型同时学两个任务——真实类别和老师的预测类别——相当于把CNN看到的“图像归纳偏置”通过soft label灌进Transformer里。

这里有一个很多复现工程容易搞错的细节:原论文用的是hard-label distillation,不是常见的那种KL散度软蒸馏。它把老师的argmax预测当作硬标签,直接做交叉熵。你翻DEiT开源训练脚本会发现它是这样写的。用硬标签的好处是稳定、不用调温度、也不会因为teacher输出分布太尖导致loss爆炸。我前几次复现时一律改成KL散度,结果在小数据集上反而更差——这个后面避坑章再展开。

推理的时候,两个head都可以用。原论文做了消融:class token的head和distillation token的head精度差不多,推理时取两者平均通常更好。在timm的实现里,eval模式下默认就是返回两者平均。

2.2 不只是蒸馏:数据增强与正则化策略为什么不能被跳过

DEiT的论文标题里“Data-efficient”其实不只靠蒸馏。Transformer没有CNN那种天然的平移不变性和局部性先验,所以它对数据增强的依赖比ResNet高得多。DEiT把当时CNN训练里一套完整增强组合直接搬了过来:RandAugment、Mixup、CutMix、随机擦除、EMA权重平均。这套策略在ImageNet上配合300个epoch,让DEiT-Small在无额外数据的情况下超过了同规模CNN。

落到你自己的数据集上时,这句话得打个折扣。DEiT原配置的RandAugment幅度是9,Mixup系数是0.8,CutMix概率1.0,这套组合在几十万张图上没问题,但当你只有几千张图时,它们就是过拟合的刹车片——强得过头。我自己的经验是,先把RandAugment降到2到3,Mixup降到0.2以内,CutMix先关掉,跑通一个baseline后再逐步往上加。这套策略在代码里的位置很关键,绝大多数训练“不收敛”的锅不在模型结构,而在增强强度和数据量不匹配。

2.3 什么时候选DEiT,什么时候老老实实用CNN

选型这件事,直接决定你这几天加班值不值。DEiT适合的场景是:数据量在1千到10万张之间,类别数5到100,且你已经决定后续要做注意力可视化、多模态融合,或者就是想从CNN切到Transformer。如果你的数据每类只有几十张、还没有预训练权重,那DEiT救不了你,ResNet50微调会是更稳的起点。

另一个务实判断是看算力。DeiT-Tiny只有5.7M参数,一张8GB显卡也能跑;DeiT-Base有86M参数,再挂一个teacher模型,显存压力不小。所以团队如果是第一次接触Transformer,我一般建议从Tiny开始,跑通整个流程再换Small或Base。下表是三个常用变体的基本盘:

模型参数量输入分辨率单卡batch=64建议显存典型用途
deit_tiny_patch16_2245.7M224x2246-8GB先跑通流程、边缘部署
deit_small_patch16_22422M224x2248-11GB精度/速度均衡
deit_base_patch16_22486M224x22412-16GB数据量大、追求上限

记住一个反直觉的结论:DEiT在数据量很小的时候,精度不一定比CNN高,它真正的优势区间是中量数据(几千到几万张)外加预训练权重。拿它硬刚几百张图的小样本任务,不是它的主场。

3. 把DEiT跑起来:环境、数据与最小训练脚本

3.1 环境准备:PyTorch、timm与GPU显存底线

这个工程包我默认你已有Linux服务器和一张NVIDIA显卡。环境按最小依赖来装:

conda create -n deit python=3.8 conda activate deit pip install torch==1.13.1 torchvision==0.14.1 pip install timm==0.9.12

装完后用一段几十秒的脚本确认CUDA可用、模型能前向跑通:

import torch, timm model = timm.create_model('deit_tiny_patch16_224', pretrained=True, num_classes=10) model = model.cuda().eval() dummy = torch.randn(4, 3, 224, 224).cuda() with torch.no_grad(): out = model(dummy) print(type(out))

这里注意timm的DeiT在训练模式下返回的是tuple,eval模式下返回的是包好的tensor,而且不同timm版本行为有差异。0.9.x的eval模式默认把class token和distillation token两个分支的输出做了平均,所以你在eval模式下拿到的就是一个已经融合好的预测。这在后面自定义训练循环时需要单独处理。装完后第一件事不是急着看准确率,而是先确认前向输出类型,避免写训练循环时在解包阶段翻车。

3.2 数据集准备:目录结构和类别均衡判断

工程包内“数据”目录的预期结构就是torchvision标准的ImageFolder形式,这一点必须建立。每个子文件夹一个类别,文件夹名就是类别名,训练和验证分开两个目录:

data/ train/ class_a/ 0001.jpg ... class_b/ 0001.jpg ... val/ class_a/ 0001.jpg ... class_b/ 0001.jpg ...

如果只有一份全量数据,常见做法是先用脚本按8:2或9:1分出一部分做验证。切分时要注意:直接全体乱序切分在图像分类里会有问题——同一张图的近邻帧可能同时出现在训练和验证,导致验证指标虚高。如果有时间戳或场景编号,最好按场景分组再切。

切完后统计一下每类数量,打印一下就够:

from collections import Counter from torchvision.datasets import ImageFolder ds = ImageFolder('data/train') cnt = Counter(ds.targets) print({ds.classes[i]: c for i, c in cnt.items()})

如果发现某类样本数是另一类的10倍以上,训练时就要处理类别不均衡。DEiT对这个问题比CNN敏感,尾部类别很容易被头部类别吃掉。治标做法是loss加权或采样器加权,治本做法是收集更多尾部类数据。这个判断放在训练脚本之前,比训完再发现问题要省一天时间。

3.3 最小训练脚本:以timm 0.9.x为例

训练脚本按“学生模型 + teacher模型”两条线写。teacher先用ResNet50在同样的数据上训好,把权重存成teacher.pth。这里展示核心训练循环:

import torch, timm, torch.nn as nn from torch.utils.data import DataLoader from torchvision import datasets, transforms from timm.data import Mixup from timm.loss import SoftTargetCrossEntropy # 数据增强:小数据集从低强度开始 transform_train = transforms.Compose([ transforms.Resize((256, 256)), transforms.RandomResizedCrop(224, scale=(0.6, 1.0)), transforms.ToTensor(), transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]) ]) train_ds = datasets.ImageFolder('data/train', transform=transform_train) train_loader = DataLoader(train_ds, batch_size=64, shuffle=True, num_workers=8, pin_memory=True) # 学生:DEiT-Tiny model = timm.create_model('deit_tiny_patch16_224', pretrained=True, num_classes=len(train_ds.classes)) model.cuda() # 老师:ResNet50,评估模式,不更新梯度 teacher = timm.create_model('resnet50', pretrained=False, num_classes=len(train_ds.classes)) teacher.load_state_dict(torch.load('teacher.pth')) teacher.cuda().eval() for p in teacher.parameters(): p.requires_grad_(False) criterion_cls = nn.CrossEntropyLoss() # class token 的真实标签监督 criterion_dist = nn.CrossEntropyLoss() # distillation token 的老师硬标签监督 optimizer = torch.optim.AdamW(model.parameters(), lr=5e-4, weight_decay=0.05) scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=30) for epoch in range(30): model.train() for images, labels in train_loader: images, labels = images.cuda(), labels.cuda() # 训练模式下 timm 返回 (class_logits, dist_logits) class_logits, dist_logits = model(images) loss_cls = criterion_cls(class_logits, labels) with torch.no_grad(): # 老师的硬标签:argmax teacher_label = teacher(images).argmax(dim=1) loss_dist = criterion_dist(dist_logits, teacher_label) loss = 0.5 * loss_cls + 0.5 * loss_dist optimizer.zero_grad() loss.backward() optimizer.step() scheduler.step() print(f"epoch {epoch+1} loss {loss.item():.4f}")

代码里有两个值得说明的点。第一,model(images)在timm 0.9.x训练模式下返回的是(class_logits, dist_logits)元组,但不同版本可能只返回一个tensor。如果你用的版本不是元组,就用model.forward_features(images)拿到特征后手动取feat[:, 0]过model.head、取feat[:, 1]过model.head_dist。第二,蒸馏loss直接用老师的argmax硬标签做交叉熵,这是DEiT原版做法,比KL散度稳定得多、也不用调温度。两路loss先各占0.5,跑完一个baseline后如果你想偏向真实标签,可以改成0.7和0.3。

teacher模型必须放在eval模式并且包进torch.no_grad()里——它每一轮前向都要算一遍,如果跟学生一样开梯度,显存和耗时直接翻倍,这是新手最容易漏的。

3.4 几个必调的超参数:lr、warmup、蒸馏温度

DEiT在ImageNet上的原始配置不能直接搬到小数据集,按经验给小数据集一套落地参数:

超参数DEiT原论文配置小数据集建议值说明
epoch30030-100几千张图300轮必过拟合
lr1e-32e-4至5e-4AdamW,batch小时lr要更小
warmup5 epoch3-5 epoch不建议省
RandAugment92-3增强强度先压下来
Mixup0.80-0.2小数据开满会拖慢收敛
CutMix1.00-0.3同上
蒸馏温度3(soft变体)3(hard模式不涉及)硬标签蒸馏不需要温度
EMA0.999960.999验证时用EMA权重更稳

这里最值得盯的是RandAugment和Mixup。DEiT这类Transformer对增强的依赖性强,但小数据集上增强过头会直接造成验证集不涨、训练集已经99%的假象。我见过不少团队把原论文ImageNet配置复制过来,然后调了三天模型结构,最后发现是Mixup开0.8把几百张图的模型搞崩了。如果计算资源有限,先跑一个不带Mixup、RandAugment=2的baseline,再削减增强一点一点往上加,比一开始就上满配要容易定位问题。

4. 验证与推理:把训练好的DEiT模型用起来

4.1 验证脚本:top-1/top-5和混淆矩阵

训练结束后,验证时要注意DEiT的eval模式和训练模式输出不一样。eval模式下timm默认把两个head的输出做了平均,也就是说你不再需要解包元组,直接拿model(images)就能得到最终预测:

model.eval() correct1 = correct5 = total = 0 all_pred, all_label = [], [] with torch.no_grad(): for images, labels in val_loader: images, labels = images.cuda(), labels.cuda() logits = model(images) # eval模式下已融合两个head pred1 = logits.argmax(dim=1) pred5 = logits.topk(5, dim=1).indices correct1 += (pred1 == labels).sum().item() correct5 += (pred5 == labels.unsqueeze(1)).any(dim=1).sum().item() total += labels.size(0) all_pred.extend(pred1.cpu().tolist()) all_label.extend(labels.cpu().tolist()) print(f"top-1 {correct1/total:.4f} top-5 {correct5/total:.4f}")

top-5计算的逻辑是:对每个样本,看真实标签是否出现在topk=5返回的索引张量里。labels.unsqueeze(1)把标签变成[batch, 1],和[batch, 5]比较后再沿第二维做any判断。验证集还建议顺手输出每类别的top-1精度,比一个全局数字更能暴露尾部类别问题。你可以把上面代码按all_label分组统计,或者直接用sklearn的classification_report。

4.2 单张图片推理与类别映射

单张推理脚本比验证更简单,但有一个坑:类别索引和类别名的映射要从训练集的classes列表里保存下来,否则预测出来一个int根本不知道是什么类。推理代码:

import torch, timm from PIL import Image from torchvision import transforms idx_to_class = {i: c for c, i in train_ds.class_to_idx.items()} model = timm.create_model('deit_tiny_patch16_224', pretrained=False, num_classes=len(idx_to_class)) model.load_state_dict(torch.load('best_model.pth')) model.cuda().eval() tf = transforms.Compose([ transforms.Resize((256, 256)), transforms.CenterCrop(224), transforms.ToTensor(), transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]) ]) img = Image.open('test.jpg').convert('RGB') x = tf(img).unsqueeze(0).cuda() with torch.no_grad(): logits = model(x) probs = torch.softmax(logits, dim=1) top_p, top_i = probs.topk(3, dim=1) for p, i in zip(top_p[0].cpu().tolist(), top_i[0].cpu().tolist()): print(f"{idx_to_class[i]}: {p:.3f}")

注意convert('RGB')不可省——灰度图或带透明通道的PNG会让前向报维度错误,这在部署现场很常见。预处理里训练用了RandomResizedCrop,推理时就换成Resize加CenterCrop,尺寸对齐224。如果想要更稳的推理,可以把10个不同crop的预测平均一下,但对一个分类任务来说收益有限,我一般只在比赛或验收时加。

4.3 导出到ONNX:注意distillation分支的取舍

DEiT导出ONNX时最需要想清楚的是:导出哪个分支。eval模式下timm返回的是两个head的平均值,但ONNX导出的是整个计算图,它会连带着把distillation head一起导出来,导致输出节点冗余、推理引擎多算一个全连接层。

常见做法是只保留class token分支,导出前手动构造推理逻辑:

model.eval() class_head = model.head def forward_single(x): feat = model.forward_features(x) return class_head(feat[:, 0]) # class token位置 dummy = torch.randn(1, 3, 224, 224).cuda() torch.onnx.export( forward_single, dummy, "deit_tiny.onnx", input_names=["input"], output_names=["logits"], opset_version=14, dynamic_axes={"input": {0: "batch"}, "logits": {0: "batch"}} )

feat[:, 0]取的是class token这个位置,因为DEiT的输入顺序是class token、distillation token、patch tokens。只导这一个分支在绝大多数部署场景都够用,distillation head的收益在推理时通常也就零点几个点。导出后用onnxruntime加载跑一遍同一张图,与PyTorch的结果做误差对比,差异小于1e-4就说明计算图没有问题。这一步被很多人跳过,等部署环境里发现输出张量多了个维度才回来补,很浪费工时。

5. DEiT实战避坑:5个翻车现场和背后的原因

5.1 现象:loss在20轮附近变成NaN,训练白跑

原因大多是混合精度下蒸馏分支的数值炸了。AMP训练时class token分支的CE还好,但teacher的logits在FP16下做argmax,或者KL散度里出现极端分布,反向传播时梯度溢出。

解决:teacher的logits计算保持在FP32,蒸馏loss用硬标签CE而不是KL散度,并且在梯度更新前加一个torch.nn.utils.clip_grad_norm_保底:

torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=5.0)

如果用了AMP,把loss计算放进scaler.scale(loss)的统一入口里,student和teacher不要分两个精度体系。这个坑在训练中段出现,排查起来最耗时间,建议一开始就做好防御。

5.2 现象:val acc全程没动,train acc却在快速逼近100%

这就是典型的增强强度过猛,模型在训练集上死记硬背,在验证集上完全泛化不了。我把RandAugment调成9、Mixup开0.8的时候,在3000张图的数据集上连续三天看到这个曲线,一度以为是代码里label shuffle了。

解决:把RandAugment默认幅度降到2或3,Mixup改成0.2以内,CutMix直接关掉。跑一个干净baseline,确认val loss在下降,再按每5个epoch为一个周期把增强往上加。如果加了之后val acc不掉反升,就说明这个增强在你的数据量级是正向的。

5.3 现象:显存不够,哪怕用Tiny也OOM

原因不是模型太大,而是teacher和student同时前向时占用了双份激活值,尤其batch size设到128以上。ResNet50做teacher虽然模型不大,但中间feature map很占显存。

解决:teacher必须eval()加torch.no_grad(),这一步能省一半显存。还不够就把batch降到32或16,不要开gradient checkpointing——在Transformer上它会拖慢速度,小batch比checkpointing更划算。另外,如果同时开了Mixup和CutMix,它们会额外多占一部分临时张量,可以暂时关掉来看显存变化,逐个定位是哪个组件在吃显存。

5.4 现象:加载预训练权重时key mismatch,代码直接崩

原因通常是num_classes不匹配导致分类头、distillation头的权重被重置,或者timm版本不同导致state_dict里多出多余参数。直接加载ImageNet的1000类权重再替换head是常见做法,但很多人忽视了distillation head也需要同步替换。

解决:用timm的create_model时直接指定num_classes,它会自动处理分类头维度的变化。如果是自己写加载逻辑,加载时要过滤掉head.和head_dist.前缀的键:

state = torch.load('deit_tiny_imagenet.pth') state = {k: v for k, v in state.items() if not k.startswith(('head.', 'head_dist.'))} model.load_state_dict(state, strict=False)

5.5 现象:类别不均衡时,少数类全被预测成头部类

DEiT在数据不均衡时比CNN更依赖标签分布。训练集里A类有5000张、B类只有50张,最终输出几乎全是A。这是Transformer在小样本尾部类上的通病,不是代码bug。

解决:先用WeightedRandomSampler把采样权重拉平,再训练一个baseline;如果还不行,给尾部类加loss权重。验证时别只看全局top-1,按类别打印结果对照。森林图像分类这类数据经常一头沉,用这个方案能把少数类准确率拉回几个点,但别指望它能从10%变成80%——样本太少时,数据增强和数据补充比任何模型技巧都管用。“血泪经验”是:与其调半天loss权重,不如先花力气多标几百张少数类数据,效果立竿见影。

6. 迁移到森林图像分类:两阶段微调技巧与最后一个参数玄学

森林图像分类是个很典型的DEiT落地场景:类别少则五六个树种、多则二三十类,每类样本从几百到几千张不等。这种数据规模完全在DEiT的舒适区里。但直接加载ImageNet权重、替换分类头、全参数微调不是最优做法。我习惯分两个阶段走。

第一阶段是linear probe:把backbone所有参数冻结,只训练分类头。用很小的学习率,比如1e-3,跑5到10个epoch,只看验证集是否在正常下降。这一步的目的不是拿到精度,而是判断预训练权重和你的目标域是否匹配。如果linear probe的验证正确率能很快冲到70%以上,说明特征提取层和森林图像域差异不大,可以继续做第二阶段。如果linear probe怎么训都不超过50%,那说明还不如从零训一个ResNet,别在Transformer上硬磨。

第二阶段再解冻后几个Transformer block和整个分类头,用5e-5到2e-4的学习率微调十几个epoch。只解冻后半段而不是全参数,是因为森林图像和ImageNet的底层纹理特征(边缘、颜色、纹理)仍然共享,真正需要适配的是高层语义。

最后说一个玄学参数——蒸馏温度。上面我用的是hard-label蒸馏所以不需要温度,但如果你想换成soft蒸馏(毕竟timm的teacher输出都是softmax概率),温度T从3起步。类间相似度高的任务——比如近缘树种区分、病害早期叶片识别——T调到4会略微改善,因为更高的温度会把teacher输出中的模糊信息保留得更充分。这是我在几个森林数据集上调参得到的感受,不算严格结论,但值得一试。我自己的教训是:不要一上来就把所有增强全开,更别一上来就换Base模型,这两件事能把一个本来半天能跑完的实验拖成一周。先把Tiny在低增强下跑通全流程,再逐步往上堆配置。希望帮到你。

本文还有配套的精品资源,点击获取

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/10/4 18:07:43

MRAM取代Flash与EEPROM:STM32掉电数据保存实战方案

做工业设备的嵌入式开发,最绕不开的老大难就是数据掉电保存。EEPROM写寿命有限,Flash要先擦后写又慢得让人心焦,特别是处理高频参数记录和故障瞬间保存这种需求,总得在容量、寿命、速度之间反复妥协。这两年我在几个项目里改用 Ev…

作者头像 李华
网站建设 2026/10/4 18:07:35

别再问“哪款AI写论文最牛”了:土木水利交通人的毕设辅助工具组合,建议这样配 [特殊字符]️

如果你读的是土木、水利与交通工程,大概率会遇到一类很典型的毕业设计:某市滨河路段雨水管网提标改造与内涝整治设计你需要分析区域降雨和道路积水情况,进行汇水区域划分、雨水流量计算、管网管径与坡度设计,必要时用 SWMM、InfoW…

作者头像 李华
网站建设 2026/10/4 18:02:41

智能公关平台MediaBee深度解析:数据驱动的媒体关系管理

1. 为什么传统发稿模式越来越“失灵”了先说一个我自己的感受。这两年在公关圈子里,有个挺明显的趋势:过去那种“写好新闻稿、群发给媒体、然后等着看剪报”的传统发稿流程,正在肉眼可见地失效。信息传播路径碎得不成样子,记者看邮…

作者头像 李华
网站建设 2026/10/4 17:58:16

DeepSeek Harness桌面端安装部署与插件Skill机制全解析

1. 从命令行到桌面窗口:DeepSeek Harness 桌面端到底解决了谁的痛点第一次听说 DeepSeek Harness 出了桌面端,我的反应是"终于有人干了这件事"。如果你之前用过命令行版本的 Harness,应该能理解那种感受——功能确实强,…

作者头像 李华
网站建设 2026/10/4 17:55:24

从零攻克Python作业:环境配置、类型转换与实战案例解析

从“python作业”这四个字,我就能感受到两种截然不同的情绪:一种是刚接触编程的兴奋,另一种是完全不知道从何下手的焦虑。作为一门语言,Python在数据处理、Web开发、自动化脚本这些领域几乎无所不能,但落到具体的“作业…

作者头像 李华