ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

ViT花卉识别实战:基于Vision Transformer的图像分类迁移学习

ViT花卉识别实战:基于Vision Transformer的图像分类迁移学习 简介这份Python项目以ViTVision Transformer为骨干实现花卉图像分类任务面向计算机视觉课程设计、期末大作业及深度学习入门者。代码覆盖完整流程数据加载与预处理DataLoader.py、模型搭建model.py、分类逻辑DoClassify.py、工具函数utils.py与训练入口train.py并附config.py统一管理参数结构清晰便于二次开发。压缩包共7个文件以6个py脚本为主另含keep占位文件整体体积约10KB轻量易部署。项目由高分开发者在97分作业基础上整理注释详尽适合基础薄弱者逐行理解ViT的token嵌入、Transformer编码器及分类头原理当前已有383人学习也可作为模型调参、数据增强或迁移学习的改造蓝本。1. 大作业里的ViT花卉识别它到底在解什么题大作业python基于ViT来进行图像分类花卉识别本质上是一套“Vision Transformer 迁移学习”的完整图像分类链路。你拿到的不只是某个模型文件而是把图片数据组织成训练集、微调一个预训练好的ViT、再对真实花卉照片做推理的过程。这类项目最典型的场景是课设、毕业设计和入门Transformer视觉方向的第一份代码作业用小规模数据、常见公开花卉集在有限算力下跑出一个能达到90%以上验证精度的可演示模型。值得注意的是在整个图像分类算法体系里ViT从来不是“数据少也能正面硬刚”的选手。它真正的优势是先在ImageNet这种大规模数据集上预训练再用花卉这种小数据做微调。微调后它的泛化能力往往比从头训练CNN更好这也是为什么它成了近年最值得放进大作业里的最新图像分类模型之一。这篇笔记会从原理、数据准备、训练参数、避坑到推理验证给你一份能照着复现的完整方案。2. ViT原理与选型图像分类算法那么多为什么大作业选它2.1 ViT到底在做什么Patch嵌入、位置编码与CLS TokenViT的中文全称是Vision Transformer它把图像分类从“卷积扫描像素”变成了“序列建模”。一张224x224的图片如果按16x16像素切一个patch会得到14x14196个patch。每个patch通过一个线性层映射成向量的过程叫Patch Embedding相当于把一张图变成了196个token组成的序列。为了让模型知道patch在原始图像里的位置这里会加一组Position Embedding。它和NLP里的token embedding是一个思路模型不仅要看懂“每个patch是什么”还得知道“它在图片哪个位置”。所以ViT的输入实际上是196个patch向量再加一个CLS Token共197个token。CLS Token专门用来汇总全图信息训练结束后分类头只接在CLS Token的输出上。你可以理解为其他token负责描述局部内容CLS Token负责“综合百家意见投出最终分类”。Transformer Encoder由多层Self-Attention和MLP组成。Self-Attention让每个patch都能直接看全所有patch这一点和卷积的局部感受野截然不同。以ViT-B/16为例它有12层Transformer、12个注意力头隐藏维度768。处理一张224x224的图自注意力的计算复杂度是序列长度平方级的196个token还能接受再放大到更大分辨率就会显著变慢。因此ViT虽然能建模全局依赖却对输入尺寸和显存更敏感。2.2 对比ResNet和MobileNetV2ViT的优势与代价很多同学会问同样是花卉识别ResNet50、MobileNetV2不是更成熟吗如果你拿MobileNetV2代码去跑同样的数据集训练速度和部署成本确实更友好。但ViT的价值在于大规模预训练之后它拿到的泛化特征比CNN更平滑在迁移到下游小数据集时往往能更快收敛。下面这个对比能帮你理解选型模型感受野相同数据量下从头训练ImageNet预训练后微调显存与推理成本ResNet50局部叠加的全局感受野相对稳定易过拟合但可控好中MobileNetV2局部感受野更快适合移动端好低ViT-B/16从头到尾的全局注意力容易欠拟合数据少时效果差通常更好高ViT-B/16 迁移学习全局注意力不建议花卉5~10类能达到高精度中高如果你在课设里遇到老师问“为什么不用CNN”答这条就够了ViT的假设空间更灵活但需要足够数据或预训练才能约束住有了预训练它在大作业级别的数据上一样能打。代价是训练时GPU更吃力batch size不能开太大。这也是为什么很多项目会配合混合精度和梯度累积来跑ViT。2.3 项目选型思路用timm加载预训练ViT是常见做法常见做法是用timm库创建ViT模型它封装了ImageNet预训练权重一行代码就能拿到一个可以直接微调的ViT-B/16。timm内部会自动下载预训练权重对PyTorch生态的兼容性也最好。你也可以选择Hugging Face的transformers库它同样支持ViT并且更容易拿到注意力矩阵适合后面做可视化。对花卉识别这类下游任务通常不会直接用ImageNet-1k的权重而是用ImageNet-21k预训练权重因为21k覆盖的类别更多提取出的特征更通用。timm里对应的模型名是“vit_base_patch16_224_in21k”输入分辨率224patch size 16。如果你的数据是细粒度花卉还可以考虑更大输入分辨率如384不过显存和训练时间都会上涨。3. 环境准备与数据集组织先让代码跑起来再谈效果3.1 从Python安装到依赖torch、torchvision、timm版本选择这套项目建议使用Python 3.8以上版本。对Windows用户第一次装PyTorch遇到最多的问题不是pip命令拼错而是底层C运行库缺失。如果你在import torch时报错“由于找不到msvcp140.dll无法继续执行代码”必须先安装Microsoft C Build Tools或Visual C Redistributable再重装torch。这属于环境前置问题跟ViT本身无关但很容易卡住新手。依赖安装建议用pip一次性装完pip install torch torchvision timm tensorboard说明torch和torchvision必须版本对应最好使用PyTorch官网生成的安装命令避免混装导致CPU/GPU版本不匹配。timm建议使用0.9.x版本这个版本里timm.create_model(..., pretrainedTrue)的参数行为比较稳定。如果你后续要用AMP混合精度还需要确认PyTorch版本不低于1.10。这里有个小小的坑timm版本升级很快旧代码里model.head nn.Linear(...)这种写法在0.9.x依然有效但更推荐用num_classes参数创建模型。因为timm会自动把分类头替换成新尺寸省去你手动适配的麻烦。3.2 目录结构与标签映射Train / Valid / Test 三分法花卉数据集按类别建子文件夹是最常见的组织方式也是PyTorchImageFolder能直接识别的结构flowers_split/ ├── train/ │ ├── daisy/ │ ├── dandelion/ │ ├── rose/ │ ├── sunflower/ │ └── tulip/ ├── valid/ │ └── ... └── test/ └── ...三位划分的理由很明确训练集调权重验证集调超参数和选模型测试集只做最终评估。很多大作业把train和valid合并最后报告精度时用同一个集虽然数字好看但无法反映真实泛化能力。你在答辩时被问“测试集怎么来的”就会很被动。所以数据集组织这一步必须单独做。注意类名尽量不要用中文和空格。因为Windows文件系统虽然支持中文但某些图像库在扫描路径时遇到中文或空格会出现编码问题导致训练中途报错。建议统一用英文小写下划线命名并在一个labels.txt里记录类别顺序。3.3 自动划分数据集一份可直接改路径使用的脚本直接从原始采集目录划分train/valid/test的脚本如下import os import random import shutil from pathlib import Path src Path(flowers_raw) # 原始数据结构为 类别文件夹/图片 dst Path(flowers_split) # 输出目录 train_ratio, valid_ratio 0.7, 0.15 # test 占比 1 - train_ratio - valid_ratio random.seed(42) for split in (train, valid, test): (dst / split).mkdir(parentsTrue, exist_okTrue) for class_dir in src.iterdir(): if not class_dir.is_dir() or class_dir.name.startswith(.): continue images [p for p in class_dir.iterdir() if p.suffix.lower() in (.jpg, .jpeg, .png, .bmp)] random.shuffle(images) n_valid int(len(images) * valid_ratio) n_test int(len(images) * (1 - train_ratio - valid_ratio)) pairs [ (train, images[n_valid n_test:]), (valid, images[:n_valid]), (test, images[n_valid:n_valid n_test]), ] for split, part in pairs: target dst / split / class_dir.name target.mkdir(parentsTrue, exist_okTrue) for img in part: shutil.copy2(img, target / img.name) print(split done)逻辑说明脚本先固定随机种子保证每次生成相同划分。对每个类别按比例算出valid和test的样本数剩下的都归train。这样做能避免某些类别数量过少时某个子集为空。它复制图片而不是移动原始数据不会被破坏方便你后续调整划分比例。参数说明train_ratio0.7, valid_ratio0.15是花卉小数据集的常见设定。如果你的每个类别只有20张图建议把valid比例降到0.1test留在0.2优先保证训练样本不要被分走太多。random.seed(42)除了让结果可复现更重要的是后续你在排查精度问题时不会因为数据划分变化而怀疑训练代码有bug。4. 训练代码与参数调优用ViT做花卉识别的可复现套路4.1 用timm创建预训练ViT模型num_classes与DropPath注意点创建模型这一步是整个训练代码的核心。常见的做法是使用ImageNet-21k预训练权重并把分类头替换成目标花卉类别数import timm import torch.nn as nn model timm.create_model( vit_base_patch16_224_in21k, pretrainedTrue, num_classeslen(class_names), drop_path_rate0.1 )逻辑说明num_classes要等于你的花卉类别总数timm会自动把预训练分类头替换成对应输出维度的新线性层。drop_path_rate0.1是Stochastic Depth你可以把它理解为“训练时随机跳过某些残差分支”能有效抑制小数据微调时的过拟合。推理时DropPath自动失效不会影响最终精度。参数说明vit_base_patch16_224_in21k表示Base规模、patch 16、输入224、使用ImageNet-21k预训练。如果你的显卡只有6GB显存可以换成vit_small_patch16_224或vit_tiny_patch16_224训练速度明显更快最终精度通常只差1到2个百分点。这里不要盲目追求大模型课设场景里“跑得动”比“模型大”更重要。注意一点timm加载预训练权重时会自动把分类头忽略掉只加载backbone部分。如果你手动写了model.head nn.Linear(768, 5)就会覆盖掉num_classes传入的分类头。建议二选一不要重复操作。4.2 数据加载与增强训练、验证管线必须分开训练和验证的数据预处理不能混用。训练集可以用随机增强来扩展数据验证集和测试集必须保持固定预处理方式否则你打印出来的验证精度会失真。以下是使用timm自带数据管线的方式from timm.data import create_dataset, create_loader IMAGE_SIZE 224 BATCH_SIZE 32 NUM_WORKERS 4 MEAN, STD (0.485, 0.456, 0.406), (0.229, 0.224, 0.225) train_ds create_dataset(folder, rootflowers_split/train) valid_ds create_dataset(folder, rootflowers_split/valid) test_ds create_dataset(folder, rootflowers_split/test) train_loader create_loader( train_ds, input_sizeIMAGE_SIZE, batch_sizeBATCH_SIZE, is_trainingTrue, use_prefetcherTrue, re_prob0.5, scale(0.6, 1.0), ratio(0.75, 1.333), hflip0.5, num_workersNUM_WORKERS, meanMEAN, stdSTD ) valid_loader create_loader( valid_ds, input_sizeIMAGE_SIZE, batch_sizeBATCH_SIZE, is_trainingFalse, num_workersNUM_WORKERS, meanMEAN, stdSTD )逻辑说明训练loader打开了随机裁剪、随机长宽比和水平翻转其中scale(0.6, 1.0)表示裁剪面积占原图的60%到100%这比固定缩放更能抗过拟合验证loader只做缩放和中心裁剪不做任何随机增强。use_prefetcherTrue是timm提供的性能优化能在GPU上提前预取数据并做归一化让数据加载不再成为训练瓶颈。参数说明BATCH_SIZE32在单张RTX 3090这类24GB显卡上很轻松但如果你只有8GB显存请降到16或8并配合梯度累积。MEAN和STD是ImageNet数据集的均值和标准差ViT的预训练权重就是基于这套归一化训练的迁移时不要随意替换否则模型会“看不懂”输入图像。4.3 训练循环AdamW、Warmup与CosineAnnealingViT微调的优化器和CNN略有不同常见推荐是AdamW配合较小学习率。因为预训练特征已经很好了学习率过大会“冲坏”这些特征。我一般把分类头学习率设为1e-4backbone设为1e-5使用不同参数组分别更新import torch optimizer torch.optim.AdamW([ {params: model.head.parameters(), lr: 1e-4}, {params: [p for n, p in model.named_parameters() if head not in n], lr: 1e-5} ], weight_decay0.05) scheduler torch.optim.lr_scheduler.CosineAnnealingLR( optimizer, T_max30, eta_min1e-6 ) criterion torch.nn.CrossEntropyLoss()逻辑说明分类头是随机初始化的需要相对大的学习率快速收敛backbone是ImageNet预训练的只做小幅度微调。冻结backbone也能跑但会让模型无法适应花卉特有的颜色和纹理通常没必要完全冻结。weight_decay0.05是AdamW的常见搭配能抑制权重增长进一步提升泛化。参数说明T_max30对应总训练轮数。如果你的数据集很小训练轮数可以降到20T_max跟着改即可。eta_min1e-6是学习率的最低值让loss在训练后期尽肯能平稳落地。不要直接使用固定学习率训完整个项目ViT的自注意力在后期很敏感学习率不衰减会导致验证loss反复震荡。训练主循环里还要加一个关键操作混合精度。ViT的计算量比ResNet大半精度训练能显著降低显存占用和训练时间scaler torch.cuda.amp.GradScaler() for epoch in range(EPOCHS): model.train() for images, targets in train_loader: images, targets images.cuda(), targets.cuda() optimizer.zero_grad() with torch.cuda.amp.autocast(): outputs model(images) loss criterion(outputs, targets) scaler.scale(loss).backward() scaler.step(optimizer) scaler.update()逻辑说明混合精度不是在代码里直接调用images.half()而是用autocast把计算图内的输入输出自动转成半精度。GradScaler负责在反向传播时缩放loss防止梯度下溢。如果你的显卡不支持AMP把这段直接换回普通loss.backward()和optimizer.step()即可。4.4 模型评估与保存只看验证集选点测试集留到最后训练过程中我会在每一轮结束后跑一遍valid_loader计算准确率并保存历史最高验证集准确率对应的参数。这样能保证最后交出去的模型不是最后一轮的而是峰值模型best_acc 0.0 for epoch in range(EPOCHS): # ... 训练循环 ... model.eval() correct 0 total 0 with torch.no_grad(): for images, targets in valid_loader: images, targets images.cuda(), targets.cuda() outputs model(images) preds outputs.argmax(dim1) correct (preds targets).sum().item() total targets.size(0) acc correct / total print(fepoch {epoch1} valid_acc{acc:.4f}) if acc best_acc: best_acc acc torch.save(model.state_dict(), best_vit_flower.pth) torch.save({class_names: class_names}, class_names.pth)逻辑说明验证集不参与梯度计算所以model.eval()之后用with torch.no_grad()节省显存和计算时间。outputs.argmax(dim1)取每个样本预测概率最大的类与真实标签比较。保存时把state_dict和class_names分开保存权重加载时就不会出现类别顺序对不上的问题。很多人喜欢在训练过程中打印最后一轮的精度但从实验结果看训练后期验证精度的锯齿往往比前几轮更明显。把best模型单独存一份相当于给自己留了一颗“后悔药”。输入一张新图片推理时加载这份best模型通常会比最后一轮更可靠。5. 花卉识别避坑指南5个真实翻车点与排查思路5.1 现象训练Loss在降验证精度却原地不动原因是数据量太小而且分布太接近模型很快记住了训练集的纹理和背景。尤其当花卉照片来自同一批爬虫数据时训练和验证的背景颜色都高度相似模型几乎没有学到真正的花瓣结构。解决方案是先增强后换模型。最直接的做法是把训练loader里的scale下限从0.6改到0.4并增加hflip到0.5随机遮挡也可以用RandomErasing。如果增强后验证精度提升不明显再检查是不是分类头学习率太大把backbone的学习率降到5e-6重试。小数据上学习率对ViT的影响通常大于模型结构。5.2 现象加载权重时报错维度不匹配或标签KeyError原因是保存的是整个模型而加载时用了torch.load直接装载或者保存state_dict时类别顺序和当前环境不一致。我在第4章里把class_names单独保存就是为了避免这种冲突。解决方法是统一使用model.load_state_dict(torch.load(best_vit_flower.pth, map_locationcpu))。如果需要把训练好的模型放到别的代码目录里推理先打印一次class_names.pth里的顺序再和预测输出对照。不要假设文件夹排序和训练时一致Windows和Linux的文件读取顺序可能不同。5.3 现象显存不足或CPU训练跑到怀疑人生原因是ViT自注意力层的中间激活值占用显存严重尤其batch size和输入尺寸都偏大时非常容易翻车。不少同学把ImageNet用的224x224当默认值却忘了自己的显卡只有6GB。解决方法是按优先级依次尝试先把BATCH_SIZE降到16再关掉use_prefetcher或者把NUM_WORKERS降到2还不够就换vit_small_patch16_224或vit_tiny_patch16_224。如果只有CPU环境必须把timm.create_model里的pretrainedTrue保留因为预训练模型在CPU上微调15到20轮也能出效果但请把输入尺寸换成vit_tiny_patch16_224且BATCH_SIZE8否则跑完一个epoch就是一场耐力测试。5.4 现象验证集准确率虚高提交后真实图片翻车原因几乎总是验证管线被错误地加上了随机增强。我在带项目时见过有同学把RandomHorizontalFlip同时写到训练和验证loader里验证精度飙到97%可一旦用手机拍一张新照片准确率立刻降到六成。随机增强会改变验证数据的分布后续选模型的基准就废了。解决方法是检查验证Loader里is_training是否严格为False或者使用timm的create_loader时不要给验证Loader传hflip、scale这类增强参数。更可靠的指标是单独准备一个test集这个集只在你最终确认模型时跑一次不许用来反复调参。5.5 现象pretrainedTrue下载失败或timm版本API报错原因是预训练权重托管在海外对象存储上国内机房或校园网经常出现连接超时另一个原因是timm升级后部分旧模型名被废弃。当下载失败时问题并不在代码逻辑而是网络和缓存。解决方法是设置timm的权重下载镜像环境变量常见做法是把HF_ENDPOINT指向国内镜像再把下载好的权重放到timm的缓存目录中。timm对每个模型有固定的缓存文件名你可以先通过小模型触发一次下载找到缓存目录结构再把权重文件放进去。另外遇到API报错时先看timm.list_models(vit_base*)输出是否包含目标模型名如果模型名变了就按列表里的名字重新创建。不要硬记网上旧版本的模型名至少在timm 0.9.x时代API变动比你想的频繁得多。6. 推理、注意力可视化和部署的一点经验训练完best模型后最需要做的是用一张训练时从未见过的新图验证整个流程。这里面的关键不是“能跑通”而是“预处理必须和验证Loader完全一致”。以下是我常用的推理代码骨架from PIL import Image from torchvision import transforms transform transforms.Compose([ transforms.Resize((224, 224)), transforms.ToTensor(), transforms.Normalize((0.485, 0.456, 0.406), (0.229, 0.224, 0.225)) ]) image Image.open(test_sunflower.jpg).convert(RGB) x transform(image).unsqueeze(0) model.eval() with torch.no_grad(): logits model(x) probs torch.softmax(logits, dim1) top2 torch.topk(probs, 2) print(class_names[top2.indices[0][0]], top2.values[0][0])推理看一眼top2而不是top1是我自己的习惯。花卉识别中很多照片同时包含两种花比如向日葵旁边挤着雏菊top1可能是错的但top2完全合理。大作业里能输出候选类别比硬给一个答案更有说服力。如果你想在报告里展示ViT的“黑匣子”到底看了哪里最直接的方法是注册一个forward hook在某个Transformer层把注意力矩阵取出来。不需要修改模型内部代码def hook_fn(module, inp, out): attn out[-1] if isinstance(out, tuple) else out global attn_map attn_map attn.detach().cpu().mean(dim1).squeeze(0) model.blocks[5].attn.register_forward_hook(hook_fn)这里mean(dim1)是把多个注意力头平均成一个矩阵你要留意打印出的shape如果它是14x14就直接当作patch坐标可视化如果是196x196就需要按patch数重排成14x14。把CLS Token对应的那行attention取出来画热力图粘贴到原图上基本就能看出模型到底在关注花瓣还是叶子。最后说一个我踩过的坑。有一次我把验证集也顺手加了随机翻转val_acc到了97%交作业前用手机拍了几张野花测试五张错三张当场翻车。后来我把验证和测试流程固定成“Resize到256再CenterCrop到224”重新评估的真实精度是91%和训练时看到的虚高数字差了6个百分点。从那次以后我每次训练完都会多做一个动作单独放一张全新的、带干扰背景的花卉照片确认推理结果符合预期再决定要不要提交。做好这步你交出去的大作业就基本站稳了。希望帮到你。本文还有配套的精品资源点击获取
返回列表