☰
ViT图像分类实战:花卉识别迁移学习与PyTorch微调指南
2026/10/7 5:51:06 网站建设 项目流程

简介:基于Vision Transformer(ViT)的花卉图像分类Python实现,面向机器学习/深度学习相关课程设计、期末大作业及入门开发者,解决从数据加载、模型训练到图像分类的完整落地问题。压缩包共7个文件,包含6个Python脚本和1个keep占位文件;代码按配置、数据读取、模型定义、训练、分类等模块划分,整体约10KB,结构轻量且便于阅读。项目源自97分高分期末作业,注释较详细,下载即可运行;同时预留了二次开发空间,可替换数据集或调整ViT参数进行更多实验,适合快速搭建图像分类基线。对于需要撰写课程报告或期末展示的同学,可以直接复用其中模块化结构和思路,减少环境搭建与排错时间。目前已有383人学习浏览,参考价值在于帮助读者把ViT原理真正落到可运行代码上。

1. 大作业用ViT做花卉识别:先把预期管住

用 python 做基于 ViT 的图像分类花卉识别大作业,最常见的翻车方式不是代码写不出来,而是「从头训练」。ViT(Vision Transformer)把图片切成 16×16 的 patch 再按序列丢给 Transformer,结构优雅,但它天生缺 CNN 那种局部归纳偏置,在几千张花卉图的小数据集上从零训练,准确率常常停在 60% 上下,怎么调都上不去。实际能跑通的路线是:拿 ImageNet 预训练权重做迁移学习,微调分类头和高层 block,在 Oxford Flowers 这类数据上把 top-1 推到 90% 以上,同时交付训练脚本、混淆矩阵、注意力可视化和推理脚本。这套方案适合课程设计、毕设起步,以及想从 CNN 换到 Transformer 的初学者。你只需要准备一份分类好的花卉图片和一张 8GB 显存的显卡,剩下的事下文拆开讲。

2. ViT图像分类的花卉数据准备:数据集划分与预训练权重选择

2.1 ViT 为什么挑数据:归纳偏置弱,必须走迁移学习

CNN 做图像分类,靠卷积核天然带了两个先验:局部性和平移等变性。一个 3×3 卷积只看一个小邻域,所以几百张图也能学到简单纹理。ViT 不是这样,它把 224×224 的图切成长度为 196(14×14)的 patch 序列,每个 patch 展平后做线性投影,再和位置编码一起进入 Transformer encoder。注意力机制理论上能建模全局关系,但代价是它并不知道「相邻 patch 更可能属于同一个物体」这件事。数据量不够时,这种自由度过高会让模型记住训练集的噪声,而不是花的结构。

所以在花卉识别这种小数据任务里,没有人从头训 ViT。常见做法是加载在 ImageNet-1k 或 ImageNet-21k 上预训练好的权重,冻结一部分参数,只微调后半段。这个选择直接决定你后面所有代码的写法。如果课程要求里明确写了「必须用 ViT」,那vit_base_patch16_224是最稳妥的答案:它有 8600 万参数,结构是 12 个 Transformer block,答辩时讲得清楚,timm 和 Hugging Face transformers 里都有现成权重。别一上来就挑战 ViT-Large,8GB 显存会很难受。

2.2 花卉数据集怎么组织:ImageFolder 目录结构是标准答案

不管你用的是 Oxford 102 Flowers(102 类,每类 40 到 258 张,总共 8000 多张),还是网上爬来的 5 类或 17 类花卉数据,建议统一整理成 torchvisionImageFolder能直接读的目录结构。这一步做对了,后面数据加载零配置。

所在目录大致长这样:

data/ ├── train/ │ ├── rose/ # 类别名作为文件夹名 │ │ ├── rose_001.jpg │ │ └── rose_002.jpg │ ├── daisy/ │ └── tulip/ └── val/ ├── rose/ ├── daisy/ └── tulip/

ImageFolder会按根目录下每个子文件夹作为一个类别,并按字母序自动分配从 0 开始的标签。这一点非常重要,后面排查标签错位时要回来找它。划分 train/val 时,我一般不用简单的random_split,而是按类别做分层划分,保证每个类在训练集和验证集里的比例一致。floral 数据集最常见的坑是同一朵花的不同角度照片被分进了两个集合,验证指标会虚高,后面第 4 章会详细说。

