基于ViT的图像分类系统实现:从Patch Embedding到部署实践
简介一套基于视觉Transformer的图像分类系统实现面向需要部署前沿深度学习模型的开发者与研究者解决了从零搭建ViT训练流程的痛点系统采用ViT-B/16预训练模型支持自定义类别数涵盖数据增强、模型构建、训练评估与可视化四大模块。资源包共九个文件以四个Python脚本为主体辅以三个pyc缓存文件、一个txt说明文件和一个docx项目文档压缩后仅二十七KB轻量而完整。目前已有五十八人学习适合快速掌握ViT图像分类的核心流程。内容上既包含随机裁剪、水平翻转、颜色抖动等专业的数据增强策略也提供交叉熵损失、Adam优化器以及准确率、精确率、召回率、F1分数和特异度等全面评估指标同时支持GPU加速、tqdm进度显示和训练曲线可视化并针对中文环境做了优化。代码模块解耦良好读者可依据项目文档快速替换数据集并迁移到其他分类任务。1. 把ViT从论文变成可落地的图像分类系统为什么它值得跑一遍很多人第一次接触ViT是被“Transformer也能做图像分类”这个结论吸引的。CNN统治视觉多年ViT却把一张图切成patch序列直接送进多头自注意力里做全局建模。我用ViT跑森林图像分类任务时遇到的第一个反直觉现象是同样一份数据ViT收敛比ResNet慢但一旦训练到位对纹理复杂、背景干扰大的场景精度上限明显更高。这套“基于ViT的图像分类系统实现”要做的就是把ViT模型从论文里的结构图变成一个能训练、能评估、能部署的完整流程包括patch embedding、位置编码、CLS token、训练超参与推理验证。适合两类人想从CNN路线切到transformer图像分类路线的算法工程师以及需要拿ViT做课程设计或开源项目复现的学生。它能让你少走弯路直接看到一条主流技术路线跑通的全过程。2. 数据与结构准备搞清楚ViT的输入脾气和数据集格式再动手2.1 ViT的核心结构拆解Patch Embedding、位置编码和CLS TokenViT的整体逻辑可以压缩成一句话把图片变成一串向量再用Transformer编码这串向量最后取一个特殊位置的输出做分类。这里面有几个关键设计决定了它对输入格式的要求和CNN完全不同。CNN天然假设邻近像素相关所以用卷积核在局部滑ViT没有这个先验它把一张224x224的图按16x16的patch切分得到196个patch每个patch展平后线性投影成一个768维的向量。这个步骤叫Patch Embedding对应代码里通常是一个kernel和stride都等于patch_size的卷积。位置编码解决的是“序列顺序”问题。自注意力本身是无序的196个patch如果不加位置信息模型看到的就只是一个集合。标准做法是初始化一个可学习的position embedding维度是(197, 768)——多出来的1是CLS token占的位置。CLS token是拼在最前面的一个特殊向量训练时它通过自注意力不断聚合全图信息最后分类头只吃这个token的输出。这也是vit结构里最常被问到的点为什么不把196个patch的输出平均一下再分类实践上CLS token让模型自己决定“哪些位置的信息值得汇总”效果通常优于简单平均。子模块作用关键参数Patch Embedding把图像切成patch并投影到embedding空间patch_size16embed_dim768CLS Token聚合全局信息作为分类特征1个可学习向量位置编码保留patch的空间顺序可学习参数形状(197, 768)Transformer Encoder多头自注意力 MLP LayerNorm 残差depth12num_heads12分类头把CLS token映射到类别概率Linear(768, num_classes)这就是vit主流技术路线里的标准骨架也是后面所有实验的基础。理解这一层之后再去看timm库里的vit_base_patch16_224实现会发现代码结构和这张表一一对应。2.2 数据集整理目录结构、标签映射与DataLoaderViT对数据集的组织方式没有特殊要求和CNN分类任务一样一个子目录一个类别即可。但有一个细节值得提前固定下来类别名必须排序后建立索引。因为在Linux上os.listdir的顺序不一定是字典序万一训练和推理两次运行读出来的类别顺序不一致后面排查会非常痛苦。我一般会在数据集类里直接sort并把这个映射保存成JSON供推理阶段复用。from torch.utils.data import Dataset from PIL import Image import os class ForestImageDataset(Dataset): def __init__(self, root_dir, transformNone): self.root_dir root_dir self.classes sorted(os.listdir(root_dir)) # 固定顺序 self.class_to_idx {cls: i for i, cls in enumerate(self.classes)} self.samples [] for cls in self.classes: cls_dir os.path.join(root_dir, cls) for fname in os.listdir(cls_dir): self.samples.append((os.path.join(cls_dir, fname), self.class_to_idx[cls])) self.transform transform def __len__(self): return len(self.samples) def __getitem__(self, idx): path, label self.samples[idx] img Image.open(path).convert(RGB) if self.transform: img self.transform(img) return img, label这段代码的核心逻辑是构造一个样本列表每个元素是(图片路径, 标签索引)。sorted()保证了class_to_idx在每次运行时结果一致这是避免推理阶段类别错位的第一个保险。图片统一用convert(RGB)处理避免灰度图或带透明通道的PNG导致通道数不一致。2.3 数据增强的取舍Resize、RandomCrop和CutMix怎么组合ViT在中小规模数据集上比CNN更容易过拟合这是它缺少卷积归纳偏置的代价。数据增强在这里不是可选项而是影响模型能不能用的关键变量。我常用的组合是训练阶段用RandomResizedCrop配合翻转和颜色抖动验证阶段只做Resize加CenterCrop保证评估结果可比较。from torchvision import transforms train_transform transforms.Compose([ transforms.RandomResizedCrop(224, scale(0.7, 1.0)), transforms.RandomHorizontalFlip(), transforms.RandomRotation(15), transforms.ColorJitter(brightness0.3, contrast0.3, saturation0.3), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]), ]) val_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]), ])RandomResizedCrop的scale范围我习惯设成(0.7, 1.0)比ImageNet常用的(0.08, 1.0)保守一些。原因是对森林图像这类目标尺度较为集中的数据裁剪太狠会切掉主体反而引入噪声。Normalize用的mean和std是ImageNet统计值如果换用专有数据集可以在训练前先算一次全量数据的均值和标准差替换这两个参数收敛会更快一点。增强策略过拟合风险验证精度趋势适用场景仅Resize到224高先升后降数据量极大Resize 翻转中平稳小数据集兜底RandomResizedCrop全家桶低稳步上升中等规模数据MixUp / CutMix低收敛慢但上限高大规模数据CutMix这类方法对ViT有效但它会让训练时间明显变长前期的loss曲线也更难看。如果数据集只有几千张先把RandomResizedCrop这组基础增强用好比盲目上CutMix更划算。3. 模型搭建与训练写一个能收敛的ViT分类器3.1 核心模型代码从PatchEmbed到Transformer Encoder不依赖timm从零写一个可训练的ViT分类器是理解这个模型最直接的方式。下面这个实现是图像分类算法里最常见的vit结构写法包含Patch Embedding、位置编码、Transformer Encoder层和分类头四个部分和timm库中vit_base_patch16_224的配置对齐。import torch import torch.nn as nn class PatchEmbed(nn.Module): def __init__(self, img_size224, patch_size16, in_chans3, embed_dim768): super().__init__() self.num_patches (img_size // patch_size) ** 2 self.proj nn.Conv2d(in_chans, embed_dim, kernel_sizepatch_size, stridepatch_size) def forward(self, x): x self.proj(x) # B, embed_dim, H/16, W/16 x x.flatten(2) # B, embed_dim, num_patches x x.transpose(1, 2) # B, num_patches, embed_dim return x class TransformerEncoderLayer(nn.Module): def __init__(self, embed_dim768, num_heads12, mlp_ratio4.0, dropout0.1): super().__init__() self.norm1 nn.LayerNorm(embed_dim) self.attn nn.MultiheadAttention(embed_dim, num_heads, dropoutdropout, batch_firstTrue) self.norm2 nn.LayerNorm(embed_dim) self.mlp nn.Sequential( nn.Linear(embed_dim, int(embed_dim * mlp_ratio)), nn.GELU(), nn.Dropout(dropout), nn.Linear(int(embed_dim * mlp_ratio), embed_dim), nn.Dropout(dropout), ) def forward(self, x): x x self.attn(self.norm1(x), self.norm1(x), self.norm1(x))[0] x x self.mlp(self.norm2(x)) return x class ViTClassifier(nn.Module): def __init__(self, img_size224, patch_size16, in_chans3, num_classes10, embed_dim768, depth12, num_heads12, dropout0.1): super().__init__() self.patch_embed PatchEmbed(img_size, patch_size, in_chans, embed_dim) num_patches self.patch_embed.num_patches self.cls_token nn.Parameter(torch.zeros(1, 1, embed_dim)) self.pos_embed nn.Parameter(torch.zeros(1, num_patches 1, embed_dim)) self.pos_drop nn.Dropout(dropout) self.blocks nn.Sequential(*[ TransformerEncoderLayer(embed_dim, num_heads, dropoutdropout) for _ in range(depth) ]) self.norm nn.LayerNorm(embed_dim) self.head nn.Linear(embed_dim, num_classes) def forward(self, x): B x.shape[0] x self.patch_embed(x) cls_tokens self.cls_token.expand(B, -1, -1) x torch.cat([cls_tokens, x], dim1) # B, N1, embed_dim x x self.pos_embed x self.pos_drop(x) x self.blocks(x) x self.norm(x) x x[:, 0] # 只取CLS token x self.head(x) return x几点逻辑值得说明。PatchEmbed用步长等于kernel的卷积实现patch切分输出形状是(B, embed_dim, 14, 14)flatten和transpose之后变成(B, 196, 768)正好是Transformer期望的序列格式。CLS token用torch.zeros初始化而不是随机初始化这是因为位置编码和embedding矩阵都是随机初始化的CLS从一个中性起点开始学习会更平稳。整个模块序列里depth12表示12层Transformer Encodernum_heads12表示每层12个注意力头这套参数组合是vit结构里最经典的一套标准配置。参数名默认值含义选型建议patch_size16每个patch的边长8更精细但计算量大32太粗embed_dim768向量维度小模型可降到384depth12Encoder层数小数据集用6层即可num_heads12注意力头数一般整除embed_dimdropout0.1全连接层失活比例过拟合时调到0.23.2 训练超参lr、batch size、warmup和weight decay怎么配ViT对超参的敏感度比CNN高AdamW是默认优化器。实践里我踩过一个坑直接把CNN任务里好用的1e-3学习率搬到ViT上loss前几个epoch直接飞到NaN。ViT的自注意力模块在训练初期方差很大学习率一高就发散所以业界普遍的做法是3e-4起步并配合warmup。下面这套配置是目前vit模型训练的常见做法在中小规模数据集上可以直接套用。import torch.optim as optim from torch.optim.lr_scheduler import CosineAnnealingLR, LinearLR, SequentialLR model ViTClassifier(img_size224, num_classeslen(train_dataset.classes)) optimizer optim.AdamW(model.parameters(), lr3e-4, weight_decay0.05) epochs 100 warmup_epochs 10 warmup LinearLR(optimizer, start_factor0.01, end_factor1.0, total_iterswarmup_epochs) cosine CosineAnnealingLR(optimizer, T_maxepochs - warmup_epochs) scheduler SequentialLR(optimizer, schedulers[warmup, cosine], milestones[warmup_epochs]) for epoch in range(epochs): model.train() for imgs, labels in train_loader: imgs, labels imgs.to(device), labels.to(device) optimizer.zero_grad() logits model(imgs) loss nn.functional.cross_entropy(logits, labels) loss.backward() optimizer.step() scheduler.step()warmup的作用是让学习率从0.01倍慢慢升到目标值前10个epoch相当于给模型一个“热身期”等自注意力的梯度方差稳定后再全力优化。CosineAnnealingLR负责把后半程学习率按余弦曲线降到接近0后期的小学习率对收敛到平滑的极小值点有帮助。SequentialLR在milestones指定的epoch处自动从warmup切换到cosine不需要手动干预。超参推荐值调节方向optimizerAdamW不用SGDbase_lr3e-4batch增大时按比例放大weight_decay0.05过拟合时升到0.1batch_size64或128显存够就取大warmup_epochs总epochs的10%数据集越小比例越高label_smoothing0.1类别多时建议开启一个很实用的缩放规则是线性缩放学习率batch size从256变成64学习率也要从3e-4降到约0.75e-4。保持恒定的lr/batch比值比单独调lr更稳。3.3 训练过程的观察方法loss曲线与验证精度怎么读ViT训练初期的loss曲线比CNN“难看”。CNN通常前几个epoch就有明显下降ViT在patch数量多、数据量小的情况下loss可能要5到10个epoch才进入快速下降通道。这不一定是代码错了而是自注意力需要先学会“看哪里”。判断训练是否正常光盯loss不够要同时看验证集Top-1精度。如果训练loss稳步下降但验证精度长时间不涨说明模型在死记训练集这是过拟合的前兆。应对顺序是先加大weight_decay到0.1再把dropout从0.1调到0.2最后才考虑减少训练epoch。反过来的情况也要注意验证精度涨但训练loss降不下去通常是数据标签噪声大或增强过强需要检查数据集里有没有标注错误的样本。4. 评估与可视化用混淆矩阵和注意力热力图看模型学到了什么4.1 评估指标Top-1/Top-5准确率与混淆矩阵训练完成后先别急着部署跑一遍评估脚本把Top-1、Top-5和混淆矩阵一次算齐。Top-1是预测概率最大的类别和真实标签一致的比例Top-5是真实标签落在前5个预测里的比例。在多类别细粒度分类场景下Top-5能反映模型“虽然没有答对但候选范围已经缩小”的能力。def evaluate(model, val_loader, device): model.eval() correct_1, correct_5, total 0, 0, 0 all_preds, all_labels [], [] with torch.no_grad(): for imgs, labels in val_loader: imgs, labels imgs.to(device), labels.to(device) logits model(imgs) pred_1 logits.argmax(dim1) pred_5 logits.topk(5, dim1).indices correct_1 (pred_1 labels).sum().item() correct_5 (pred_5 labels.view(-1, 1)).sum().item() total labels.size(0) all_preds.extend(pred_1.cpu().tolist()) all_labels.extend(labels.cpu().tolist()) return correct_1 / total, correct_5 / total, all_preds, all_labelsTop-5计算里labels.view(-1, 1)是关键pred_5的形状是(B, 5)labels需要变成(B, 1)才能广播比较直接比较会得到全False。这个细节我第一次写时翻过车算出的Top-5精度是0排查半天才发现是形状问题。混淆矩阵用sklearn计算它能告诉你的信息比一个准确率数字多得多。from sklearn.metrics import confusion_matrix, ConfusionMatrixDisplay import matplotlib.pyplot as plt cm confusion_matrix(all_labels, all_preds) disp ConfusionMatrixDisplay(cm, display_labelstrain_dataset.classes) disp.plot(cmapBlues, xticks_rotation45) plt.savefig(confusion_matrix.png, dpi150)看混淆矩阵时重点观察两件事第一哪些类被系统性误判成另一个类比如森林图像里“落叶林”和“针叶林”互相混淆说明这两类在纹理上确实接近需要考虑增加类别内部的细粒度标注第二有没有某一类几乎不被预测说明它的样本量太少或特征被其他类覆盖了。这两条比总精度更能指导下一步迭代。4.2 注意力可视化把CLS Token的注意力变成热力图ViT的可解释性来自注意力权重。CLS token在每一层都会对整张图的196个patch做注意力加权把这些权重还原成热力图叠加在原图上就能直观看到模型分类时“看”的是哪里。实现上要注意一点nn.MultiheadAttention默认在调用时不返回注意力权重需要自己控制计算。class ViTWithAttn(ViTClassifier): def forward_with_attn(self, x): B x.shape[0] x self.patch_embed(x) cls_tokens self.cls_token.expand(B, -1, -1) x torch.cat([cls_tokens, x], dim1) x x self.pos_embed x self.pos_drop(x) attn_list [] for block in self.blocks: normed block.norm1(x) attn_out, attn_w block.attn(normed, normed, normed, average_attn_weightsFalse) x x attn_out x x block.mlp(block.norm2(x)) attn_list.append(attn_w) x self.norm(x) return x, attn_listaverage_attn_weightsFalse让返回的attn_w保持(B, num_heads, N1, N1)的形状这样可以看到每个头各自的注意力分布。可视化时取某一层某个头在CLS位置上的权重attn_w[0, head_index, 0, 1:]形状是(196,)reshape成(14, 14)再用bilinear插值放大到和原图相同尺寸叠加成热力图。我一般看最后一层第一个头的效果这个位置的注意力通常已经比较聚焦。如果热力图显示模型在背景上花了大面积注意力说明数据增强里的RandomResizedCrop裁剪范围太保守模型没学会忽略无关区域。4.3 单张图片推理与模型保存评估完之后保存模型和做单张推理是两个高频动作。保存模型时不要只存state_dict把class_to_idx映射一起存成JSON这个习惯能省掉后续无数排查时间。import json torch.save(model.state_dict(), vit_forest.pt) with open(class_to_idx.json, w) as f: json.dump(train_dataset.class_to_idx, f, indent2) def predict_one(model, img_path, class_names, device, transformval_transform): img Image.open(img_path).convert(RGB) x transform(img).unsqueeze(0).to(device) with torch.no_grad(): logits model(x) prob torch.softmax(logits, dim1) idx logits.argmax(dim1).item() return class_names[idx], prob[0, idx].item()推理时有一个常见坑加载模型后如果重新用os.listdir生成类别列表顺序可能和训练时不一样导致输出类别名错位。正确做法永远是读回训练时保存的class_to_idx.json用它来映射预测索引。我曾经在这上面吃过一次亏模型精度明明很高推理结果却全是乱的最后发现是两次运行读目录的顺序不同从那以后这个映射文件成了训练产物里的必备项。5. ViT实战避坑五个高频翻车点与排查记录5.1 现象一loss不降或者先降几个epoch然后直接变成NaN原因学习率太大是首要嫌疑。ViT里多头上注意力的梯度在初始阶段变化剧烈AdamW虽然能自适应调整但3e-4以上的学习率仍然容易让loss冲上NaN。另一个常见原因是位置编码和CLS token没有参与梯度计算比如误用了requires_gradFalse。解决把学习率降到1e-4到3e-4之间同时加warmup。训练循环里加梯度裁剪把梯度范数限制在1.0以内成本极低但能兜住大部分发散风险。torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm1.0)5.2 现象二训练精度很高验证精度却低得离谱原因这就是典型的过拟合而且ViT在小数据集上比CNN发生得更快。CNN靠卷积的局部性和权重共享天然抗过拟合ViT的全局注意力在数据少时会去“背”训练样本的位置组合。解决按顺序做三件事。第一把weight_decay从0.05加大到0.1第二把dropout从0.1调到0.2第三确认训练数据增强是否足够强尤其是RandomResizedCrop的scale范围。如果数据集只有几千张建议直接换用预训练权重做微调从头训练ViT在中小规模数据上性价比很低。5.3 现象三训练时显存溢出OOM原因ViT的激活值占用远高于同参数量的CNN。输入分辨率每增加一倍patch数量变成四倍注意力矩阵就是二次方膨胀。224x224输入算196个patch384x384输入就变成576个patch注意力矩阵占用直接翻了约8倍。解决先减batch_size再把训练改成混合精度。AMP在A100上能明显省显存在消费级显卡上至少能省三分之一。from torch.cuda.amp import autocast, GradScaler scaler GradScaler() for imgs, labels in train_loader: optimizer.zero_grad() with autocast(): logits model(imgs) loss nn.functional.cross_entropy(logits, labels) scaler.scale(loss).backward() scaler.step(optimizer) scaler.update()5.4 现象四训练时精度正常部署推理时类别名对不上原因推理代码重新扫描目录生成了新的category顺序和训练时的映射不一致。这个坑在Windows上尤其容易出现因为os.listdir在Windows和Linux下的默认排序规则不同。解决训练结束后立刻把class_to_idx保存成JSON推理时只从JSON读映射不依赖任何目录扫描。with open(class_to_idx.json, r) as f: class_to_idx json.load(f) idx_to_class {v: k for k, v in class_to_idx.items()}5.5 现象五换用更高分辨率输入时模型直接报shape不匹配原因位置编码的序列长度是定死的。基于ViT的图像分类系统在初始化时pos_embed的形状由(224/16)^21确定换成384x384输入后patch数量变成(384/16)^2576位置编码长度对不上PyTorch会立刻报错。解决对位置编码做双线性插值把它从196个patch扩展到576个patch。注意cls_token对应的位置编码要单独拆出来不能也跟着插值。def resize_pos_embed(pos_embed, new_num_patches): old_num_patches pos_embed.shape[1] - 1 cls_pos pos_embed[:, :1] patch_pos pos_embed[:, 1:] # 1, old_num, D side int(old_num_patches ** 0.5) patch_pos patch_pos.transpose(1, 2).reshape(1, pos_embed.shape[-1], side, side) new_side int(new_num_patches ** 0.5) new_patch_pos torch.nn.functional.interpolate( patch_pos, size(new_side, new_side), modebicubic) new_patch_pos new_patch_pos.flatten(2).transpose(1, 2) return torch.cat([cls_pos, new_patch_pos], dim1)插值之后最好用低学习率微调20到30个epoch让位置编码适应新的分辨率分布。直接固定位置编码推理也能出结果但精度通常会掉1到2个百分点。6. 从训练到部署静态化导出与多尺度推理的两个实用习惯6.1 用TorchScript和ONNX固定模型输入训练完的PyTorch模型依赖训练环境部署时要么用TorchScript做静态化要么导出ONNX走推理引擎。ViT的Transformer结构在导出时比CNN更容易踩坑原因是nn.MultiheadAttention内部实现可能包含动态控制流。TorchScript导出用trace模式最稳它会用dummy输入把计算图固定下来。model.eval() model model.cpu() dummy torch.randn(1, 3, 224, 224) traced torch.jit.trace(model, dummy) traced.save(vit_forest.pt)ONNX导出时dynamic_axes只动batch维度不要动高和宽维度。位置编码存了固定patch数量动态H/W会让尺寸不匹配的问题在部署时爆发。torch.onnx.export( model.cpu(), dummy, vit_forest.onnx, input_names[image], output_names[logits], dynamic_axes{image: {0: batch}, logits: {0: batch}}, )6.2 固定输入尺寸与多尺度验证的习惯ViT在部署时通常固定输入尺寸常见做法是把图片最短边缩放到256再中心裁剪224。这个组合在大多数场景下比直接拉伸到224好因为拉伸会改变物体长宽比干扰patch的语义。如果目标场景里物体尺寸变化大可以在保存模型前做一个多尺度验证把验证集分别用224、256、288三种尺寸跑一遍看精度变化趋势。精度掉了说明模型对尺度敏感这时再回头调整数据增强里的RandomResizedCrop的scale范围。从那以后我每次换数据集都强制走一遍固定流程先排序类别并保存映射再固定增强策略训练完立刻跑全量验证集算混淆矩阵最后用三种尺寸做多尺度对比再导出。这套流程救过我很多次尤其是第5章里那几个翻车点每一步都是血泪经验换来的。希望帮到你。本文还有配套的精品资源点击获取