拓十年匠心定制 · 商业建站与技术教学双线并行 咨询热线:400-886-1026 service@lmnt.cn
ARTICLE DETAIL

资讯详情

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

基于Transformer的木薯叶病虫害分类实战:从ViT原理到训练部署

基于Transformer的木薯叶病虫害分类实战:从ViT原理到训练部署

简介:一份基于Transformer模型的木薯叶病虫害分类Python源码,适合机器学习、深度学习课程期末大作业或毕业设计参考。资源难度适中,代码已通过本地编译验证,可直接运行,包含模型定义、数据集处理、全局变量配置、GPU调用等6个Python脚本,另有5个编译缓存文件与1个Markdown说明文档,压缩包整体约11KB,结构简洁。已有199人学习下载。源码经过助教老师审定,目录规划清晰,能帮助读者快速理解Transformer在图像分类任务中的落地流程,也可在现有模型与主程序基础上调整参数、更换数据集或加入更多评估指标,便于二次开发与答辩展示。

1. 基于 Transformer 的木薯叶病虫害分类:这个高分源码包到底解决了什么

做 python 木薯叶病虫害分类,真正卡人的往往不是模型理论,而是从数据读取到训练调参这条链路能不能一次跑通。这份基于 transformer 模型的木薯叶病虫害分类源码,把数据加载、设备选择、模型定义、训练循环、断点保存拆成了独立模块,本地编译可运行,难度适中,是典型的期末大作业高分结构。

它解决的具体问题是:给你 5 类木薯叶病害图像,怎么用 Vision Transformer(ViT)结构搭分类器,并把训练流程写成能看懂、能复现的工程代码。木薯是热带地区重要粮食作物,细菌性疫病、褐条病、绿斑驳病、花叶病这几类病害每年造成大量减产,图像分类是自动化诊断的基础方案,竞赛里常用的数据也是两万余张标注图像、5 个类别。

适合三类人:交课程设计的学生、刚学完 transformer 想做图像分类实战的开发者、想拿现成 baseline 迁移到自有数据的从业者。后面我按「工程骨架 → 数据闭环 → 调参策略 → 踩坑记录 → 推理落地」的顺序拆开讲。

2. 拆解工程骨架:从 main.py 入口到 Model.py 的 ViT 实现

拿到这个 zip,别急着跑 run.py,先看 README,再把包里的文件按依赖关系排一遍。顺序基本是 Global_Variable.py → Gpu.py → CassavaDataset.py → Model.py → run.py → main.py。main.py 是总入口,run.py 是训练逻辑主体,前四个文件分别是配置、设备、数据、模型,__pycache__里的 pyc 是本地编译缓存,可以忽略。

这个分层对课程设计来说是最稳的写法:老师查重看结构,答辩问细节,你都能按模块讲清楚;自己调试时,改任何一块都不用动其他文件。

2.1 先看 main.py 和 run.py:入口与训练循环的分工

main.py 的职责是「组装」:解析参数、初始化设备、构建数据加载器、创建模型,然后把控制权交给 run.py。这类小型项目里最常见的入口写法是这样:

# main.py(常见写法,与包内文件结构对应) import argparse from Global_Variable import * from Gpu import setup_device from CassavaDataset import build_dataloader from Model import create_model from run import train if __name__ == '__main__': parser = argparse.ArgumentParser() parser.add_argument('--epochs', type=int, default=EPOCHS) parser.add_argument('--batch_size', type=int, default=BATCH_SIZE) parser.add_argument('--resume', type=str, default=None) args = parser.parse_args() device = setup_device() train_loader, val_loader = build_dataloader( batch_size=args.batch_size, num_workers=4) model = create_model(num_classes=NUM_CLASSES, pretrained=True) train(model, train_loader, val_loader, device, args)

逻辑说明:build_dataloader 返回训练和验证两个 loader,create_model 负责构造 ViT,train 函数在 run.py 里跑完整训练循环。args.resume 用于断点续训,这个点后面避坑章节会专门讲。

