☰
Transformer图像分类实战:木薯叶病虫害ViT源码复现与调优指南
2026/10/1 13:19:23 网站建设 项目流程

简介:一份基于Python与Transformer模型的木薯叶病虫害分类源码,面向深度学习初学者与期末大作业场景,用于解决木薯叶图像识别与病害分类问题,项目难度适中,代码均已在本地编译运行。压缩包共12个文件,含6个.py脚本、5个.pyc预编译文件及1个README说明,整体仅11KB,已有199人学习下载。源码经助教老师审定,模块清晰:主程序调度、全局变量配置、GPU适配、数据加载与Transformer分类网络一应俱全,README给出运行指引。既适合期末大作业完整参考,也便于迁移至其他农作物病害分类项目,对理解注意力机制在视觉任务中的应用尤为有益。

1. 木薯叶病虫害分类:为什么这个视觉任务绕不开 Transformer

transformer 图像分类在木薯叶病虫害数据上比 ResNet 能高出 2~5 个点的准确率,这件事在公开比赛和不少落地项目里都被反复验证过。木薯叶的病斑往往很小、分布零散,常规卷积网络容易把注意力浪费在整片叶子的纹理背景上;而 transformer 把图片切成 patch token 之后,靠自注意力在全局范围建立长距离关联,对“一个病斑与另一处病斑相互印证”这类特征特别擅长。这个源码包的组件不算复杂:python 读取数据、ViT 模型定义、训练与可视化脚本,正好覆盖课程设计和论文复现的所有环节。适合两类人:需要交高分课设的在校生,以及想在农业视觉方向快速搭一个能出图的基线模型的工程师。

2. 拿到源码包先别跑:目录、环境与数据集的三个准备步骤

拿到一个 python 实现的 transformer 木薯叶病虫害分类源码包,我的习惯是先解压,把每个 .py 文件的结构摸一遍,再谈训练。直接运行 train.py 通常会被数据路径、依赖不对齐这类问题卡住,而那跟模型本身没关系。先把项目读明白,后面每一步才不会被黑匣子牵着走。

2.1 从入口文件反推项目结构:train.py、model.py、data_utils.py 各管什么

解压之后先看一层目录,很多源码包的结构是类似这样的:

unzip 木薯叶_transformer分类源码.zip -d cassava_project cd cassava_project tree -L 2
cassava_project/ ├── train.py ├── model.py ├── data_utils.py ├── config.py ├── utils.py ├── requirements.txt ├── data/ └── weights/

train.py 是入口,model.py 放 ViT 模型或预训练封装,data_utils.py 做 Dataset 和增强,config.py 集中放超参数。拿到包先看这三个文件就能知道该项目的“体质”:如果 model.py 里只有 create_model 函数,多半走 timm 封装路线;如果出现 PatchEmbed、MultiHeadAttention 这类类名,就是手写实现路线。两种写法在答辩时讲法完全不同,前者强调迁移学习,后者强调架构细节。

如果你手里的包结构混乱,可以用一个小脚本把每个文件的顶层函数和类扫出来,比逐个打开省时间:

import ast from pathlib import Path root = Path("cassava_project") for py in root.rglob("*.py"): try: tree = ast.parse(py.read_text(encoding="utf-8")) except SyntaxError: continue items = [] for node in tree.body: if isinstance(node, (ast.FunctionDef, ast.ClassDef)): items.append(node.name) if items: print(f"{py.relative_to(root)}: {', '.join(items)}")

这段代码只提取函数和类名,不执行任何逻辑,能快速判断项目是“timm 封装派”还是“手写 ViT 派”。这一步花两分钟,后面复现时能少走很多弯路。顺带说一句,很多高分项目源码包会把权重保存目录和日志目录也带上,如果你看到 weights/ 或 logs/,说明包作者自己调试过,这种包的可信度通常高一些。

2.2 环境依赖:先装对 PyTorch、timm 与 CUDA 的搭配

依赖清单常见写法是:

python -m venv .venv source .venv/bin/activate pip install "torch>=1.13" "timm>=0.9" opencv-python pandas tqdm einops scikit-learn matplotlib

