
简介面向Python图像分类课程设计与期末大作业提供一套基于Vision TransformerViT的花卉识别完整代码适合具有基础Python语法和深度学习概念的学生参考。项目按功能拆分为数据读取、模型构建、训练与分类预测等模块并配有清晰注释便于初学者理解与二次开发压缩包共7个文件以6个Python脚本为主涵盖数据集处理、ViT模型定义、训练流程与识别调用另含keep占位文件整体仅10KB代码结构清晰下载后可直接运行。目前已有383人学习特别适合需要快速搭建图像分类项目或借鉴高分课程设计思路的读者。通过这份代码可以获取完整工程结构、核心实现逻辑和端到端的运行方案并在此基础上扩展花卉种类、调整模型参数或更换数据集进行迁移实验。1. 用ViT做图像分类花卉识别为什么适合拿来做大作业和入门实践花卉识别是一个看起来简单、做起来却很诚实的图像分类场景公开数据集类别多同类样本之间角度、光照和品种差异大。直接用CNN从头训练要调的东西很多而基于ViTVision Transformer做图像分类借助ImageNet预训练权重做迁移学习往往在较少的训练轮次内就能跑到90%以上的验证准确率。这里说的“基于ViT来做图像分类”核心其实就三步把图像切成patch序列、过Transformer编码器、接一个分类头。对大作业和入门项目来说这套路代码量不大对显卡的要求也比从头训练一个Transformer低得多因为真正参与微调的参数远少于模型总量。它适合写过基础Python分类代码的从业者也适合第一次接触Transformer视觉任务的人。你只需要搞懂三件事ViT的结构怎么理解、数据怎么喂、参数怎么调才不翻车。2. ViT图像分类的结构原理从Patch到全局注意力2.1 ViT核心结构Patch Embedding、位置编码、Transformer编码层和CLS TokenViT的输入不是整张图的像素矩阵而是把图像切成固定大小的patch序列。常见配置是输入224×224patch size为16那么图像被切成(224/16)²196个不重叠的16×16小块。以ViT-Base/16为例hidden size是768每个小块展平后经过一个线性投影变成768维向量这一步叫Patch Embedding。它等价于一个stride等于kernel size的卷积timm内部就是这么实现的。import torch import torch.nn as nn class PatchEmbed(nn.Module): def __init__(self, in_ch3, patch16, embed_dim768): super().__init__() self.proj nn.Conv2d(in_ch, embed_dim, kernel_sizepatch, stridepatch) def forward(self, x): # x: [B, 3, 224, 224] x self.proj(x) # [B, 768, 14, 14] return x.flatten(2).transpose(1, 2) # [B, 196, 768]这个代码块里最关键的是stride和kernel_size都等于patch所有patch彼此不重叠空间信息先被压缩成14×14的网格。B是batch size196是序列长度。之所以用Conv2d实现而不是手动切片是因为卷积在GPU上有更高效的内存布局实际项目里没人真的去for循环切图。切完patch之后还需要两个东西。一个是位置编码因为Transformer本身没有空间顺序概念196个token如果不加位置信息模型分不清左上角和右下角的patch。ViT用的是可学习的一维位置编码加到所有patch token上这也是为什么输入尺寸尽量固定——换了输入尺寸位置编码就不匹配了。另一个是CLS Token序列开头额外拼接一个可学习向量让它在编码过程中聚合整张图的全局信息最后分类只用这个token的输出。真正干活的是一层层Transformer Encoder由多头自注意力、MLP、LayerNorm和残差连接组成。自注意力机制让每一层的任意一个patch都能看到其余195个patch这给了ViT天然的全局感受野。CNN要堆很深才能看到大范围上下文而ViT第一层就能做全图交互。花卉识别的很多特征恰恰是全局性的花瓣排列、花型轮廓、花蕊与花瓣的位置关系这些都依赖跨patch的信息融合。2.2 为什么选ViT而不是ResNet小样本花卉任务的选型理由不少人会质疑花卉数据集样本量不大ViT参数多不是更容易过拟合吗这个疑问放到“从头训练”的场景里是对的但实际代码包走的是迁移学习路线结论会反过来。预训练权重决定了ViT的起点。从ImageNet-21k或ImageNet-1k上预训练出来的ViT已经学会了大量纹理、边缘、形状和物体部件概念。花卉识别和ImageNet里的自然图像分布非常接近微调时只需要在最后分类头上做较大调整前面的Transformer层改动很小。这意味着有效参数量远小于模型总参数量。从结构上看ResNet等CNN模型具备平移等变性和局部先验在数据少的时候更稳但它的卷积核感受野受限高层语义依赖层叠堆出来的“视野”。花卉的判别信息往往不是局部的比如菊花和蒲公英的整体形态差异、玫瑰和月季的花型差异都需要跨区域对比。ViT的全局注意力在这个任务上天然占优公开数据集上微调后的结果ViT-Base/16通常能比ResNet-50高2到5个点。选型还要看生态成本。timm一行代码就能加载预训练ViT数据预处理接口和CNN模型完全一致训练代码只需要改模型名。和需要设计增强策略、调整卷积核的CNN方案相比ViT在代码维护上更省心。真正的代价是显存和推理速度ViT-Base/16在224分辨率下batch size 32大约要6到7GB显存而ResNet-50同样batch只需要一半左右。如果只有CPU或者老显卡建议直接用timm里更小的模型比如vit_small_patch16_224或者vit_base_patch16_224但把batch降到16。用一个简表来总结两种路线的差别对比项ViT-Base/16ResNet-50模型参数量约8600万约2500万感受野第一层就是全局靠堆层扩大小数据微调稳定性依赖预训练微调需小lr相对更稳同batch显存占用较高较低代码复杂度timm一行加载也一行加载花卉公开集上微调效果通常更高略低这个表不是让你无脑选ViT而是说明在“已有预训练权重”的前提下ViT值得为花卉识别付出额外的显存成本。如果数据集每类只有二三十张还建议用冻结backbone的线性探测策略后面第6章会展开讲。3. 花卉识别数据准备目录结构、增强策略与Dataloader配置3.1 数据集选型与目录组织按类别分文件夹是通用做法这个代码包里的花卉识别没有限定数据集最常见的配套选择是Oxford 102 Flowers公开集一共102类、大约8000多张图单类样本量从40到200多张不等分布不均衡但有挑战性。也有一部分人会直接拍自己身边的花来建数据集。无论哪一种目录结构都建议按“一个类别一个文件夹”来组织这在后续用torchvision的ImageFolder或者自写Dataset时都顺畅。常见目录结构如下flower_data/ ├── train/ │ ├── 0_bluebell/ │ ├── 1_buttercup/ │ ├── ... │ └── 101_water_lily/ └── val/ ├── 0_bluebell/ ├── 1_buttercup/ └── ...这里的类别目录名前缀了数字编号目的是让排序稳定避免出现“第0类”和“第10类”因为字符串排序顺序错乱的问题。如果你下载的数据集是图片全在一个文件夹、类别存在CSV里建议花几分钟把它重排成上面的格式后续所有训练代码都不需要再处理标签映射。拆分数据时我一般保持train和val比例在8:2左右。102类数据集如果原本没划分可以先把每类图片打乱取前80%进train后20%进val。不要用随机把整库切成两份的方式因为那样可能让某类在训练集被极度稀释。另外提醒一句训练集和验证集千万不能有相同图片这个看起来是废话但从网上下载的整理好的花卉数据集里有时会混入重复图建议跑完训练后抽查验证集里top-1错误的样本。3.2 数据增强与归一化ViT微调的标准Transform配置ViT不像CNN那样对输入尺寸有很强的柔性patch size固定为16timm里预训练权重默认吃224×224所以transform里最终输出尺寸必须是224。下面是这套代码包里常见的增强配置也是我实际用下来比较稳的一组参数from torchvision import transforms as T train_tf T.Compose([ T.RandomResizedCrop(224, scale(0.6, 1.0), ratio(0.75, 1.333)), T.RandomHorizontalFlip(p0.5), T.ColorJitter(brightness0.2, contrast0.2, saturation0.2), T.ToTensor(), T.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]) ]) val_tf T.Compose([ T.Resize(256), T.CenterCrop(224), T.ToTensor(), T.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]) ])RandomResizedCrop的scale参数值得多说一句。ImageNet的默认scale是(0.08, 1.0)那适合物体占比不确定的通用场景。花卉数据里花通常是画面主体占比往往超过一半如果裁到0.08很容易把花切掉只剩叶子模型会被迫从残片里猜类别。把scale下限提到0.6之后裁剪保留的主体更多实测收敛更快。如果数据集里花比较小比如远景拍摄的公园花海可以把下限降到0.3再试。Normalize必须用ImageNet统计的mean和std也就是[0.485, 0.456, 0.406]这一组。原因很简单预训练权重里的BatchNorm或LayerNorm早就适应了这个分布如果换成自己算的花卉数据集均值和方差输入分布偏移会让微调效果明显变差。不要在花卉数据上重新统计mean和std这是我踩过的坑省那点计算时间不值得。验证集用Resize(256)加CenterCrop(224)而不是直接Resize到224是因为直接缩放会把非正方形的图片拉伸变形CenterCrop保留中心区域更接近测试时的真实分布。图像先缩放成256再居中裁剪224也保留了少量上下文信息对花卉这种中心主体明确的图效果更好。3.3 类别读取与Dataloader避免标签顺序错位类别读取的关键是顺序必须全局一致。训练、验证、预测时如果用了不同的排序方式标签就会错位结果看起来准确率很高但实际分类全部偏移。通用做法是用排序后的目录列表作为唯一类别索引基准import os from torch.utils.data import Dataset, DataLoader from PIL import Image CLASSES sorted(os.listdir(flower_data/train)) NUM_CLASSES len(CLASSES) print(CLASSES, NUM_CLASSES) class FlowerDataset(Dataset): def __init__(self, root, transform): self.samples [] for cls_id, cls_name in enumerate(sorted(os.listdir(root))): cls_dir os.path.join(root, cls_name) for fname in os.listdir(cls_dir): self.samples.append((os.path.join(cls_dir, fname), cls_id)) 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) return self.transform(img), label train_loader DataLoader(FlowerDataset(flower_data/train, train_tf), batch_size32, shuffleTrue, num_workers4, pin_memoryTrue)这里有两个容易被忽略的细节。第一个是Image.open之后要加.convert(RGB)因为花卉图片里可能出现带透明通道的PNG或灰度图不转换的话后面ToTensor会得到4通道或1通道输入模型forward直接炸。第二个是DataLoader里num_workers不要贪多CPU核数一半左右比较稳妥。Windows系统上num_workers设成4以上容易报BrokenPipe错误如果碰到直接改成0是最省事的办法。pin_memoryTrue可以加速GPU拷贝但如果机器内存很紧张关了反而更稳。训练过程中如果发现loss震荡不收敛优先检查增强是否太强、lr是否太大而不是Dataset有没有写错。验证集loss比训练集高得离谱时再回头检查样本是否有重复、目录划分是否有泄漏。4. 用Python训练ViT花卉分类模型模型加载、超参数与训练循环4.1 用timm加载ViT预训练权重并替换分类头训练环境建议用Python 3.8以上版本安装torch和timm。torch的安装方式取决于你的CUDA版本这里不再展开重点说timm加载模型的写法。老版本代码包里常有人手动从huggingface下载权重再加载新一点的代码包基本都用timm因为它会把权重下载、patch对齐、分类头替换一次性做完。import timm import torch.nn as nn NUM_CLASSES 102 # 由前面CLASSES列表长度得到 model timm.create_model( vit_base_patch16_224, pretrainedTrue, num_classesNUM_CLASSES, ) model.train() print(model.head)create_model里的num_classes参数非常关键。如果你不传它模型默认保留原来ImageNet的1000类分类头训练时loss维度会和标签对不上传了102之后timm会自动把分类头换成新的全连接层。vit_base_patch16_224这个模型名拆开看vit是模型族base是规模patch16表示patch size为16224表示预训练时的输入分辨率。timm里还有vit_base_patch16_224_in21k它是在ImageNet-21k上预训练的特征更通用花卉这类和ImageNet分布接近的任务微调效果往往更好缺点是权重文件更大。如果加载之后还想手动改分类头常见写法是model.head nn.Linear(in_featuresmodel.embed_dim, out_featuresNUM_CLASSES)使用model.embed_dim而不是硬编码768是因为timm不同版本、不同模型的维度不同这么写能避免以后换模型时漏改。打印model.head确认最后一层输出维度等于类别数再开始训练这一步能省掉后面至少半小时的debug时间。4.2 完整训练循环损失函数、优化器、warmup与余弦退火训练ViT和训练CNN在代码层面几乎一样但有几个参数需要区别对待。优化器用AdamW而不是SGD这是ViT论文和timm默认配置的主流路线。微调阶段的学习率必须比从头训练小一个数量级常见范围在1e-4到3e-4之间超过5e-4大概率loss震荡。import torch from torch import nn from torch.optim import AdamW from torch.optim.lr_scheduler import LinearLR, CosineAnnealingLR, SequentialLR from torch.cuda.amp import autocast, GradScaler EPOCHS 30 lr 3e-4 model model.cuda() opt AdamW(model.parameters(), lrlr, weight_decay0.05) criterion nn.CrossEntropyLoss() iters_per_epoch len(train_loader) warmup_iters 2 * iters_per_epoch warmup LinearLR(opt, start_factor0.1, end_factor1.0, total_iterswarmup_iters) cosine CosineAnnealingLR(opt, T_max(EPOCHS - 2) * iters_per_epoch, eta_min1e-6) sched SequentialLR(opt, [warmup, cosine], milestones[warmup_iters]) scaler GradScaler() for ep in range(EPOCHS): model.train() running_loss 0.0 for X, y in train_loader: X, y X.cuda(), y.cuda() opt.zero_grad() with autocast(): logits model(X) loss criterion(logits, y) scaler.scale(loss).backward() scaler.step(opt) scaler.update() sched.step() running_loss loss.item() print(fepoch {ep1}/{EPOCHS} loss {running_loss / len(train_loader):.4f})这里有几个参数需要说明白。weight_decay用0.05而不是CNN常用的1e-4是因为ViT预训练阶段就是按比较大的weight decay训练的微调时保持一致更不容易破坏原有权重分布。warmup设两个epoch目的是让优化器在初期不要直接以满学习率更新否则预训练权重会被前几个batch冲坏。余弦退火的T_max覆盖剩余所有iter学习率会平滑地降到1e-6这个底线让模型在后期做精细调整。autocast和GradScaler是一对混合精度训练组合。混合精度在A卡和N卡上都可用能明显降低显存占用大约能省25%到30%速度也有提升。注意GradScaler的step和update必须在backward之后按顺序调用少一个都会让训练悄悄变慢甚至不收敛。如果你的环境比较老不支持amp可以去掉这两行把X、y转成float32其他逻辑不变。4.3 验证、混淆矩阵与模型保存策略训练循环之外还需要一个验证函数在每个epoch结束时评估当前模型并决定是否覆盖保存。花卉识别的评价指标最直观的是top-1准确率但102类里有一些外观非常接近的品种比如某种郁金香的不同变种可以顺带看top-5。验证时必须关掉梯度、切换成eval模式还要把autocast一起带上保证推理和训练时的数值路径一致。torch.no_grad() def evaluate(model, loader): model.eval() correct total 0 all_preds, all_labels [], [] for X, y in loader: X, y X.cuda(), y.cuda() with autocast(): logits model(X) pred logits.argmax(dim1) correct (pred y).sum().item() total y.size(0) all_preds.extend(pred.cpu().numpy()) all_labels.extend(y.cpu().numpy()) model.train() return correct / total, all_preds, all_labels acc, preds, labels evaluate(model, val_loader) print(fval acc: {acc:.4f})保存模型时不要每个epoch都存否则几十个epoch下来磁盘会堆满。我习惯用验证准确率作为门槛只保存最高分best_acc 0.0 # 放在验证之后 if acc best_acc: best_acc acc torch.save(model.state_dict(), fvit_flowers_{acc:.4f}.pth)保存state_dict而不是整个model文件更小加载时需要先build出同样的模型结构再load_state_dict。如果是为了快速部署或交作业演示直接torch.save(model, vit_full.pth)更省事但换环境加载时容易因为timm版本不一致出问题。训练完成后用sklearn的confusion_matrix打印几个混淆严重的类别能直观看出模型到底分不清什么这比只看一个准确率数字有用得多。5. ViT花卉识别避坑与排查5条血泪经验帮你省一天时间5.1 解压后路径多了一层目录训练直接FileNotFoundError现象按代码包默认路径运行报FileNotFoundError提示找不到flower_data/train目录。原因zip解压后文件夹往往嵌套了一层实际结构是某个外层目录/ flower_data/train而训练脚本默认从解压根目录开始找。这是代码包最常见的问题不是代码bug纯粹是路径层级不一致。解决先用一行命令看清目录层级find . -maxdepth 2 -type d | head -20然后把脚本里的ROOT改成实际数据所在路径或者在解压时把外层目录去掉。判断路径是否正确的标准是ROOT/train这个路径下直接是类别文件夹再下一层才是图片。如果train下面还有一层train把代码里的路径拼接改成“数据根目录/训练目录”时多写一层或手动调整目录结构都行。5.2 num_classes没传报错说shape不匹配现象训练第一个batch时loss报错提示logits是[32, 1000]而标签是[32]期望的维度对不上。原因create_model时没有传num_classes参数模型还保留着ImageNet的1000类分类头。花卉识别数据集类别数通常是102或17和1000不匹配。解决先确认类别数来源是从CLASSES列表动态计算的不要写死在代码里。比如NUM_CLASSES len(CLASSES)这样即使换数据集也不用改训练代码。加载模型后打印model.head.shape确认最后一维等于你的类别数再开训。timm里create_model传num_classes会自动替换head但仍然建议打印检查timm版本升级时行为偶尔有变化。5.3 训练loss不降或剧烈震荡明明模型和数据都没问题现象loss在0.5附近上下弹或者前几个epoch下降之后又开始回升验证准确率始终上不去。原因多数情况是学习率过大了。ViT微调不是从头训练预训练权重已经处于一个不错的局部最优附近学习率过大会直接把它推出原来的盆地。另一个常见原因是没有warmup第一批数据就以满学习率更新几个batch之后权重就被带偏。解决把学习率降到1e-4到3e-4区间并确保warmup覆盖前2个epoch左右。如果用的是线性探测或冻结部分层的方案学习率可以适当提到1e-3因为更新的参数少了。还有一个玄学但有效的小技巧如果loss卡在非常小的值不动比如0.2附近先把最后几层解冻只训分类头看有没有变化再逐步放开更多的层用渐进式解冻定位问题。5.4 显存OOMbatch size和混合精度的平衡现象训练到第二个epoch直接报CUDA out of memory或者开了某些数据增强后显存暴涨。原因ViT-Base/16在224分辨率下batch size 32大约需要6到7GB显存算上中间激活和梯度很多8GB卡会很吃力。增强操作里的RandomResizedCrop本身不占额外显存但更高分辨率或更大batch会成倍放大问题。解决最直接的办法是把batch size降到16确认单卡能跑通后再往上加。同时打开autocast混合精度能省大约25%到30%显存。如果降batch后影响收敛用梯度累积保持等效batch sizeaccum_steps 2 opt.zero_grad() for i, (X, y) in enumerate(train_loader): X, y X.cuda(), y.cuda() logits model(X) loss criterion(logits, y) / accum_steps loss.backward() if (i 1) % accum_steps 0: opt.step() opt.zero_grad()这个写法里把loss除以累积步数相当于每accum_steps次迭代才做一次参数更新。注意学习率不需要按等效batch size调整因为AdamW本身对batch size的敏感度不像SGD那么高。还有一个更彻底的办法换更小的模型比如vit_small_patch16_224或vit_base_patch32_224准确率会略降但显存压力小很多。5.5 训练准确率高但预测结果张冠李戴问题出在标签映射现象验证集准确率90%以上但随便拿一张图进去预测打印出的类别名明显不对。更诡异的是某些类别的错误是整体偏移一位的。原因训练、验证、推理三处用了不同的标签映射方式。比如训练时用sorted(os.listdir())得到索引推理时却手写了类别列表或者验证集和训练集的类别目录顺序不一致。字符串排序下0_bluebell和10_x的顺序和纯数字排序是不一样的这种错位很隐蔽。解决在项目里定义一个唯一的CLASSES变量所有Dataset、预测脚本、展示脚本都引用它。推理时只通过CLASSES做索引到名称的映射不另行维护列表。每次训练前打印每个训练样本的类别名和索引对应关系随机抽5张图人工看一遍再开始训练。这个检查只需要一分钟但能避免整轮训练白跑。6. 进阶思路线性探测与注意力可视化验证模型是否真的在看花6.1 用冻结backbone的方式做线性探测适合极小数据集如果花卉数据集只有几百张图微调整个模型容易过拟合一个稳妥的替代方案是冻结绝大多数层只训练最后的分类头。这种策略叫线性探测在ViT上尤其有效因为预训练特征的判别力已经很强很多任务只训一个线性层就能达到不错的效果。for name, param in model.named_parameters(): if head not in name: param.requires_grad False opt AdamW(filter(lambda p: p.requires_grad, model.parameters()), lr1e-3, weight_decay0.05)这里过滤掉了head以外的所有参数实际参与更新的只剩分类头那几十万个参数比全模型微调小两个数量级以上。学习率可以放宽到1e-3因为不更新backbone就不怕破坏预训练特征。等到线性探测跑通、确认数据没问题之后再放开backbone做全量微调往往比直接全量微调更稳。6.2 Attention Map可视化看模型关注哪里而不是只看准确率准确率是结果注意力可视化是过程证据。ViT的self-attention权重天然可以作为解释依据把最后一层CLS Token对其他patch的注意力取平均reshape回14×14就能看出模型在分类时关注了图像哪些区域model.eval() attention_map {} def hook_fn(module, inp, out): # timm不同版本输出格式略有差异这里打印结构确认后取attention矩阵 print(hook out length:, len(out)) for o in out: if isinstance(o, torch.Tensor): print(o.shape) handle model.blocks[-1].attn.register_forward_hook(hook_fn) with torch.no_grad(): logits model(img.unsqueeze(0).cuda()) handle.remove()更常见的做法是直接在forward里拿到attention权重但用hook的好处是不改动原模型前向逻辑。拿到权重后形状一般是[B, heads, seq, seq]其中seq等于197包含了1个CLS Token和196个patch。取CLS那行对所有头像求平均得到196个值reshape成14×14上采样到224再叠加到原图上就能生成类似热力图的输出。如果热力图集中在花瓣和花蕊上说明模型学到了正确的判别信息如果热力图大面积落在背景叶片或地面就要怀疑数据增强的裁剪策略是否让模型走偏了。我自己做花卉识别时最开始也只看accuracy后来发现有两类花训练集和验证集准确率都很高但注意力热力图全部聚焦在叶片纹理上其实是数据集里这两类花的背景差异太明显模型在偷懒按背景分类。从那以后我养成了一个习惯每次训练完至少随机抽20张验证集图片画注意力热力图和预测结果放在一起肉眼检查一遍。这个过程不复杂但能帮你发现准确率掩盖的很多真实问题。希望这些参数和组织方式能让你少走弯路真正把代码跑起来而不是停留在改bug上。本文还有配套的精品资源点击获取