2.3 backbone 选择对比:vit-base 还是 deit 还是 mobilevit

timm 里能直接用的 ViT 变体很多,大作业场景下我建议看这张表来定:

timm 模型名参数量输入分辨率单卡训练显存感受适合场景
vit_base_patch16_22486M2248GB 显存能跑 bs=16大作业首选,默认不会错
vit_small_patch16_22448M2246GB 显存比较轻松显存紧张或想快速迭代
deit_base_distilled_patch16_22487M224与 base 相当想利用蒸馏 token 提升一点精度
mobilevit_s5.6M2562GB 就能跑想在报告里讨论轻量化部署

选vit_base_patch16_224的理由很实际:它是最容易被搜索到复现结果的组合,网上讨论多,遇到报错好查;另外 ViT 在小数据集上微调,base 规模不会像从头训练那样欠拟合,因为预训练已经给了很强的特征。如果老师要求对比实验,你可以加一个 ResNet50 或者 MobileNetV2 做 baseline,说明 ViT 在全局特征上的优势,这也是这类大作业最常见的写法。森林图像分类、卫星图分类这类任务的组织方式完全一样,只是类别文件夹换成地貌类型而已。

2.4 加载预训练权重的两种姿势:timm 与 transformers

timm 写法更简洁,适合训练脚本:

import timm model = timm.create_model( "vit_base_patch16_224", pretrained=True, num_classes=102, drop_path_rate=0.1 ) print(model.head)

num_classes=102传进去之后,timm 会自动把原来的 ImageNet 分类头替换成新的全连接层,不需要你手动改model.head。drop_path_rate=0.1是 Stochastic Depth 的丢弃比例,微调小数据集时建议开着,能明显抑制过拟合。

Hugging Face transformers 的写法更适合你想顺便展示「用AutoModelForImageClassification做迁移」的场景:

from transformers import ViTForImageClassification model = ViTForImageClassification.from_pretrained( "google/vit-base-patch16-224", num_labels=102, ignore_mismatched_sizes=True )

两种写法训练时前向输出不一样:timm 直接返回[B, num_classes]的 logits;transformers 返回一个ImageClassifierOutput对象,要取.logits。如果中途换库,务必检查这里,不然你会看到loss.backward()报NoneType错。我一般推荐 timm,因为它代码量少、配套的create_model参数体系统一,适合交作业。

3. 基于ViT的花卉识别微调:从数据装载到训练闭环

3.1 数据装载与增强:这个 transform 直接决定涨不涨点

ViT 对输入尺寸敏感,预训练权重默认吃 224×224。训练时不要只做 Resize,至少要有 RandomResizedCrop 和 ColorJitter,否则花卉这种类间差异小、类内差异大的数据很容易过拟合。

# transform.py from torchvision import transforms train_transform = transforms.Compose([ transforms.RandomResizedCrop(224, scale=(0.08, 1.0), ratio=(0.75, 1.333)), transforms.RandomHorizontalFlip(p=0.5), transforms.RandomRotation(15), transforms.ColorJitter(brightness=0.3, contrast=0.3, saturation=0.3, hue=0.05), 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.08, 1.0)表示裁剪面积占原图的 8% 到 100%,模拟不同远近的拍摄;ColorJitter 的 hue 只给了 0.05,因为花的颜色本身是重要类别特征,动太狠会把玫瑰和月季搞混。Normalize 用的均值标准差必须和预训练时一致,ImageNet 就是上面这一组,不要自己统计花卉数据集的均值来代替,迁移学习里这一步属于「不该省的事」。

3.2 用 timm 构建 ViT 并设置冻结策略

模型搭好之后,下一步是决定哪些层参与训练。小数据集(几千张)上微调,常见做法是冻结前面一半的 block,只解冻后面几个 block 和分类头。