Python 版本建议 3.8~3.10,vscode 里把解释器切到刚建的 .venv 即可。torch 的版本注意和 nvidia-smi 看到的 CUDA 版本匹配,如果你在本地 Windows 环境配置 python 环境,直接安装对应 cu118/cu121 的版本更省事,不建议用默认源硬装再回头查驱动。timm 不是必须,但大多数源码包会用它加载 ImageNet 预训练权重;如果你手头包是自己手写 ViT 的,torch 之外只需要 einops 做张量维度变换。还有一个容易忽略的点:transformers 这个库其实不一定用得上,只有通过 HuggingFace 接口加载模型时才需要;requirements.txt 里即使写了也可以不装。

装完先验证环境:

python -c "import torch; print(torch.__version__, torch.cuda.is_available())"

输出里出现 True 再继续。很多新手在这里卡住,不是模型代码问题,而是 torch 装成了 CPU 版,后续脚本执行到 .cuda() 时直接抛错。这个检查只要十秒钟,但能拦住一半以上的“复现翻车”。

2.3 数据集组织:按类别建目录还是用 CSV 标注

木薯叶病虫害分类最常见的公开数据是 Kaggle 上的 Cassava Leaf Disease,五分类,约两万一千张图片,原始标注在 train.csv 里,每行是 image_id 和 label 两列。源码包如果直接读 CSV,数据目录不一定要按类别分;但多数训练脚本为了方便用 ImageFolder,会要求 data/train/0、data/train/1 这种布局。两种组织方式都能跑,关键是训练脚本里用的是哪种数据接口。

先写一个把 CSV 转成按类别目录的脚本:

import pandas as pd import shutil from pathlib import Path df = pd.read_csv("train.csv") src_dir = Path("train_images") dst_root = Path("data/train") for image_id, label in df[["image_id", "label"]].values: src = src_dir / f"{image_id}.jpg" dst = dst_root / str(label) / f"{image_id}.jpg" dst.parent.mkdir(parents=True, exist_ok=True) shutil.copy(src, dst)

注意 image_id 不带后缀,手动拼 .jpg 几乎必做;如果数据集解压出来本身就带后缀,把这一行改掉即可。然后是划分训练验证集,这一步要用分层抽样,而不是普通随机划分:

from sklearn.model_selection import train_test_split train_idx, val_idx = train_test_split( df.index, test_size=0.2, stratify=df["label"], random_state=42, ) df.loc[train_idx].to_csv("train_split.csv", index=False) df.loc[val_idx].to_csv("val_split.csv", index=False)

stratify=df["label"] 是关键。如果直接随机划分,五类样本不均衡会导致少数类在验证集里只有几十张,准确率抖动得不敢信,而且这种现象在后面对比实验时特别明显。木薯叶这类农业数据集几乎都有“健康叶是多数类、病斑类是少数类”的分布特点,训练集和验证集的比例必须在划分阶段就固定住,后面所有实验才可比。

3. 用 Python 手写 ViT 分类器:Patch Embedding 到 Encoder 的完整拆解

这个源码包的核心模型不管叫 vit_base 还是 LeafViT,骨架基本都是 transformer 架构。这一章我不念结构图,而是把 Patch Embedding、位置编码、Encoder 堆叠这三段拆到能直接抄代码的粒度,顺序和训练时前向传播的顺序一致。这样你读源码好比跟着数据流走一遍,而不是对着类名猜槽位。

3.1 图片如何变成 token 序列:patch embedding 与位置编码

transformer 自己只能吃 token 序列,所以图像要先切成 patch。常见做法是把 448 × 448 的木薯叶原图切成 28 × 28 个 16 × 16 的小块,每一块通过卷积映射成 768 维向量,这个操作就是 PatchEmbed:

import torch from torch import nn from einops import rearrange class PatchEmbed(nn.Module): def __init__(self, in_chans=3, patch_size=16, embed_dim=768): super().__init__() self.patch_size = patch_size self.proj = nn.Conv2d( in_chans, embed_dim, kernel_size=patch_size, stride=patch_size, ) def forward(self, x): x = self.proj(x) # B, 768, H/p, W/p x = rearrange(x, "b d h w -> b (h w) d") return x

用 stride=patch_size 的卷积一次完成“切块加线性映射”,这是源码包里最常见的实现。输入 448、patch 16,输出 28 × 28 等于 784 个 token,每个 token 的语义是原图一个 16 × 16 局部区域。patch_size 是一个要反复权衡的参数:越小越能看清病斑细节,但 token 数平方增长,显存不答应;16 是绝大多数预训练模型默认配置,先不要动它。

