尧图网站设计 尧图网站设计YAOTU DESIGN
ARTICLE DETAIL

资讯详情

深耕网站设计与一线实操的经验洞察。

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

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预训练权重结构、蒸馏头和训练参数全是未知数。DEiTData-efficient Image Transformers就是冲这个场景来的它用知识蒸馏加一个额外的蒸馏token让Transformer在中小数据集上也能训出接近CNN的精度。解压这个zip工程包后你拿到的是一套把DEiT用于图像分类的最小工程数据准备、训练、验证、推理都在里面。适合谁手里有几千到几万张图的算法工程师、做边缘设备分类方案的人以及想从CNN切到Transformer但怕翻车的人。这篇笔记按你最基本的诉求来写这是什么、怎么跑通、参数怎么调、坑在哪。2. DEiT的两个关键机制蒸馏token与训练策略2.1 distillation tokenCNN老师怎么教Transformer学生DEiT在结构上最直观的变化是在patch embedding之后、输入Transformer encoder之前除了标准的class token就是ViT里那个[CLS]之外再多拼接一个蒸馏tokendistillation 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幅度是9Mixup系数是0.8CutMix概率1.0这套组合在几十万张图上没问题但当你只有几千张图时它们就是过拟合的刹车片——强得过头。我自己的经验是先把RandAugment降到2到3Mixup降到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。下表是三个常用变体的基本盘模型参数量输入分辨率单卡batch64建议显存典型用途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 python3.8 conda activate deit pip install torch1.13.1 torchvision0.14.1 pip install timm0.9.12装完后用一段几十秒的脚本确认CUDA可用、模型能前向跑通import torch, timm model timm.create_model(deit_tiny_patch16_224, pretrainedTrue, num_classes10) model model.cuda().eval() dummy torch.randn(4, 3, 224, 224).cuda() with torch.no_grad(): out model(dummy) print(type(out))这里注意timm的DeiT在训练模式下返回的是tupleeval模式下返回的是包好的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, transformtransform_train) train_loader DataLoader(train_ds, batch_size64, shuffleTrue, num_workers8, pin_memoryTrue) # 学生DEiT-Tiny model timm.create_model(deit_tiny_patch16_224, pretrainedTrue, num_classeslen(train_ds.classes)) model.cuda() # 老师ResNet50评估模式不更新梯度 teacher timm.create_model(resnet50, pretrainedFalse, num_classeslen(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(), lr5e-4, weight_decay0.05) scheduler torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max30) 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(dim1) 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(fepoch {epoch1} 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-4AdamWbatch小时lr要更小warmup5 epoch3-5 epoch不建议省RandAugment92-3增强强度先压下来Mixup0.80-0.2小数据开满会拖慢收敛CutMix1.00-0.3同上蒸馏温度3soft变体3hard模式不涉及硬标签蒸馏不需要温度EMA0.999960.999验证时用EMA权重更稳这里最值得盯的是RandAugment和Mixup。DEiT这类Transformer对增强的依赖性强但小数据集上增强过头会直接造成验证集不涨、训练集已经99%的假象。我见过不少团队把原论文ImageNet配置复制过来然后调了三天模型结构最后发现是Mixup开0.8把几百张图的模型搞崩了。如果计算资源有限先跑一个不带Mixup、RandAugment2的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(dim1) pred5 logits.topk(5, dim1).indices correct1 (pred1 labels).sum().item() correct5 (pred5 labels.unsqueeze(1)).any(dim1).sum().item() total labels.size(0) all_pred.extend(pred1.cpu().tolist()) all_label.extend(labels.cpu().tolist()) print(ftop-1 {correct1/total:.4f} top-5 {correct5/total:.4f})top-5计算的逻辑是对每个样本看真实标签是否出现在topk5返回的索引张量里。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, pretrainedFalse, num_classeslen(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, dim1) top_p, top_i probs.topk(3, dim1) 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_version14, 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_norm5.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或3Mixup改成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, strictFalse)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在低增强下跑通全流程再逐步往上堆配置。希望帮到你。本文还有配套的精品资源点击获取
返回列表