# build_model.py import timm import torch.nn as nn num_classes = 102 model = timm.create_model( "vit_base_patch16_224", pretrained=True, num_classes=num_classes, drop_path_rate=0.1 ) # 冻结前 6 个 transformer block,解冻 head 和后面的 block freeze_until = 6 for name, param in model.named_parameters(): if name.startswith("blocks"): block_idx = int(name.split(".")[1]) # blocks.5.attn.qkv.weight -> 5 if block_idx < freeze_until: param.requires_grad = False elif "head" not in name: param.requires_grad = False

这段代码里name.split(".")[1]取到的是 block 编号,比如blocks.5.attn.qkv.weight的编号是 5。冻结前 6 个 block 的理由是:浅层学到的是边缘、颜色块这类通用纹理,深层才接近具体语义,Flowers 数据集的拍照风格和 ImageNet 差异不大,底层特征直接复用即可。如果你的数据集有 1 万张以上,或者拍摄角度很特殊(俯拍、微距、逆光),可以把freeze_until改成 0,全部解冻,学习率仍然用 5e-5 量级。

3.3 优化器、学习率与训练循环:AMP、梯度累积、早停

ViT 微调最常用的优化器组合是 AdamW + 小学习率 + 余弦退火。不要照搬 CNN 分类那套 lr=1e-3 的配置,ViT 的 attention 层对学习率很敏感,我见过太多人把 lr 设成 1e-3 然后 loss 直接发散。推荐 5e-5 起步。

# train.py import torch import torch.nn as nn from torch.cuda.amp import GradScaler, autocast from torch.optim import AdamW from torch.optim.lr_scheduler import CosineAnnealingLR from torch.utils.data import DataLoader from torchvision.datasets import ImageFolder device = "cuda" if torch.cuda.is_available() else "cpu" epochs = 30 lr = 5e-5 batch_size = 16 accum_steps = 2 # 等效 batch = 16 * 2 = 32 patience = 5 train_ds = ImageFolder("data/train", transform=train_transform) val_ds = ImageFolder("data/val", transform=val_transform) train_loader = DataLoader(train_ds, batch_size=batch_size, shuffle=True, num_workers=4, pin_memory=True) val_loader = DataLoader(val_ds, batch_size=batch_size, shuffle=False, num_workers=4, pin_memory=True) optimizer = AdamW(model.parameters(), lr=lr, weight_decay=0.01) scheduler = CosineAnnealingLR(optimizer, T_max=epochs) criterion = nn.CrossEntropyLoss() scaler = GradScaler() best_acc = 0.0 bad_epochs = 0 for epoch in range(epochs): model.train() running_loss = 0.0 for i, (images, labels) in enumerate(train_loader): images, labels = images.to(device), labels.to(device) with autocast(): logits = model(images) loss = criterion(logits, labels) / accum_steps scaler.scale(loss).backward() if (i + 1) % accum_steps == 0: scaler.step(optimizer) scaler.update() optimizer.zero_grad() running_loss += loss.item() * accum_steps # 每个 epoch 结束后验证 model.eval() correct = 0 total = 0 with torch.no_grad(): for images, labels in val_loader: images, labels = images.to(device), labels.to(device) with autocast(): logits = model(images) preds = logits.argmax(dim=1) correct += (preds == labels).sum().item() total += labels.size(0) acc = correct / total scheduler.step() if acc > best_acc: best_acc = acc torch.save({ "model_state_dict": model.state_dict(), "class_to_idx": train_ds.class_to_idx, "acc": best_acc, }, "best.pth") bad_epochs = 0 else: bad_epochs += 1 if bad_epochs >= patience: print(f"early stop at epoch {epoch}, best acc {best_acc:.4f}") break

这段代码的关键点有三个。第一,autocast只包前向和 loss 计算,反向传播scaler.scale(loss).backward()照常在 FP32 下进行,混合精度的意义是让激活值和梯度在 FP16 下计算、优化器状态保持在 FP32,显存可以省下近一半。第二,loss / accum_steps是为了梯度累积,每 2 个 step 才更新一次参数,相当于把 batch 从 16 放大到 32,在 8GB 显存上这是白捡的稳定性。第三,scheduler.step()放在验证之后,保证最后一个 epoch 也用到了完整的 cosine 区间,不会出现学习率还没降到底就停止的尴尬。