参数说明:--epochs 和 --batch_size 的默认值直接吃 Global_Variable 里的全局变量,命令行传参时优先,这样不用为跑一次小实验就改配置文件。

把入口和训练循环拆开还有个实际好处:你可以在 main.py 里自由替换数据源、模型、优化器,而训练逻辑本身完全不动,这对后续换自己的数据集非常关键。

2.2 Global_Variable.py:超参数集中管理为什么值得

Global_Variable.py 是这份源码里最不起眼但最实用的文件。它把所有超参数集中在一个位置,而不是散落在各函数里。助教审这类作业时,第一个看的就是这里是否清晰。

# Global_Variable.py 关键配置(常见参数值) IMG_SIZE = 224 # 输入图像统一缩放到 224x224 PATCH_SIZE = 16 # 每个 patch 的边长,224/16=14,共 14x14=196 个 patch EMBED_DIM = 768 # patch 投影后的 embedding 维度 NUM_HEADS = 12 # 多头注意力头数 NUM_LAYERS = 12 # Transformer 编码器层数 MLP_RATIO = 4 # FFN 隐藏层是 embed 维度的 4 倍 DROPOUT = 0.1 # 注意力与 FFN 里的 dropout NUM_CLASSES = 5 # 木薯叶病害类别数 BATCH_SIZE = 8 # 显存不够先降到 4 EPOCHS = 30 LEARNING_RATE = 1e-4 # ViT 比 CNN 更吃小学习率 WEIGHT_DECAY = 1e-4

逻辑说明:IMG_SIZE 和 PATCH_SIZE 决定了 patch 序列长度,这两个数必须能整除,否则 num_patches 算出来不是整数,Model.py 里直接报错。EMBED_DIM、NUM_HEADS、NUM_LAYERS 是决定模型容量的三件套,木薯叶这种 5 分类中等规模任务,用 768/12/12 这套 ViT-Base 量级配置已经偏大,想省显存可以降到 384/8/8。

参数说明:NUM_CLASSES 是唯一和任务强绑定的参数,换成你自己的数据集时只改它和路径两个地方。LEARNING_RATE 对 Transformer 尤其敏感,后面专门展开。

这种集中管理的模式,对课程设计最大的价值是答辩时能直接讲「我把学习率从 1e-3 调到 1e-4,收敛稳定了」,而不是支支吾吾说不清参数在哪改的。

2.3 Model.py:一个能跑通的简化 ViT 是怎么搭出来的

Model.py 是技术核心。基于 transformer 做图像分类,标准做法是 Vision Transformer:把图像切成 patch,每个 patch 线性投影成一个 token,送进标准 Transformer 编码器。Swin Transformer 是改进版,用了窗口注意力,但工程复杂度高不少;这个项目用原始 ViT,理由很实际——代码短、好讲、好调。

第一步是 Patch Embedding,用卷积实现是最简洁的写法:

class PatchEmbed(nn.Module): def __init__(self, img_size=224, patch_size=16, in_channels=3, embed_dim=768): super().__init__() self.num_patches = (img_size // patch_size) ** 2 self.proj = nn.Conv2d(in_channels, embed_dim, kernel_size=patch_size, stride=patch_size) def forward(self, x): # x: [B, 3, 224, 224] x = self.proj(x) # [B, embed_dim, 14, 14] x = x.flatten(2) # [B, embed_dim, 196] x = x.transpose(1, 2) # [B, 196, embed_dim] return x

逻辑说明:一个 16x16 卷积核、步长 16,等价于把图像切成 14x14 个不重叠 patch,每个被投影成 768 维向量。flatten 和 transpose 把卷积输出的 [B, 768, 14, 14] 重排成 Transformer 需要的 [B, 196, 768],196 是序列长度,768 是每个 token 的维度。

第二步是 Transformer 编码器层。PyTorch 里有现成的 nn.MultiheadAttention,但手写一层更容易看清结构,答辩也更好讲:

class TransformerEncoderLayer(nn.Module): def __init__(self, embed_dim, num_heads, mlp_ratio=4, dropout=0.1): super().__init__() self.norm1 = nn.LayerNorm(embed_dim) self.attn = nn.MultiheadAttention(embed_dim, num_heads, dropout=dropout) self.norm2 = nn.LayerNorm(embed_dim) self.mlp = nn.Sequential( nn.Linear(embed_dim, embed_dim * mlp_ratio), nn.GELU(), nn.Dropout(dropout), nn.Linear(embed_dim * mlp_ratio, embed_dim), nn.Dropout(dropout), ) def forward(self, x): x = x + self.attn(self.norm1(x), self.norm1(x), self.norm1(x))[0] x = x + self.mlp(self.norm2(x)) return x

逻辑说明:这是 Pre-LN 结构,先 LayerNorm 再进注意力,和原始 Transformer 论文的 Post-LN 相反。Pre-LN 在图像任务里收敛更稳,ViT 官方实现也是这么做的,这个细节值得在答辩时主动提一句。两个残差连接把梯度直接传给浅层,12 层堆叠也不容易梯度消失。

参数说明:nn.MultiheadAttention 的默认输入排列是 [seq_len, B, embed_dim],所以中间层传进去的 x 是 [196, B, 768],和 CNN 习惯的 [B, C, H, W] 完全不同,维度对不上时先检查是不是忘了 transpose。

第三步是组装完整 ViT,包含 cls token 和位置编码:

class ViT(nn.Module): def __init__(self, img_size=224, patch_size=16, num_classes=5, embed_dim=768, num_heads=12, num_layers=12): super().__init__() self.patch_embed = PatchEmbed(img_size, patch_size, 3, embed_dim) num_patches = self.patch_embed.num_patches self.cls_token = nn.Parameter(torch.zeros(1, 1, embed_dim)) self.pos_embed = nn.Parameter(torch.zeros(1, num_patches + 1, embed_dim)) self.blocks = nn.ModuleList([ TransformerEncoderLayer(embed_dim, num_heads) for _ in range(num_layers) ]) self.norm = nn.LayerNorm(embed_dim) self.head = nn.Linear(embed_dim, num_classes) def forward(self, x): x = self.patch_embed(x) # [B, 196, 768] cls_token = self.cls_token.expand(x.shape[0], -1, -1) x = torch.cat([cls_token, x], dim=1) # [B, 197, 768] x = x + self.pos_embed # 位置编码加到 token 上 for block in self.blocks: x = block(x) x = self.norm(x) return self.head(x[:, 0]) # 取 cls token 分类

逻辑说明:cls token 是 ViT 分类任务的关键设计——序列前拼接一个可学习向量,经过所有编码器层后,它的输出汇聚了全局信息,再接 Linear 头做分类。位置编码表有 197 行,对应 196 个 patch token 加 1 个 cls token,长度必须和序列一致,否则相加直接 shape 报错,这是新手最容易踩的坑。

整个 Model.py 控制在 100 行左右,可读性强,也方便换成 Swin Transformer 或加 DropPath 正则化。纯手写的好处是对每个张量的 shape 都有掌控,不像直接调 timm 库里的现成模型,出了问题只能当黑匣子猜。

3. 数据加载与设备调度:CassavaDataset 和 Gpu.py 的落地细节

数据集准备阶段通常有一个 train.csv,两列:image_id 和 label,label 是 0 到 4 的整数;图片放在 train_images 目录,文件名就是 image_id。CassavaDataset.py 的核心是把这两者对应起来。

3.1 CassavaDataset.py:图像读取、标签映射与数据增强

from torch.utils.data import Dataset from PIL import Image import os class CassavaDataset(Dataset): def __init__(self, img_dir, label_df, transform=None): self.img_dir = img_dir self.label_df = label_df.reset_index(drop=True) self.transform = transform def __len__(self): return len(self.label_df) def __getitem__(self, idx): img_name = self.label_df.loc[idx, 'image_id'] label = self.label_df.loc[idx, 'label'] img_path = os.path.join(self.img_dir, img_name) image = Image.open(img_path).convert('RGB') if self.transform: image = self.transform(image) return image, label

逻辑说明:getitem每次返回一个 (image, label) 对,DataLoader 会把 batch 个样本自动堆成张量。注意两个细节:reset_index 防止 csv 原始索引不连续导致 loc 取错行;convert('RGB') 强制三通道,防止数据里混进灰度图或带透明通道的 PNG,否则训练时通道数不一致会直接报错。

transform 是分类任务里容易被低估的部分。木薯叶拍摄环境差异很大,光照、角度、叶片遮挡都真实存在,增强策略直接影响泛化:

from torchvision import transforms train_transform = transforms.Compose([ transforms.Resize((224, 224)), transforms.RandomHorizontalFlip(p=0.5), transforms.RandomRotation(15), transforms.ColorJitter(brightness=0.2, contrast=0.2, saturation=0.2), transforms.ToTensor(), transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]), ]) val_transform = transforms.Compose([ transforms.Resize((224, 224)), transforms.ToTensor(), transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]), ])

