ARTICLE DETAIL

资讯详情

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

Transformer图像分类实战:木薯叶病虫害ViT源码复现与调优指南

Transformer图像分类实战:木薯叶病虫害ViT源码复现与调优指南 简介一份基于Python与Transformer模型的木薯叶病虫害分类源码面向深度学习初学者与期末大作业场景用于解决木薯叶图像识别与病害分类问题项目难度适中代码均已在本地编译运行。压缩包共12个文件含6个.py脚本、5个.pyc预编译文件及1个README说明整体仅11KB已有199人学习下载。源码经助教老师审定模块清晰主程序调度、全局变量配置、GPU适配、数据加载与Transformer分类网络一应俱全README给出运行指引。既适合期末大作业完整参考也便于迁移至其他农作物病害分类项目对理解注意力机制在视觉任务中的应用尤为有益。1. 木薯叶病虫害分类为什么这个视觉任务绕不开 Transformertransformer 图像分类在木薯叶病虫害数据上比 ResNet 能高出 25 个点的准确率这件事在公开比赛和不少落地项目里都被反复验证过。木薯叶的病斑往往很小、分布零散常规卷积网络容易把注意力浪费在整片叶子的纹理背景上而 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 2cassava_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(encodingutf-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 torch1.13 timm0.9 opencv-python pandas tqdm einops scikit-learn matplotlibPython 版本建议 3.83.10vscode 里把解释器切到刚建的 .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(parentsTrue, exist_okTrue) 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_size0.2, stratifydf[label], random_state42, ) df.loc[train_idx].to_csv(train_split.csv, indexFalse) df.loc[val_idx].to_csv(val_split.csv, indexFalse)stratifydf[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 维向量这个操作就是 PatchEmbedimport torch from torch import nn from einops import rearrange class PatchEmbed(nn.Module): def __init__(self, in_chans3, patch_size16, embed_dim768): super().__init__() self.patch_size patch_size self.proj nn.Conv2d( in_chans, embed_dim, kernel_sizepatch_size, stridepatch_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用 stridepatch_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, std0.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、valuequery 和所有 key 做点积并归一化得到注意力权重再对 value 加权求和。多头就是把 768 维切成 12 个 64 维子空间让每个头各管一种关系。代码实现如下class Attention(nn.Module): def __init__(self, dim, num_heads12): 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, biasFalse) 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 × NN 是 token 数这也是 ViT 显存开销大的直接原因。448 × 448 输入下 N 是 785batch 一放大注意力层立刻变成显存大户。因此后面训练章给的梯度累积方案就是为了缓解这个问题。多头注意力在木薯叶图片上具体捕捉什么以我的观察有些头会把病斑边缘与相邻叶片区域的 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_ratio4, dropout0.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-LNPre-LN 在深网络里梯度更稳这也是 ViT 能堆 12 层还能从预训练继续微调的原因。源码包里判断是手写还是封装看这个 Block 的实现就能一眼分辨。把前面几段组装成完整模型class ViT(nn.Module): def __init__(self, img_size448, patch_size16, embed_dim768, depth12, num_heads12, num_classes5): 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], dim1) 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_size384, num_classes5) pretrained timm.create_model(vit_base_patch16_224, pretrainedTrue).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_classes5) 更可控答辩时也能讲清楚“哪部分用了预训练、哪部分从零学”。如果输入不是 224上面的插值处理要留到后面。4. 木薯叶分类训练参数调优让准确率从 80 到 90 的配置与手段模型结构定了之后分数高不高基本全看训练策略。这个项目要从“能跑”到“高分”下面几个设置按顺序加效果远比反复换模型骨架来得快。参数不是一拍脑袋定的是 ViT 在中小规模图像分类任务上很通用的一组基线抄的时候照着调即可。4.1 数据增强CutMix 适合重叠病斑MixUp 会模糊健康叶木薯叶数据集的难点是病斑密度高、叶片相互遮挡常规的随机裁剪翻转不够用。源码包里常见的增强组合是 RandomResizedCrop、RandomHorizontalFlip、ColorJitter再配 CutMix。CutMix 把一张图的某个矩形区域换成另一张图标签按面积比例混合让模型必须同时关注叶片局部和整体。实现如下import torch def cutmix(x, y, alpha1.0): idx torch.randperm(x.shape[0], devicex.device) y_b y[idx] lam torch.distributions.Beta(alpha, alpha).sample() cx, cy torch.randint(x.shape[2], (2,), devicex.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(), lr1.5e-4, weight_decay0.05) warmup LinearLR(optimizer, start_factor0.1, end_factor1.0, total_iters1000) cosine CosineAnnealingLR(optimizer, T_max49 * total_steps_per_epoch, eta_min1e-5) scheduler SequentialLR(optimizer, [warmup, cosine], milestones[1000])这里的 warmup 是 1000 个 step不是 1000 个 epoch如果脚本里一个 epoch 只有四百多个 batchwarmup 大约覆盖两个 epoch。total_iters 和 milestones 都要按“实际 optimizer.step 的次数”填这是最容易抄错的地方。lr 取 1.5e-4 到 2e-4weight_decay 取 0.05是 ViT 微调里很常见的一组配置。如果你担心 warmup 步数不准也可以用一个更省心的写法warmup 步数固定等于一个 epoch 的 step 数乘以 2T_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_typecuda, dtypetorch.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_locationcuda) 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_size16 时448 对应 784 个 tokenpatch_size168 更细但显存和计算量非线性增长batch_size16配合梯度累积等效到大 batchlr1.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 显存直接 OOMbatch_size 调到 4 都跑不动现象训练第一个 step 就报 CUDA out of memory调小 batch_size 后仍然崩溃。原因ViT 的注意力矩阵是 N×N448 × 448 输入带 [CLS] 是 785 个 token单卡 batch 16 的注意力图就要吃掉好几 GB 显存。很多人以为换小 batch 就完事但 transformer 的显存峰值经常出现在反向传播保存的中间张量上batch4 也救不回来。解决先确认是不是显存不够把 batch_size 调到 4、关闭混合精度做基准然后用梯度累积补回等效 batchaccum_steps4 等效 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-4warmup 保持 1000 step 不放宽。如果验证仍然抖回第 2 章查验证集是不是 stratify 划分验证集图片太少同样会造成抖动。5.3 五类不均衡健康叶 recall 很高病斑类被压得厉害现象整体准确率看着还行按类别看混淆矩阵时发现 CBSD、HCBM 这类少数类 recall 只有 60% 上下。原因数据集中健康叶是多数类标准交叉熵损失会偏向样本多的类别。病斑类样本少梯度贡献被多数类淹没学不好。解决给 CrossEntropyLoss 传 class weightimport torch.nn as nn weights torch.tensor([0.8, 1.2, 1.2, 1.2, 0.6]).cuda() criterion nn.CrossEntropyLoss(weightweights, label_smoothing0.1)更稳的做法是统计每类样本数的倒数再归一化不要手写固定值。class weight 和 label smoothing 可以同时用前者平衡样本量后者抑制过拟合。改完后盯少数类的 recall一般五到十个 epoch 就能看到明显改善。5.4 数据加载把 GPU 饿死训练一卡一卡GPU 利用率上不去现象nvidia-smi 里 GPU 利用率在 30% 和 70% 之间来回跳一个 epoch 耗时接近纯模型计算的两倍。原因木薯叶图片原始分辨率不低每个 epoch 都要重新读 JPG、随机裁剪、resize 到 448。这些操作全在 CPU 数据管线里一步慢就拖垮整体吞吐。解决把解码后的张量缓存进内存。常见做法是训练开始时把所有图像预读成 tensor 存进 listgetitem只做轻量操作class CachedCassavaDataset(torch.utils.data.Dataset): def __init__(self, df, img_dir, size448, cacheTrue): 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 的位置编码插值到 448import torch.nn.functional as F def interpolate_pos_embed(pos_embed, new_num_patches): # pos_embed: 1, N1, 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), modebicubic, align_cornersFalse) pos_embed torch.cat([cls_pos, new_grid.flatten(2).permute(0, 2, 1)], dim1) 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(dim1).cpu().tolist()) labels.extend(y.tolist()) print(classification_report(labels, preds, target_names[fclass_{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(dim0) # 所有头对 [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 × 448opset 设 12 以上导出后用 onnxruntime 跑一张图确认精度一致再谈部署。我个人的习惯是每次实验都留一份文本记录时间、数据划分 seed、学习率、增强开关、最后的混淆矩阵、这次和上次唯一改了什么。木薯叶病虫害分类这类项目验收时经常被问的不是“准确率多少”而是“位置编码怎么处理的”“验证集怎么划的”。把前面的插值逻辑和分层划分写进文档比多压两个点的准确率更能体现工程完成度。希望这些经验能帮你在复现这个 transformer 源码包时少走几段弯路。本文还有配套的精品资源点击获取
返回列表