4. 提取花卉识别训练高频排查:损失不降、OOM 与过拟合的表现

4.1 标签错位:训练集 loss 正常下降,验证集准确率稳在 50% 附近

现象:训练集 loss 从 4.6 慢慢降到 0.8,看起来一切正常,但验证集 top-1 一直在 50% 上下晃,甚至换了几种模型都一样。

原因:ImageFolder按文件夹名排序生成标签,比如daisy是 0、rose是 1、tulip是 2,是按字母序。如果你的labels.txt是按数据集原来给的顺序写的,或者推理脚本里手动硬编码了["daisy", "rose", ...],两端一旦错一位,整个评估就是错的。另一个常见来源是训练代码里DataLoader的shuffle=True带到了验证集,导致每轮验证的标签和图片对应关系每次都在变。

解决:训练结束后打印train_ds.class_to_idx,和你的labels.txt逐行比对。建议把class_to_idx和模型权重一起保存在 checkpoint 里,推理时直接读,不要靠手写类名列表。验证集的DataLoader务必设置shuffle=False,这是一条不会亏的纪律。

4.2 显存不足:batch size 一调大就 CUDA OOM

现象:batch_size=32直接报CUDA out of memory,显存 8GB,训练中断。

原因:ViT-Base 有 8600 万参数,加上 AdamW 的动量状态,一个 step 的显存占用是模型本身的好几倍。很多人把 batch 调到 4 硬跑,结果 BN 层(这里其实是 LayerNorm)统计不稳,loss 抖动剧烈。

解决:优先用混合精度 + 梯度累积的组合。torch.cuda.amp能把激活显存砍半,梯度累积用 2 步等效放大 batch,这两招同时用,bs=16 + accum=2 在 8GB 卡上能稳定跑完 30 个 epoch。如果还撑不住,再把输入分辨率从 224 降到 192。注意:ViT 的 position embedding 是预训练时定死的,降低分辨率需要把pos_embed做插值,timm 里可以通过timm.models.vision_transformer.pos_embed处理,但这属于高阶操作,大作业不建议碰,直接换vit_small_patch16_224更省事。

4.3 过拟合:训练集 top-1 到 99%,验证集停在 70%

现象:训练集准确率一路涨到 99%,验证集却从第 8 个 epoch 开始不再上升,甚至往下掉。

原因:三类问题叠加。一是数据增强太弱,只有 Resize 和水平翻转,模型把花的背景、拍摄角度都背下来了;二是drop_path_rate=0,ViT 的深层次结构在几千张图上没有正则;三是冻结层数不对,如果只训练分类头,浅层特征和当前数据分布不匹配,训练集暴涨、验证集不涨是必然结果。

解决:把RandomResizedCrop(224, scale=(0.2, 1.0))加上,ColorJitter 也打开;drop_path_rate设成 0.1;解冻 block 数量从 0 调整到 6 或 12。改完这三处再观察 loss 曲线,训练集和验证集的差距通常会从 20 个百分点缩到 5 个百分点以内。另外,如果验证集里存在同一朵花的多张连拍,属于数据泄漏,需要先按文件名或者图像感知哈希去重再划分,不然你模型的「泛化能力」是假的。

4.4 加载权重报 size mismatch 或 Unknown Pytorch header

现象:在另一台电脑上跑torch.load("best.pth"),报KeyError: 'model_state_dict'或者size mismatch for head.weight。

原因:训练脚本里如果直接torch.save(model, "best.pth"),保存的是整个模型对象,换环境后类定义位置不对就报 header 错。size mismatch则是加载脚本里num_classes和训练时不一致,比如训练用了 102 类花卉,推理脚本里timm.create_model(..., num_classes=5),最后一层形状对不上。

解决:统一用第 3 章的写法,保存state_dict而不是整个模型对象,并顺手把class_to_idx也存进去。加载时用:

checkpoint = torch.load("best.pth", map_location="cpu") model.load_state_dict(checkpoint["model_state_dict"])