参数说明:Resize 必须和 Global_Variable 里的 IMG_SIZE 一致。RandomRotation 角度不建议超过 20,转太多会把叶片背景噪声也学进去。Normalize 用 ImageNet 的 mean/std,因为常见做法是加载 ImageNet 预训练权重,迁移学习时归一化必须和预训练一致,否则前几层收到的分布和预训练时完全不同,微调效果会打折扣。

提示:验证集的 transform 里不要出现任何 Random 开头的增强,随机翻转和旋转只属于训练集,否则验证指标每次跑都不一样。

3.2 Gpu.py:设备选择的兜底逻辑

Gpu.py 在包里是个小文件,作用却很关键。课程设计的运行环境五花八门,有的有 NVIDIA 显卡,有的只有 CPU,Gpu.py 做的就是自动检测:

import torch def setup_device(): if torch.cuda.is_available(): device = torch.device('cuda') print('Using GPU:', torch.cuda.get_device_name(0)) else: device = torch.device('cpu') print('Using CPU, 训练会很慢') return device

逻辑说明:get_device_name 只是打印信息。实际开发里我一般还会补一行torch.backends.cudnn.benchmark = True,对固定输入尺寸的 ViT 能自动选最快卷积算法。如果是纯 CPU 环境,建议把 BATCH_SIZE 调到 4、EPOCHS 减半,先跑通流程再谈精度。

数据加载器这边的参数同样有讲究:

train_loader = DataLoader(train_dataset, batch_size=BATCH_SIZE, shuffle=True, num_workers=4, pin_memory=True) val_loader = DataLoader(val_dataset, batch_size=BATCH_SIZE, shuffle=False, num_workers=4, pin_memory=True)

参数说明:shuffle=True 只用于训练集,验证集必须 False,否则每轮验证的样本顺序都在变,不方便对比。num_workers 在 Windows 上经常踩坑,如果报 spawn 相关错误,说明主程序没包在if __name__ == '__main__':里,或者直接设 0 最省心。pin_memory=True 能加速 GPU 训练时的 Host 到 Device 拷贝,CPU 环境开了也无害。

3.3 run.py 训练循环:从 loss 到 checkpoint

run.py 里的训练循环是项目里最常被改的部分,核心是标准监督训练四步:

import torch import os from Global_Variable import EPOCHS, MODEL_SAVE_DIR, LOG_INTERVAL, LEARNING_RATE, WEIGHT_DECAY def train(model, train_loader, val_loader, device, args): criterion = torch.nn.CrossEntropyLoss() optimizer = torch.optim.AdamW(model.parameters(), lr=LEARNING_RATE, weight_decay=WEIGHT_DECAY) start_epoch = 0 best_acc = 0.0 for epoch in range(start_epoch, EPOCHS): model.train() running_loss = 0.0 correct = 0 total = 0 for batch_idx, (images, labels) in enumerate(train_loader): images, labels = images.to(device), labels.to(device) outputs = model(images) loss = criterion(outputs, labels) optimizer.zero_grad() loss.backward() optimizer.step() running_loss += loss.item() _, predicted = outputs.max(1) total += labels.size(0) correct += predicted.eq(labels).sum().item() if (batch_idx + 1) % LOG_INTERVAL == 0: print(f'epoch {epoch+1} batch {batch_idx+1} ' f'loss {loss.item():.4f}') train_acc = 100.0 * correct / total print(f'epoch {epoch+1}/{EPOCHS} avg_loss ' f'{running_loss/len(train_loader):.4f} acc {train_acc:.2f}%') if epoch % 5 == 0 or epoch == EPOCHS - 1: torch.save({ 'epoch': epoch + 1, 'model_state_dict': model.state_dict(), 'optimizer_state_dict': optimizer.state_dict(), 'best_acc': best_acc, }, os.path.join(MODEL_SAVE_DIR, 'checkpoint.pth')) return model