然后是位置编码。你问 transformer 的位置信息怎么计算,ViT 的答案最直接:自注意力本身对 token 顺序不敏感,patch token 打乱顺序后注意力结果完全不变,所以必须额外注入位置信息。ViT 的做法是设置一个可学习的参数矩阵,形状是(token 数, embed_dim),直接加到 patch token 上:

class PositionalEncoding(nn.Module): def __init__(self, num_patches, embed_dim): super().__init__() self.pos_embed = nn.Parameter( torch.zeros(1, num_patches + 1, embed_dim) ) nn.init.trunc_normal_(self.pos_embed, std=0.02) def forward(self, x): return x + self.pos_embed

加一是给 [CLS] 分类 token 留位置。ViT 会在 patch token 序列最前面拼一个可学习的 [CLS] token,最终只用这个位置的输出做分类。这个设计从 BERT 沿袭而来,比把所有 patch 平均更稳。如果你把输入分辨率从 224 改成 448,可学习位置编码的个数就和预训练不一致,直接使用会报 shape 错误;正规做法是把预训练 pos_embed 按间距插值成新尺寸,这个坑后面有完整解法。

3.2 多头注意力在叶片病害图中具体捕捉什么

多头注意力的计算不复杂:每个 token 生成 query、key、value,query 和所有 key 做点积并归一化得到注意力权重,再对 value 加权求和。多头就是把 768 维切成 12 个 64 维子空间,让每个头各管一种关系。代码实现如下:

class Attention(nn.Module): def __init__(self, dim, num_heads=12): super().__init__() self.num_heads = num_heads head_dim = dim // num_heads self.scale = head_dim ** -0.5 self.qkv = nn.Linear(dim, dim * 3, bias=False) def forward(self, x): B, N, D = x.shape qkv = self.qkv(x).reshape(B, N, 3, self.num_heads, D // self.num_heads) qkv = qkv.permute(2, 0, 3, 1, 4) # 3, B, H, N, head_dim q, k, v = qkv[0], qkv[1], qkv[2] attn = (q @ k.transpose(-2, -1)) * self.scale attn = attn.softmax(dim=-1) out = attn @ v out = out.transpose(1, 2).reshape(B, N, D) return out

注意力矩阵尺寸是 B × H × N × N,N 是 token 数,这也是 ViT 显存开销大的直接原因。448 × 448 输入下 N 是 785,batch 一放大,注意力层立刻变成显存大户。因此后面训练章给的梯度累积方案,就是为了缓解这个问题。

多头注意力在木薯叶图片上具体捕捉什么?以我的观察,有些头会把病斑边缘与相邻叶片区域的 patch 拉近,有些头会长距离关联两个分离但相似的病斑。这种跨 patch 建模能力正是小病斑、散分布场景需要的,也是它比 CNN 局部感受野更适合叶片病害识别的原因。如果项目层面想进一步提升,可以对比 Swin Transformer 的窗口注意力,但实现复杂度明显更高;对一份以“基于 transformer 的图像分类”为主题的源码包,ViT 已经足够撑起高分局了。

3.3 Encoder 堆叠与分类头:预训练权重如何接进来

单个 Encoder Block 由注意力、LayerNorm、MLP 和残差连接组成:

class Block(nn.Module): def __init__(self, dim, num_heads, mlp_ratio=4, dropout=0.1): super().__init__() self.norm1 = nn.LayerNorm(dim) self.attn = Attention(dim, num_heads) self.norm2 = nn.LayerNorm(dim) self.mlp = nn.Sequential( nn.Linear(dim, int(dim * mlp_ratio)), nn.GELU(), nn.Dropout(dropout), nn.Linear(int(dim * mlp_ratio), dim), nn.Dropout(dropout), ) def forward(self, x): x = x + self.attn(self.norm1(x)) x = x + self.mlp(self.norm2(x)) return x

这是 Pre-LN 结构,LayerNorm 放在注意力之前,残差连接在外面。相比 Post-LN,Pre-LN 在深网络里梯度更稳,这也是 ViT 能堆 12 层还能从预训练继续微调的原因。源码包里判断是手写还是封装,看这个 Block 的实现就能一眼分辨。

把前面几段组装成完整模型:

class ViT(nn.Module): def __init__(self, img_size=448, patch_size=16, embed_dim=768, depth=12, num_heads=12, num_classes=5): super().__init__() n_patches = (img_size // patch_size) ** 2 self.patch_embed = PatchEmbed(3, patch_size, embed_dim) self.cls_token = nn.Parameter(torch.zeros(1, 1, embed_dim)) self.pos_embed = PositionalEncoding(n_patches, embed_dim) self.blocks = nn.Sequential(*[ Block(embed_dim, num_heads) for _ in range(depth) ]) self.norm = nn.LayerNorm(embed_dim) self.head = nn.Linear(embed_dim, num_classes) def forward(self, x): x = self.patch_embed(x) cls = self.cls_token.expand(x.shape[0], -1, -1) x = torch.cat([cls, x], dim=1) x = self.pos_embed(x) x = self.blocks(x) return self.head(self.norm(x[:, 0]))

注意分类头直接输出五类 logits,没有手动 softmax。CrossEntropyLoss 内部自带 log_softmax,手动加 softmax 反而会让数值变差,这是新手最容易写错的位置。预训练权重加载用 timm 很直接:

import timm model = ViT(img_size=384, num_classes=5) pretrained = timm.create_model("vit_base_patch16_224", pretrained=True).state_dict() state = model.state_dict() for k, v in pretrained.items(): if k in state and state[k].shape == v.shape: state[k] = v model.load_state_dict(state)

形状相同的层直接复制,包括 patch_embed、attention、mlp;分类头和 CLS token 的 shape 不匹配会自动跳过。这种方式比直接 timm.create_model(num_classes=5) 更可控,答辩时也能讲清楚“哪部分用了预训练、哪部分从零学”。如果输入不是 224,上面的插值处理要留到后面。

4. 木薯叶分类训练参数调优:让准确率从 80 到 90 的配置与手段

模型结构定了之后,分数高不高基本全看训练策略。这个项目要从“能跑”到“高分”,下面几个设置按顺序加,效果远比反复换模型骨架来得快。参数不是一拍脑袋定的,是 ViT 在中小规模图像分类任务上很通用的一组基线,抄的时候照着调即可。

4.1 数据增强:CutMix 适合重叠病斑,MixUp 会模糊健康叶

木薯叶数据集的难点是病斑密度高、叶片相互遮挡,常规的随机裁剪翻转不够用。源码包里常见的增强组合是 RandomResizedCrop、RandomHorizontalFlip、ColorJitter,再配 CutMix。CutMix 把一张图的某个矩形区域换成另一张图,标签按面积比例混合,让模型必须同时关注叶片局部和整体。实现如下:

import torch def cutmix(x, y, alpha=1.0): idx = torch.randperm(x.shape[0], device=x.device) y_b = y[idx] lam = torch.distributions.Beta(alpha, alpha).sample() cx, cy = torch.randint(x.shape[2], (2,), device=x.device) r = int(x.shape[2] * (1 - lam.sqrt())) x1, x2 = max(0, cx - r // 2), min(x.shape[2], cx + r // 2) y1, y2 = max(0, cy - r // 2), min(x.shape[3], cy + r // 2) x[:, :, x1:x2, y1:y2] = x[idx, :, x1:x2, y1:y2] lam = 1 - ((x2 - x1) * (y2 - y1)) / (x.shape[2] * x.shape[3]) return x, y, y_b, lam

调用时,loss 变成两项的加权:

logits = model(x) loss = lam * criterion(logits, y) + (1 - lam) * criterion(logits, y_b)

CutMix 的 lam 从 Beta(1, 1) 采样,等价于均匀分布。相比 MixUp 把两张图像素级叠加,CutMix 保留清晰的空间结构,对“小病斑加局部判断”更友好;MixUp 会把健康叶片也叠加模糊,在这个任务里我一般不做首选的强增。如果机器能承受,CutMix 之后还可以加 RandAugment,幅度控制在 5 以内,再大容易让叶片颜色失真。

4.2 学习率策略:warmup + cosine decay 的具体数字

Transformer 对学习率比 CNN 敏感得多,AdamW 下直接上 1e-3 大概率训练震荡。常见做法是先 warmup 让模型用小学习率适应预训练权重和新分类头的组合,再用余弦退火慢慢收敛。具体配置:

from torch.optim import AdamW from torch.optim.lr_scheduler import LinearLR, CosineAnnealingLR, SequentialLR optimizer = AdamW(model.parameters(), lr=1.5e-4, weight_decay=0.05) warmup = LinearLR(optimizer, start_factor=0.1, end_factor=1.0, total_iters=1000) cosine = CosineAnnealingLR(optimizer, T_max=49 * total_steps_per_epoch, eta_min=1e-5) scheduler = SequentialLR(optimizer, [warmup, cosine], milestones=[1000])

这里的 warmup 是 1000 个 step,不是 1000 个 epoch;如果脚本里一个 epoch 只有四百多个 batch,warmup 大约覆盖两个 epoch。total_iters 和 milestones 都要按“实际 optimizer.step 的次数”填,这是最容易抄错的地方。lr 取 1.5e-4 到 2e-4,weight_decay 取 0.05,是 ViT 微调里很常见的一组配置。

如果你担心 warmup 步数不准,也可以用一个更省心的写法:warmup 步数固定等于一个 epoch 的 step 数乘以 2,T_max 等于总步数减 warmup 步数。这样不管 batch_size 怎么改,学习率曲线的形状都相对稳定。

4.3 训练主循环:混合精度、梯度累积与断点续训

小显存也能把 ViT 跑起来,靠的是混合精度和梯度累积。混合精度把大部分计算切成 float16,显存和速度都受益;梯度累积把若干 batch 的梯度攒起来再更新一次,等效 batch size 变大。训练循环核心:

scaler = torch.cuda.amp.GradScaler() accum_steps = 4 model.train() for step, (x, y) in enumerate(train_loader): x, y = x.cuda(), y.cuda() with torch.autocast(device_type="cuda", dtype=torch.float16): logits = model(x) loss = criterion(logits, y) scaler.scale(loss).backward() if (step + 1) % accum_steps == 0: scaler.step(optimizer) scaler.update() optimizer.zero_grad() scheduler.step()

注意每个子 batch 前不要调用 optimizer.zero_grad(),只在累计完成后清零;如果写错,梯度在求和之外又被清零,训练就完全乱掉。混合精度出现问题不收敛时,先关掉做对照实验,不要在 fp16 下盲目调参。

断点续训是“后悔药”,检查点里至少保存 model、optimizer、scheduler、epoch 四样东西:

def save_checkpoint(ckpt_path, model, optimizer, scheduler, epoch): torch.save({ "model": model.state_dict(), "optimizer": optimizer.state_dict(), "scheduler": scheduler.state_dict(), "epoch": epoch, }, ckpt_path) def load_checkpoint(ckpt_path, model, optimizer, scheduler): ckpt = torch.load(ckpt_path, map_location="cuda") model.load_state_dict(ckpt["model"]) optimizer.load_state_dict(ckpt["optimizer"]) scheduler.load_state_dict(ckpt["scheduler"]) return ckpt["epoch"] + 1

只存 model 权重是很多源码包的偷懒写法,一旦你改了 optimizer 参数想回退,就只能从头训。把 scheduler 存进去,复现时才能严格接上学习率曲线,这是高分工程度上的重要差别。

4.4 关键超参速查表

参数常用值说明
img_size448 / 384patch_size=16 时,448 对应 784 个 token
patch_size168 更细但显存和计算量非线性增长
batch_size16配合梯度累积等效到大 batch
lr1.5e-4 ~ 2e-4AdamW 下超过 5e-4 就容易发散
warmup_steps1000按 optimizer.step 次数计
weight_decay0.05除分类头外都建议应用
label_smoothing0.1缓解过拟合,对易混病种友好
epochs50第 30 轮后准确率进入明显上升期是常态
cutmix alpha1.0从 Beta(1, 1) 采样

前几轮 loss 降得慢不用急,ViT 在小数据集上普遍要 20 轮之后才进入上升通道。如果加载了 ImageNet 预训练,50 epoch 内从 85% 左右冲到 90% 是很常见的走势。中途出现平台期就是没调好 warmup 和 lr 的典型信号,先回 4.2 检查 scheduler 而不是回模型结构里找问题。

5. 复现这个项目最容易翻车的 5 个坑:现象、原因与修复

模型和参数都到位,剩下的就是“复现成功”和“复现翻车”的分界。下面五条是我做木薯叶分类时几乎每次都会撞上的,按现象、原因、解决三步写清楚。

5.1 显存直接 OOM,batch_size 调到 4 都跑不动

现象:训练第一个 step 就报 CUDA out of memory,调小 batch_size 后仍然崩溃。

原因:ViT 的注意力矩阵是 N×N,448 × 448 输入带 [CLS] 是 785 个 token,单卡 batch 16 的注意力图就要吃掉好几 GB 显存。很多人以为换小 batch 就完事,但 transformer 的显存峰值经常出现在反向传播保存的中间张量上,batch=4 也救不回来。

解决:先确认是不是显存不够,把 batch_size 调到 4、关闭混合精度做基准;然后用梯度累积补回等效 batch,accum_steps=4 等效 batch 16;再不够就降输入分辨率,384 输入的 token 数从 784 降到 576,显存立刻降约 40%。如果还要激进,用 torch.utils.checkpoint 对 encoder 的 Block 做激活重计算,用计算换显存。按这个顺序能把项目跑起来,才算可交付。

5.2 训练 loss 一直在降,验证准确率却在震荡甚至下降

现象:loss 从 1.0 降到 0.4,验证集 top-1 却忽高忽低,每个 epoch 波动超过两个点。

原因:通常是学习率过大或没有 warmup。ViT 的高层分类头是随机初始化的,前几个 epoch 用大学习率会把预训练的特征分布撞歪;另一个高频原因是验证时忘了把模型切到 eval 模式,dropout 还在随机丢弃,验证结果自然抖。

解决:先确认训练循环里 model.train() 和 model.eval() 放对了位置;然后做一个小步长实验,把 lr 降到 1e-4,warmup 保持 1000 step 不放宽。如果验证仍然抖,回第 2 章查验证集是不是 stratify 划分,验证集图片太少同样会造成抖动。

5.3 五类不均衡:健康叶 recall 很高,病斑类被压得厉害

现象:整体准确率看着还行,按类别看混淆矩阵时发现 CBSD、HCBM 这类少数类 recall 只有 60% 上下。

原因:数据集中健康叶是多数类,标准交叉熵损失会偏向样本多的类别。病斑类样本少,梯度贡献被多数类淹没,学不好。

解决:给 CrossEntropyLoss 传 class weight:

import torch.nn as nn weights = torch.tensor([0.8, 1.2, 1.2, 1.2, 0.6]).cuda() criterion = nn.CrossEntropyLoss(weight=weights, label_smoothing=0.1)

更稳的做法是统计每类样本数的倒数再归一化,不要手写固定值。class weight 和 label smoothing 可以同时用,前者平衡样本量,后者抑制过拟合。改完后盯少数类的 recall,一般五到十个 epoch 就能看到明显改善。

5.4 数据加载把 GPU 饿死:训练一卡一卡,GPU 利用率上不去

现象:nvidia-smi 里 GPU 利用率在 30% 和 70% 之间来回跳,一个 epoch 耗时接近纯模型计算的两倍。

原因:木薯叶图片原始分辨率不低,每个 epoch 都要重新读 JPG、随机裁剪、resize 到 448。这些操作全在 CPU 数据管线里,一步慢就拖垮整体吞吐。

解决:把解码后的张量缓存进内存。常见做法是训练开始时把所有图像预读成 tensor 存进 list,getitem只做轻量操作:

class CachedCassavaDataset(torch.utils.data.Dataset): def __init__(self, df, img_dir, size=448, cache=True): self.files = [img_dir / f"{i}.jpg" for i in df["image_id"]] self.labels = df["label"].values self.size = size self.cache = cache and len(self.files) < 30000 if self.cache: self.images = [self._load(f) for f in self.files] def _load(self, path): import cv2 img = cv2.imread(str(path)) img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB) img = cv2.resize(img, (self.size, self.size)) return torch.from_numpy(img).permute(2, 0, 1) def __getitem__(self, idx): img = self.images[idx].clone() if self.cache else self._load(self.files[idx]) if torch.rand(1) > 0.5: img = torch.flip(img, dims=[-1]) return img, self.labels[idx]

缓存是拿空间换时间,16G 内存基本够放两万张 resize 后的整图。clone 不能省,否则随机翻转会污染缓存数据。内存不够时退一步,把预处理后的 .npy 落盘,原理一样,读写从 CPU 内存换到磁盘而已。

5.5 加载了 ImageNet 预训练却不如 ResNet50

现象:用 timm 加载预训练权重后微调 30 个 epoch,最后 top-1 比 ResNet50 还低两个点。

原因:多半是预训练权重加载时把 224 × 224 位置编码直接丢掉了,或者你的输入是 448 但没做插值,ViT 的位置信息整个乱掉。其次是新分类头从随机值起步,学习率又太小,新头发育太慢,拖累了整体收敛。

解决:不要因为 shape 不匹配就跳过 pos_embed,而是把 224 的位置编码插值到 448:

import torch.nn.functional as F def interpolate_pos_embed(pos_embed, new_num_patches): # pos_embed: 1, N+1, D old_tokens = pos_embed.shape[1] - 1 old_h = old_w = int(old_tokens ** 0.5) cls_pos = pos_embed[:, :1] grid_pos = pos_embed[:, 1:].reshape(1, old_h, old_w, -1).permute(0, 3, 1, 2) new_h = new_w = int(new_num_patches ** 0.5) new_grid = F.interpolate(grid_pos, size=(new_h, new_w), mode="bicubic", align_corners=False) pos_embed = torch.cat([cls_pos, new_grid.flatten(2).permute(0, 2, 1)], dim=1) return pos_embed

插值的作用是让新位置编码仍然保持“空间近邻的 patch 编码相近”这一语义。同时把新 head 的初始化方差适当调大,或者前几个 epoch 单独给 head 高一点学习率,收敛速度会明显变快。这条解决完,ViT 的优势才会真正体现出来,也才是你选择 transformer 而不是 ResNet 的理由。

6. 最后的进阶:用混淆矩阵与注意力热图验证模型到底学到了什么

训练完不要只打印 accuracy,用两个工具把模型“打开”看看。第一个是验证集上的混淆矩阵,第二个是注意力热图。

6.1 混淆矩阵:找出最容易混淆的叶片病害对

用 scikit-learn 一次跑完:

from sklearn.metrics import confusion_matrix, classification_report preds = [] labels = [] model.eval() with torch.no_grad(): for x, y in val_loader: logits = model(x.cuda()) preds.extend(logits.argmax(dim=1).cpu().tolist()) labels.extend(y.tolist()) print(classification_report(labels, preds, target_names=[f"class_{i}" for i in range(5)])) cm = confusion_matrix(labels, preds)

重点看矩阵里哪两类互相错得多。木薯叶数据中症状接近的病害类别天然难分,如果某两类混淆严重,下一步就是针对这两类加样本或做定向增强,而不是盲目调全局学习率。

6.2 注意力热图:模型到底在看病斑还是看背景

ViT 的可视化用最后一层 attention map 比 Grad-CAM 更自然。把最后一层注意力权重拿出来,取所有头对 [CLS] token 的权重平均,再 reshape 回二维:

attn_map = attn_weights[0, :, 0, 1:].mean(dim=0) # 所有头对 [CLS] 的平均 grid = int(attn_map.numel() ** 0.5) attn_img = attn_map.reshape(grid, grid).cpu().numpy()

把 attn_img resize 回原图大小,用 matplotlib 叠加到原图上。如果热图中心集中在病斑位置,说明模型在用叶片特征;如果集中在背景边缘,就去检查数据集是否存在背景偏差。这一步在答辩和汇报里是最能体现项目深度的素材。

6.3 从课设高分到能落地的边缘部署

如果项目是课程设计,做到混淆矩阵和热图已经足够;想再往前一步,把训练好的 ViT 导出为 ONNX 或 TorchScript。ViT 导出最容易踩的是动态尺寸问题,建议固定输入 448 × 448,opset 设 12 以上,导出后用 onnxruntime 跑一张图确认精度一致再谈部署。

我个人的习惯是每次实验都留一份文本记录:时间、数据划分 seed、学习率、增强开关、最后的混淆矩阵、这次和上次唯一改了什么。木薯叶病虫害分类这类项目,验收时经常被问的不是“准确率多少”,而是“位置编码怎么处理的”“验证集怎么划的”。把前面的插值逻辑和分层划分写进文档,比多压两个点的准确率更能体现工程完成度。希望这些经验能帮你在复现这个 transformer 源码包时少走几段弯路。

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

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

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

立即咨询