如果size mismatch仍然出现,先打印checkpoint["model_state_dict"]["head.weight"].shape,和当前模型的model.head.weight.shape对比,多半是 num_classes 设错了。这条排查顺序能帮你节省至少一晚上。

4.5 同一条代码跑两遍结果差 2 到 3 个点:随机种子没固定

现象:同一份训练代码,昨天跑 top-1 是 93.2%,今天跑变成了 90.8%,其他什么都没改。

原因:PyTorch 默认不固定随机数,数据打乱顺序、drop path 的丢弃位置、GPU 上的非确定性算子都会引入波动。ViT 本身带 dropout,对种子敏感程度比 CNN 更高,这不是玄学,是浮点运算的客观现象。

解决:在训练脚本最开头写一个seed_everything(42),同时固定 torch、numpy、random 三个库,并把cudnn.deterministic打开:

import random import numpy as np import torch def seed_everything(seed=42): random.seed(seed) np.random.seed(seed) torch.manual_seed(seed) torch.cuda.manual_seed_all(seed) torch.backends.cudnn.deterministic = True torch.backends.cudnn.benchmark = False

需要说明的是,固定种子之后,两次训练结果仍然可能有 0.2 左右的波动,这是 GPU 算子在 fp16 下的正常误差,不影响结论。真正该做的不是追求完全复现,而是固定种子后至少保证你报告里写的指标是稳定可复现的,答辩时不会被问倒。

5. ViT花卉识别效果验证:混淆矩阵与注意力热力图

5.1 用混淆矩阵找出模型到底错在哪些类

训练完成不等于大作业完成。用验证集做一次完整评估,输出混淆矩阵,这能直接告诉评审「模型在哪两类花之间犹豫」。Flowers 数据集常见错误是玫瑰和月季、雏菊和蒲公英这种视觉近邻类。

# evaluate.py import torch import numpy as np import matplotlib.pyplot as plt import seaborn as sns from sklearn.metrics import confusion_matrix, classification_report from torch.utils.data import DataLoader model.eval() all_preds = [] all_labels = [] with torch.no_grad(): for images, labels in val_loader: images = images.to(device) logits = model(images) preds = logits.argmax(dim=1).cpu().numpy() all_preds.extend(preds) all_labels.extend(labels.numpy()) cm = confusion_matrix(all_labels, all_preds) plt.figure(figsize=(12, 10)) sns.heatmap(cm, annot=True, fmt="d", cmap="Blues", xticklabels=class_names, yticklabels=class_names) plt.xlabel("Predicted") plt.ylabel("True") plt.tight_layout() plt.savefig("confusion_matrix.png", dpi=150) print(classification_report(all_labels, all_preds, target_names=class_names))

如果你的类别数超过 30,整张矩阵会挤成一团,我一般会额外打印每一行「错误最多的前 3 个类别」,比如「玫瑰误判为月季 18 次」,这个信息比热力图更适合写进报告结论。classification_report会给出每个类的 precision、recall、f1-score,答辩时你可以挑三个最容易被认错的花,说明数据层面的原因,比如颜色接近、拍摄角度多变。

5.2 把 CLS 注意力画出来:让评审看到模型在看花心还是背景

ViT 的最后一层 attention 可以告诉我们模型做决策时在关注哪些 patch。这一步是最能体现「你理解 ViT」的证据,也是大作业加分项。timm 的 ViT 前向不直接暴露 attention,需要用 hook 把最后一层 Transformer block 的 attention 权重捞出来。

# visualize_attention.py import torch import torch.nn.functional as F from PIL import Image from torchvision import transforms # 取最后一层 block 的 attention 输出 attention_map = {} def get_attention(name): def hook(module, input, output): # output shape: [B, num_heads, N, N] attention_map[name] = output[0].detach() return hook model.blocks[-1].attn.register_forward_hook(get_attention("last")) img = Image.open("sample_rose.jpg").convert("RGB") x = val_transform(img).unsqueeze(0).to(device) with torch.no_grad(): model(x) attn = attention_map["last"] # [1, 12, 197, 197] cls_attn = attn[0, :, 0, 1:].mean(dim=0) # 平均所有 head,去掉 CLS 自身 grid = cls_attn.reshape(14, 14).cpu().numpy() grid = F.interpolate( torch.tensor(grid).unsqueeze(0).unsqueeze(0), size=(224, 224), mode="bilinear", align_corners=False ).squeeze().numpy() img_np = np.array(img.resize((224, 224))) / 255.0 plt.imshow(img_np) plt.imshow(grid, alpha=0.5, cmap="jet") plt.axis("off") plt.savefig("attention_vis.png", dpi=150)