逻辑说明:这是监督分类的标准四步——前向算 loss、zero_grad 清梯度、backward 反传、step 更新参数。注意 loss.item() 会把值同步回 CPU,如果直接 print loss 张量,每一步都触发一次 GPU 同步,训练速度会肉眼可见变慢。

参数说明:CrossEntropyLoss 内部已经含 softmax 和 log,所以 model 最后一层 Linear 的输出直接喂 loss 即可,不用手动过 softmax。这也是它和 BCE 的根本区别——BCE 是二分类用的,这个项目是 5 分类多类任务,用 BCE 会出现 loss 在降但准确率永远不对的诡异现象。

checkpoint 保存成 dict 而不是只存 model,是为了把 epoch、优化器状态、历史 best_acc 都带上,这样才支持断点续训。每 5 个 epoch 存一次是防止频繁写盘;真正效果最好的模型建议单独存一份 best_model.pth,第 6 章的推理脚本会用到。

4. 参数调优与训练策略:让 ViT 在木薯叶数据上稳定收敛

4.1 学习率与 warmup:Transformer 最敏感的两个旋钮

ViT 和 CNN 调参差别很大。ResNet 用 1e-3 的 SGD 配动量能跑得不错,ViT 用同样配置大概率发散。常见做法是把学习率降到 1e-4 级别,优化器换 AdamW。原因是 Transformer 的注意力机制对梯度尺度更敏感,Adam 系优化器逐参数自适应,能显著降低这种敏感性。

实际项目中我一般还会加 warmup:前几个 epoch 让学习率从很小线性爬升到目标值,然后余弦衰减:

import math from torch.optim.lr_scheduler import LambdaLR def lr_lambda(epoch, warmup_epochs=5, total_epochs=30): if epoch < warmup_epochs: return (epoch + 1) / warmup_epochs progress = (epoch - warmup_epochs) / (total_epochs - warmup_epochs) return 0.5 * (1 + math.cos(math.pi * progress)) scheduler = LambdaLR(optimizer, lr_lambda=lr_lambda) for epoch in range(EPOCHS): train_one_epoch(...) scheduler.step()

逻辑说明:warmup 阶段系数小于 1,比如第 0 个 epoch 实际学习率是 1e-4/5=2e-5,到第 5 个 epoch 才达到完整学习率。余弦阶段让学习率平滑降到接近 0,避免训练后期学习率过大导致 loss 震荡。warmup_epochs 一般取总 epoch 的 10%~20%,30 个 epoch 取 3~5 个就够。

4.2 batch size 与显存取舍:先定 patch size,再定 batch

ViT 的显存占用是三层叠加的:patch 序列长度、batch size、embedding 维度。224 分辨率、patch 16 时序列长度是 196,12 层 768 维编码器的显存占用明显高于同规模 CNN。显存不够,优先降 BATCH_SIZE,其次换成小配置。

显存配置 EMBED/NUM_HEADS/NUM_LAYERSbatch sizepatch size
6GB384 / 8 / 83216
6GB768 / 12 / 8816
8GB768 / 12 / 12816
12GB768 / 12 / 121616

显存还是不够时的保底方案是梯度累积:

accum_steps = 2 optimizer.zero_grad() for i, (images, labels) in enumerate(train_loader): outputs = model(images) loss = criterion(outputs, labels) / accum_steps loss.backward() if (i + 1) % accum_steps == 0: optimizer.step() optimizer.zero_grad()

逻辑说明:每 2 个 batch 才更新一次参数,等效 batch size 翻倍,显存占用不变。loss 除以 accum_steps 是为了让累积后的梯度量级和正常 batch 一致。代价是训练时间变长,但 ViT 用的 LayerNorm 不受 batch size 影响,所以梯度累积对它是成立的,这也是它比 BatchNorm 类 CNN 更适合这个技巧的原因。

4.3 训练过程怎么看:loss 曲线与每类准确率

只看整体准确率很容易被类别不平衡骗过去。木薯叶数据里花叶病 CMD 的样本量明显多于其他几类,模型可能把多数类学得很好、少数类几乎全错,但整体 acc 依然好看。我习惯在每个 epoch 后单独统计每类准确率:

class_correct = [0] * NUM_CLASSES class_total = [0] * NUM_CLASSES model.eval() with torch.no_grad(): for images, labels in val_loader: images, labels = images.to(device), labels.to(device) outputs = model(images) _, predicted = outputs.max(1) for pred, target in zip(predicted, labels): class_total[target.item()] += 1 if pred.item() == target.item(): class_correct[target.item()] += 1 for c in range(NUM_CLASSES): print(f'class {c} acc: {class_correct[c]/max(class_total[c], 1):.2%}')

逻辑说明:这段必须在 model.eval() 和 torch.no_grad() 下执行,否则 dropout 的随机性会让验证结果不准,还会白白占用显存记录计算图。类别名建议同时打印出来对照:0 对应细菌性疫病 CBB,1 对应褐条病 CBSD,2 对应绿斑驳病 CGM,3 对应花叶病 CMD,4 对应健康叶片。

如果发现某类 acc 明显偏低,常见处理是给 CrossEntropyLoss 的 weight 参数按类别样本数反比加权,或对少数类做额外增强。优先用前者,代码改动最小。

5. 避坑排查:Transformer 分类项目五个常见翻车现场

这五条是我复现这类项目攒下来的血泪经验,每条都真实发生过,按「现象 → 原因 → 解决」写,你可以对照自己的报错快速定位。

5.1 训练 loss 不降反升

现象:第一个 epoch loss 就在 3.5 以上,之后一路涨或震荡,acc 徘徊在 20% 附近。

原因:学习率太大是首因。ViT 对学习率比 CNN 敏感得多,1e-3 的 AdamW 在能跑 CNN 的机器上直接让 ViT 发散。其次是优化器选错,用了 SGD 配动量而没有逐参数自适应,注意力权重更新步长失控。还有一种隐蔽情况:标签不是从 0 开始的连续整数,CrossEntropyLoss 默认按 0 到 num_classes-1 编码,标签里混进越界值会让 loss 异常高。

解决:把 LEARNING_RATE 降到 1e-4,优化器换成 AdamW;打印set(labels)确认标签集合是 {0,1,2,3,4}。仍然发散就加 warmup。

5.2 显存 OOM

现象:程序跑了几个 batch 后抛出 CUDA out of memory,有时报错在 loss.backward() 那行。

原因:ViT 反向传播要保存每层注意力矩阵,显存占用与序列长度平方相关。224 分辨率、patch 16 是 196 个 token,改成 384 分辨率就是 576 个 token,占用接近翻三倍。另一个高发原因是验证阶段忘了包 torch.no_grad(),计算图被完整保留。

解决:先降 BATCH_SIZE 到 4 验证能跑,再用梯度累积补回等效 batch;确认验证循环包了 no_grad()。还不行就把模型配置降到 384/8/8。

5.3 训练集 acc 95%,验证集只有 50%

现象:训练损失一路走低,训练 acc 逼近 90% 以上,val acc 却卡在 50% 左右不涨。

