尧图精选

DEiT图像分类实战:小数据集上高效训练Transformer模型

🕒 发布时间:2026/10/1 11:02:46 📁 来源:尧图网络
简介DEiT实战资源以Facebook提出的DeiT模型为核心面向想在图像分类任务中应用Transformer的深度学习开发者尤其适合已掌握PyTorch基础、希望深入理解蒸馏训练技巧的读者。压缩包共2445个文件以2437张png图表为主辅以6个Python脚本、1个json配置和1个txt说明整体约737MB内容覆盖数据集预览、训练过程可视化及代码配置等环节。已有870人学习浏览可对照原理解析进行复现。资源内不仅提供分类模型实现还包含日志、图表与类别映射文件方便读者对比不同阶段的准确率变化梳理训练流程可作为图像分类实战的完整参考。1. 拿着 DEiT 做图像分类先搞懂它为什么“省数据”在图像分类模型里翻来翻去最后让我决定用 DEiT 而不是直接套 ViT 的是它那个“数据高效”的标签。DEiT 是 Data-efficient Image Transformers 的缩写本质是一个用蒸馏训练出来的 transformer 图像分类模型。它解决的不是“分类精度上限”问题而是“只有一两万张图时transformer 图像分类模型怎么才能不翻车”的问题。森林图像分类、工业零件缺陷分类、医学切片小样本分类这类场景最典型的需求就是数据量有限却想吃到 ViT 在归纳偏置上的红利。这套方案适合谁适合手里有自定义图像分类数据集、想低成本从 CNN 迁移到 transformer又不打算靠海量预训练数据硬扛的人。下面我就按自己实际做项目的顺序把环境、训练、推理和坑都摊开讲。2. 跑通 DEiT 图像分类的最小环境与数据准备2.1 为什么选 DEiT蒸馏训练补上 ViT 的小数据短板ViT 在小数据集上不如 ResNet这是公开的结论原因在于 ViT 缺少 CNN 的局部先验全靠数据量把注意力模式“喂”出来。DEiT 的做法很直接引入一个 teacher 模型一般是 RegNetY 这类带卷积结构的强分类器训练时让 DEiT 同时学真实标签和 teacher 的预测分布。这样学生模型不光是模仿答案还顺带学会了 teacher 对“哪些像素区域对分类更关键”的偏好相当于把 CNN 的归纳偏置间接搬进了 transformer。这就是 DEiT 能拿更少数据逼近 ViT-400M 预训练效果的核心原因也是它和普通 ViT 在架构上最大的差异点多了一个 distillation token。我在实际项目里选 DEiT 还有一个很现实的理由它只有三四个直接可用的预训练尺寸tiny 版参数量不到 600 万一张 24GB 显存的卡能塞下十几个并行实验。对森林图像分类这种自定义任务先跑 tiny 找出数据增强和超参的规律再切 small 或 base 提精度这个阶梯用起来很顺手。对比 Swin TransformerDEiT 的预训练权重在 ImageNet 之外的小数据集上更“皮实”不容易出现一微调就掉点的玄学问题。2.2 环境准备用 PyTorch 和 timm 五分钟把模型装进手里常见做法是直接用 PyTorch 加 timmtimm 里已经内置了 deit_tiny_patch16_224、deit_small_patch16_224 和 deit_base_patch16_224 这几个标准结构。我一般会用 conda 建一个干净环境避免 torchvision 和 timm 互相打架。命令如下conda create -n deit python3.9 -y conda activate deit pip install torch torchvision --index-url https://download.pytorch.org/whl/cu118 pip install timm0.9.2 einops第一行创建 Python 3.9 的独立环境第二行激活第三行装 CUDA 11.8 版 PyTorch第四行装 timm 和 einops。einops 是 timm 内部依赖缺了会报ModuleNotFoundError: einops所以这里提前装好。不指定 CUDA 版本的话默认装到的可能是 CPU 版或者与显卡驱动不匹配的构建训练时device检查会是cpu速度直接差二十倍以上。装完可以用下面这行命令验证模型能不能正常构建python -c import timm; model timm.create_model(deit_tiny_patch16_224, pretrainedTrue, num_classes2); print(sum(p.numel() for p in model.parameters()))如果能看到输出参数总量而不是报错说明环境和权重下载链路都通了。这里pretrainedTrue会去下载 ImageNet 预训练权重第一次跑需要联网之后 timm 会缓存到本地。如果公司内网离线需要提前把权重文件拷到~/.cache/huggingface/hub或~/.cache/torch的对应目录路径不对会在加载时报No such file or directory。2.3 把“森林图像分类”做成自己的数据集图像分类数据集下载后最常见的形式是按类别分目录存放。我以森林场景二分类为例目录结构通常是forest_dataset/ ├── train/ │ ├── Forest/ │ │ ├── 001.jpg │ │ └── ... │ └── NonForest/ │ └── ... └── val/ ├── Forest/ └── NonForest/读取这种结构不需要自己写复杂的解析逻辑直接用torchvision.datasets.ImageFolder就能搞定。但如果你拿到的数据是 CSV 或者单个标签文件那就需要自定义 Dataset 类from torch.utils.data import Dataset from PIL import Image import pandas as pd import os class ForestCsvDataset(Dataset): def __init__(self, csv_path, img_root, transformNone): self.df pd.read_csv(csv_path) # 列: image_path,label self.img_root img_root self.transform transform # 把类别字符串映射成整数索引 self.classes sorted(self.df[label].unique()) self.label2id {c: i for i, c in enumerate(self.classes)} def __len__(self): return len(self.df) def __getitem__(self, idx): rel_path self.df.iloc[idx][image_path] label self.df.iloc[idx][label] img Image.open(os.path.join(self.img_root, rel_path)).convert(RGB) if self.transform: img self.transform(img) return img, self.label2id[label]这里有两个参数必须注意convert(RGB)是为了统一通道数很多相机拍的图片是 RGBA 四通道不转换会在进模型时维度报错label2id的映射顺序用sorted固定下来保证训练和推理时的标签编号一致。用 CSV 管理标签比纯目录结构多一步读取操作但后续做分层采样、按类别统计样本数会方便很多。数据增强方面我一般会做Resize(256) CenterCrop(224)再叠加随机翻转和光度抖动。DEiT 预训练时用的是 224×224 输入如果你用 384 或其他尺寸需要同步调整 Patch Embedding 的投影矩阵否则预训练权重加载时会遇到形状冲突。对小数据集RandomResizedCrop(224, scale(0.6, 1.0))比固定中心裁剪更能提升模型对目标尺度变化的适应能力。3. 用 DEiT 训一个自定义图像分类模型训练脚本与参数解读3.1 从 ViT 到 DEiT理解三个 token 与蒸馏损失的设计ViT 的输入序列里有一个 class token用来汇总整张图的信息并给分类头输出预测。DEiT 在这个基础上多塞了一个 distillation token它不参与最终分类而是在训练时承担“跟 teacher 学习”的通道。前向传播时图片经过 Patch Embedding 切成 16×16 的 patch tokenclass token 和 distillation token 各占序列开头一个位置三层 token 一起进入 Transformer Encoder。训练时 loss 由两部分组成真实标签的交叉熵加上 teacher 预测结果的蒸馏损失。硬蒸馏hard distillation把 teacher 的 argmax 当作一个额外标签和真实标签一起算交叉熵软蒸馏soft distillation则用 KL 散度去匹配 teacher 输出的概率分布。硬蒸馏实现简单、收敛快适合类别间边界清晰的任务软蒸馏保留了 teacher 对“相似类别”的判断信息在森林图像分类这种类别间存在大量纹理重叠的场景下精度通常比硬蒸馏高一个点左右。我习惯把蒸馏 loss 的权重 alpha 设在 0.5也就是说真实标签和 teacher 各占一半监督信号。alpha 太低模型退化成普通 ViT小数据过拟合风险回来alpha 太高模型会忽略真实标签跟着 teacher 的错误走。这个取舍没有标准答案只能在小验证集上试。3.2 训练脚本硬蒸馏、软蒸馏与正则化参数下面是一个可以直接改路径就开跑的简化训练脚本核心代码没有做额外的包装方便对照参数排查问题import torch import torch.nn as nn from timm.models import create_model from torch.utils.data import DataLoader from torchvision import transforms, datasets def build_teacher(num_classes): # RegNetY-160 是 DEiT 论文里推荐的 teacher teacher create_model(regnety_160, pretrainedTrue, num_classesnum_classes) teacher.eval() for p in teacher.parameters(): p.requires_grad False return teacher def train_one_epoch(model, teacher, loader, optimizer, criterion, device, alpha0.5): model.train() total_loss 0 for images, labels in loader: images, labels images.to(device), labels.to(device) optimizer.zero_grad() # 学生模型前向 logits model(images) # teacher 模型前向只做推断 with torch.no_grad(): t_logits teacher(images) # 真实标签损失 loss_ce criterion(logits, labels) # 蒸馏损失这里用硬标签的方式 t_labels t_logits.argmax(dim1) loss_distill criterion(logits, t_labels) loss alpha * loss_ce (1 - alpha) * loss_distill loss.backward() optimizer.step() total_loss loss.item() return total_loss / len(loader) epochs 50 batch_size 64 device cuda if torch.cuda.is_available() else cpu transform transforms.Compose([ transforms.RandomResizedCrop(224, scale(0.6, 1.0)), transforms.RandomHorizontalFlip(), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]), ]) train_ds datasets.ImageFolder(forest_dataset/train, transformtransform) train_loader DataLoader(train_ds, batch_sizebatch_size, shuffleTrue, num_workers4) model create_model(deit_tiny_patch16_224, pretrainedTrue, num_classes2).to(device) teacher build_teacher(2).to(device) optimizer torch.optim.AdamW(model.parameters(), lr5e-4, weight_decay0.05) criterion nn.CrossEntropyLoss(label_smoothing0.1) for epoch in range(epochs): loss train_one_epoch(model, teacher, train_loader, optimizer, criterion, device) print(fepoch {epoch 1}/{epochs}, loss: {loss:.4f})这段脚本里的logits model(images)在 timm 0.9.2 的默认配置下取的是 class token 经过分类头之后的结果尺寸是(batch, num_classes)。如果你用官方 DEiT 仓库的模型定义输出可能是(batch, 2, num_classes)或者把 class 和 distillation token 拼接需要拆开处理这是迁移代码时最容易翻车的地方。所有超参里lr5e-4是一个相对保守的起点。DEiT 官方在 ImageNet 上用的是5e-4 * batch_size / 256的线性缩放也就是说 batch 从 256 改成 64学习率大约降到 1.25e-4。我实际做小数据集时发现这个线性缩放规则在 batch 很小时会过度压低学习率导致模型五六轮都不见 loss 下降所以更常用的是固定 5e-4只调 warmup。weight_decay0.05是 ViT 系模型的标配比 CNN 常用的 1e-4 大很多主要用来约束 attention 矩阵里的冗余连接。label_smoothing0.1能让类别 logit 不过度自信对防止过拟合有直接帮助。3.3 三个必调参数与一组能直接开跑的默认值很多人拿到 DEiT 直接套用自己以前训练 ResNet 的参数发现精度不升反降问题几乎都出在这三个参数上。第一个是drop_path。DEiT 里它的别名是drop_path_ratetimm 创建模型时可以通过drop_path_rate传入。这个随机丢弃的是整条残差路径而不是单个神经元它能显著提升 transformer 在小数据上的泛化能力。我的经验值tiny 用 0.1small 用 0.2base 用 0.3。设置过高会让模型欠拟合训练 loss 降得很慢这个“玄学”参数在实践中比 dropout 更值得先调。第二个是 warmup epoch。transformer 的优化器是 AdamW它在训练初期如果没有 warmup梯度更新幅度容易被较大学习率带偏导致 loss 一开始就冲高。DEiT 官方训练 300 epoch 时用了 5 个 epoch 的 warmup。我建议如果只训 50 epochwarmup 用 5 到 10 epoch如果训 100 epoch 以上warmup 可以放到 10 到 20。判断原则是warmup 阶段结束后训练 loss 应该还在平稳下降而不是出现一个明显跳变。第三个是混合增强。DEiT 的官方强增强配置是 mixup 0.8、cutmix 1.0。在小数据集上直接套这么激进的增强容易让模型“看不清”原图验证集指标反而波动很大。森林图像分类里背景和目标的耦合本来就很强我通常把 mixup 压到 0.4、cutmix 设为 0.5。一组可以直接试跑的默认值大概是lr5e-4batch_size64drop_path_rate0.1mixup0.4cutmix0.5label_smoothing0.1distillation_alpha0.5。这个组合在 2 万张以内的自定义数据集上通常比裸 ViT 高出 2 到 4 个点。4. 用训练好的 DEiT 做预测与评估找出模型“真会了”还是“背下来”4.1 推理脚本class token 与 distillation token 怎么配合训练时蒸馏 token 参与了梯度更新但很多推理脚本只取 class token 的输出做预测这等于丢弃了模型一半的学习成果。我在实际代码里常见的做法是把 class token 和 distillation token 分别通过分类头得到两套 logits然后在 softmax 层面取平均。这样做的依据是两个 token 各自关注了不同的判别特征class token 偏向整体语义distillation token 偏向 teacher 迁移过来的局部纹理融合后对森林这类纹理密集场景更稳。一个可用的推理函数长这样def predict(model, image_tensor, device): model.eval() with torch.no_grad(): _, dist_logits model.forward_features(image_tensor.to(device)) cls_logits model.head(model.head_dropout(model.fc_norm(model.norm(model.patch_embed(image_tensor.to(device)))))) # 这里为了演示拆开两个 token 的预测分支 # timm 不同版本的结构字段名略有差异以实际模型为准 cls_prob torch.softmax(cls_logits, dim1) dist_prob torch.softmax(dist_logits, dim1) final_prob 0.5 * cls_prob 0.5 * dist_prob return final_prob上面的写法刻意暴露了一个问题DEiT 内部结构在不同版本的 timm 里字段名不完全一致直接按函数名一层层调用非常容易踩坑。更稳的通用做法是用 hook 从模型内部把两个 token 的输出接出来或者直接看model.forward的返回值。timm 在forward里默认只返回 class token 的 logits如果你不想拆内部结构就老老实实只用这个输出做预测精度差距大约在 0.5 到 1 个点倒也不会致命。项目上线要快时我不建议在这个问题上花太多时间先把 class token 的输出跑通后续再优化融合策略。4.2 评估指标在森林图像分类场景下精度够不够看了训练集准确率再高只能说明模型把训练样本记住了验证集才是判断泛化能力的准绳。我一般会算三个指标top-1 准确率、top-5 准确率和每类别的召回率。森林图像分类如果只有两个类别top-5 意义不大可以换成 F1 和混淆矩阵。用 sklearn 一行就能输出分类报告from sklearn.metrics import classification_report, confusion_matrix # val_preds 是概率矩阵val_labels 是整数标签 pred_ids val_preds.argmax(axis1) print(classification_report(val_labels, pred_ids, target_names[Forest, NonForest])) print(confusion_matrix(val_labels, pred_ids))这个报告的读法有个容易被忽略的地方如果 Forest 类样本数是 NonForest 的三倍整体准确率很容易虚高但 Forest 类召回率可能反而低因为模型倾向于把有难度的样本全部判成多数类。对这类数据不平衡问题训练时给损失函数加类别权重或者对少数类做过采样通常比调模型结构更有用。DEiT 的蒸馏机制会把 teacher 的偏差也带进来所以 teacher 在验证集上的表现也值得单独看teacher 都分不对的样本学生大概率也学不好这不是训不出来的问题是标签本身存在歧义。5. DEiT 实战踩坑记录五个常见问题的排查清单5.1 现象预训练权重加载报 shape mismatch加载deit_tiny_patch16_224的权重到自定义类别数时分类头的权重维度是 (1000, 192)而你定义的模型是num_classes2timm 会自动丢弃不匹配的分类头这时控制台会打印一条警告。很多人没注意以为预训练权重全部加载了结果模型实际上是从随机初始化的分类头开始训练导致前面若干 epoch 的 loss 偏高。解决方法是训练前打印一下model.load_state_dict的返回值确认missing_keys里只有 head 相关的键其他主干网络的键都完整加载。如果发现 patch_embed、blocks 这些层也有缺失要检查模型结构的 name 是否匹配比如deit_tiny_patch16_224和deit_tiny_distilled_patch16_224是不同的结构前者没有 distillation token后者有强行互换会大面积报键名对不上。5.2 现象在自建图片集上精度比 ResNet 还低最常见的原因是输入尺寸和预处理不对齐。DEiT 预训练时用的是 ImageNet 的mean[0.485, 0.456, 0.406]和std[0.229, 0.224, 0.225]如果你用自己数据集的均值和标准差去归一化模型看到的数据分布和预训练时完全不一致等于让预训练权重失效。另一个高频原因是图片内容本身和 ImageNet 差异太大比如红外图像、显微图像这种时候用 ImageNet 预训练权重的收益本来就有限不如直接在目标数据上从零训练或者先做一层简单的灰度直方图均衡化再进网络。5.3 现象训练 loss 下降但验证 loss 缓慢上升这是典型的过拟合但在 DEiT 上有一个容易被忽略的诱因数据增强太弱。transformer 的容量大如果只有随机翻转和裁剪模型很快能把训练集细节背下来。解决方法是先加 mixup 和 cutmix再加 RandAugment最后才考虑缩小模型尺寸。如果增强已经加上去了验证 loss 还是涨那就要看drop_path_rate是不是设成了 0这个参数在 timm 里默认就是 0很多人不知道要手动调。把drop_path_rate提到 0.1 到 0.2通常能压住过拟合比单纯加 weight decay 效果更明显。5.4 现象推理时把 distillation token 丢进全连接层报维度错误出现这个错误通常是你自己改了模型 forward想强制返回两个 token 的 logits但分类头的输入维度设计错了。distillation token 和 class token 的维度是一样的都是 192 或 384所以不是维度本身的问题而是拼接顺序错误。如果直接把两个 token 拼接成(batch, 2, 192)送到一个接受(batch, 192)输入的线性层PyTorch 会告诉你mat1 and mat2 shapes cannot be multiplied。解决方法是分别过分类头或者先把两个 token 在特征维度上求平均再进分类头。实在要拼接也要先reshape(batch, -1)再定义一个输入维度为 384 的线性层不要和默认的 head 混用。5.5 现象同一张图换增强方式结果差几个百分点这个现象在分辨率比较低的图上特别明显。DEiT 的 patch size 是 16×16224×224 的输入会被切成 14×14 的网格。如果原图只有 100×100缩放上去之后很多细节会糊掉class token 和 distillation token 都学不到可靠的判别特征。验证集和测试集如果用了不同的 resize 策略精度波动两个点以上是非常正常的。我一般会在项目一开始就固定一组评估用的 transform测试集用Resize(256) CenterCrop(224)训练集再单独用随机增强不要图省事让训练和测试共用一套逻辑。把图像分类数据集下载下来之后先抽几张图按这个流程跑一遍确认尺寸没有异常再开始大规模训练。6. 进阶技巧把 DEiT 的隐藏状态变成可解释的“后悔药”评估集精度只能告诉你模型对不对不能告诉你它为什么对。DEiT 的一个隐藏优势在于Transformer Encoder 每一层的注意力矩阵都能导出来用来定位模型到底看的是树的纹理还是背景天空。我用过一个很实用的排查方式把最后一层注意力图叠在原图上如果模型把森林图像分类成非森林而这个注意力图高亮的是天空区域那就说明训练数据里背景泄漏严重原始标签本身有问题。这个步骤相当于给模型一个“后悔药”不用重训就能发现数据清洗方向。抽取注意力图不需要改模型结构用 hook 最省事import torch attention_store {} def hook_fn(module, input, output): # output 形状通常为 (B, H, N, N)N 是 token 数 attention_store[last] output model.blocks[-1].attn.register_forward_hook(hook_fn) with torch.no_grad(): model(image_tensor.unsqueeze(0)) attn attention_store[last][0] # (H, N, N) cls_attn attn[:, 0, 1:] # class token 对所有 patch 的注意力这段代码里model.blocks[-1].attn是最后一个 Transformer Block 的注意力模块output在 timm 的部分版本里是线性变换后的注意力分数形状是(B, H, N, N)如果拿到的是元组一般取第一个元素。cls_attn的每个值代表 class token 对某个 patch 的注意力权重把它 reshape 成 14×14 再上采样到原图尺寸就能画出热力图。这个方法只适用于调试和验证阶段模型部署时不需要保留 hook。另外一条经验是调参顺序。很多人拿到 DEiT 就同时改学习率、蒸馏 alpha、mixup 和 drop_path出了问题根本不知道是谁导致的。我踩过几次这个坑后固定了一套流程先固定硬蒸馏加默认增强跑出基线然后只切软蒸馏对比变化再逐步增加 mixup 和 drop_path每一步只看一个变量。如果某个改动让精度掉了一个点以上就回滚到上一个配置并用小验证集上的混淆矩阵看具体是哪个类别变差了。这样听起来慢实际比盲目调参省时间得多至少你手里的每一版模型都能说清楚它为什么好用或者为什么翻车。DEiT 在自定义图像分类任务上的价值不在于它能把精度推到排行榜第一而在于它在中小数据量下给了你一个从 CNN 平滑切换到 transformer 的通道。先用 tiny 把数据链路跑通再用 small 或 base 提上限配合蒸馏 token 的融合输出和注意力可视化这套组合足够应对绝大多数落地场景。希望这篇能帮你少走几步弯路把时间花在真正影响结果的参数和数据处理上。本文还有配套的精品资源点击获取
上一篇/下一篇内容由系统自动关联 返回资讯列表 →