注意这里attn[0, :, 0, 1:]的含义:ViT 输入序列是[CLS] + 196 个 patch,所以 attention 矩阵是 197×197。取[CLS]这一行,就是每个 patch 对最终分类向量的贡献权重;mean(dim=0)是平均 12 个 head。分辨率还原时,因为 patch_size=16,196 个 patch 天然是 14×14 网格,用双线性插值放大到 224 再叠加到原图上。如果热力图集中在一片绿色背景而不是花心,说明模型在用背景做分类,这类样本正是你报告里该分析的反例。

5.3 一个能直接交付的推理脚本:返回 top-5 和置信度

大作业代码里通常会要求一个「输入一张图,输出类别」的脚本。建议直接做成 top-5 输出,比只给一个类别更有说服力。

# predict.py import torch from PIL import Image def predict_one_image(model, img_path, class_names, device="cuda"): model.eval() img = Image.open(img_path).convert("RGB") x = val_transform(img).unsqueeze(0).to(device) with torch.no_grad(): logits = model(x) probs = torch.softmax(logits, dim=1)[0] top5_prob, top5_idx = torch.topk(probs, 5) result = [] for prob, idx in zip(top5_prob.tolist(), top5_idx.tolist()): result.append((class_names[idx], round(prob, 4))) return result

这个函数在推理时的关键点是model.eval()。timm 的 ViT 包含 Dropout 和 DropPath,如果漏了eval(),同样的图每次预测结果可能不同,这也是很多人最终演示时翻车的地方。class_names请务必从 checkpoint 里的class_to_idx反推得到,不要自己重新写一遍类名列表,除非你想再体验一次第 4.1 节的标签错位。

6. 进阶:把 ViT 花卉模型导出 ONNX 并接一个 Web 调用

如果还想再加一个亮点,把训练好的 PyTorch 模型转成 ONNX,用一个简单的 Flask 或者 FastAPI 接口接收图片、返回类名,整套大作业的完整度就上来了。导出代码很短:

# export_onnx.py model.eval() dummy = torch.randn(1, 3, 224, 224).to(device) torch.onnx.export( model, dummy, "vit_flower.onnx", input_names=["input"], output_names=["logits"], dynamic_axes={"input": {0: "batch"}, "logits": {0: "batch"}}, opset_version=14 )

这里有几个实践经验。第一,dynamic_axes只把 batch 维度设成动态,不要动 224×224 的空间维度,否则推理框架要做 resize 插值,容易和训练时的预处理不一致。第二,ONNX 导出后必须用onnxruntime验证一遍数值,常见做法是拿同一张图分别走 PyTorch 和 ONNX,比较 top-1 结果。如果发现 ONNX 结果明显偏离,优先检查model.eval()有没有在导出前调用,以及输入张量的均值方差归一化有没有被推理端重复执行。第三,LayerNorm 在 fp16 量化下容易出抖动,如果 Web 端推理用 CPU,我建议保持 fp32 的 ONNX,不要为了省 5MB 去做 int8 量化,ViT 这种结构量化掉点常常超过 3 个点,不值得。

最后说一个我自己的教训式习惯:每次训练开始前第一行就是seed_everything(42),这个习惯帮我避开了无数次「为什么演示时结果和报告不一样」的尴尬。当初我为了赶进度跳过了固定种子,答辩现场模型跑出来比报告低了 2 个百分点,被评委追问了整整五分钟。从那以后,我拿到任何视觉训练代码,第一件事都是检查随机种子和验证集顺序。既然你选择了 ViT 做花卉识别,也请顺手把这两个习惯带上。希望帮到你。

本文还有配套的精品资源,点击获取

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询