原因:典型过拟合,尤其在不加载预训练权重、从头训练 ViT 时更容易出现。木薯叶训练集只有两万余张,ViT-Base 容量大,从头训很容易把训练集细节背下来。另一个隐蔽原因是验证集 transform 里混进了随机增强。

解决:加载 ImageNet 预训练权重,只微调最后的分类头;加大 Dropout 和 WEIGHT_DECAY;验证集 transform 删掉所有 Random 开头的操作。

5.4 断点续训后指标突然回退

现象:用 --resume 加载 checkpoint 继续训练,第一个 epoch loss 明显高于上次记录,acc 也掉了一截。

原因:checkpoint 只存了 model_state_dict,没存 optimizer_state_dict 和 scheduler 状态。续训时 AdamW 的动量、学习率调度位置全部重置,相当于换了个新优化器重新起步,指标回退是必然的。

解决:保存时把优化器和调度器状态一并存进去,加载时同步恢复:

checkpoint = torch.load('checkpoint.pth') model.load_state_dict(checkpoint['model_state_dict']) optimizer.load_state_dict(checkpoint['optimizer_state_dict']) scheduler.load_state_dict(checkpoint['scheduler_state_dict']) start_epoch = checkpoint['epoch']

5.5 FileNotFoundError:数据路径拼接的坑

现象:Windows 上报路径不存在,但肉眼看着路径明明正确。

原因:csv 里的 image_id 可能已带扩展名(如 123.jpg),代码又拼了一次 .jpg,变成 123.jpg.jpg;或者 csv 解析出来带 BOM 头或首尾空格,字符串里藏着不可见字符导致路径无效;还有 Windows 和 Linux 路径分隔符混用的问题。

解决:打印拼接后的完整路径,用 os.path.exists() 逐个验证前 5 个样本;统一用 os.path.join 拼路径,不要手写字符串加斜杠;读 csv 指定 encoding='utf-8-sig' 去 BOM,并对 image_id 做 strip()。

注意:csv 的 label 列如果是从 Excel 导出的,有可能被存成了文本格式,读进来变成字符串 "0" 而不是整数 0,训练时 CrossEntropyLoss 会直接类型报错,顺手统一转 int。

6. 模型落地:单图推理脚本与混淆矩阵验证

6.1 写一个独立的单张图像推理函数

模型训完别急着收工。我习惯把 checkpoint 装回模型,写一个独立于训练脚本的推理函数,先证明它真的「用得上」:

def predict_single(model, image_path, device, transform=val_transform): image = Image.open(image_path).convert('RGB') image = transform(image).unsqueeze(0).to(device) model.eval() with torch.no_grad(): logits = model(image) pred = logits.argmax(dim=1).item() return pred class_names = ['CBB', 'CBSD', 'CGM', 'CMD', 'Healthy'] print(class_names[predict_single(model, 'test_images/100.jpg', device)])

逻辑说明:推理时直接看 logits 的 argmax,不需要先过 softmax,因为 argmax 是单调变换,不影响结果。加载权重时用model.load_state_dict(torch.load('best_model.pth', map_location=device)['model_state_dict']),map_location 保证在 CPU 机器上也能加载。切换到 eval 模式是必须的,否则 dropout 会让同一张图每次预测结果不同。

6.2 用混淆矩阵看 acc 背后的真实短板

单张预测只能证明脚本能跑,真正检验模型的是验证集上的混淆矩阵。如果健康叶片频频被误判成花叶病 CMD,说明模型在病征不明显的样本上偏向多数类,这种信息整体 acc 根本看不出来。把每个验证样本的 (label, pred) 收集起来,统计成 5x5 矩阵后按行归一化,看每一类的查全率,比单个 acc 数字直观得多。

从那以后我每次做完分类项目,都会强制走一遍「单图推理、混淆矩阵、每类 acc」三步验证,看起来多花十分钟,但能拦住大多数「acc 好看、实际不能用」的假成功。这份源码拆下来,真正有价值的不是 ViT 本身,而是把训练、验证、落地的链路完整走通——你照着第 2 章到第 5 章的顺序检查一遍,跑通是大概率事件,希望能帮到你。

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

返回列表