基于Transformer的木薯叶病虫害分类实战:从ViT原理到训练部署
发布时间:2026/10/1 13:19:54来源:尧图网络
简介一份基于Transformer模型的木薯叶病虫害分类Python源码适合机器学习、深度学习课程期末大作业或毕业设计参考。资源难度适中代码已通过本地编译验证可直接运行包含模型定义、数据集处理、全局变量配置、GPU调用等6个Python脚本另有5个编译缓存文件与1个Markdown说明文档压缩包整体约11KB结构简洁。已有199人学习下载。源码经过助教老师审定目录规划清晰能帮助读者快速理解Transformer在图像分类任务中的落地流程也可在现有模型与主程序基础上调整参数、更换数据集或加入更多评估指标便于二次开发与答辩展示。1. 基于 Transformer 的木薯叶病虫害分类这个高分源码包到底解决了什么做 python 木薯叶病虫害分类真正卡人的往往不是模型理论而是从数据读取到训练调参这条链路能不能一次跑通。这份基于 transformer 模型的木薯叶病虫害分类源码把数据加载、设备选择、模型定义、训练循环、断点保存拆成了独立模块本地编译可运行难度适中是典型的期末大作业高分结构。它解决的具体问题是给你 5 类木薯叶病害图像怎么用 Vision TransformerViT结构搭分类器并把训练流程写成能看懂、能复现的工程代码。木薯是热带地区重要粮食作物细菌性疫病、褐条病、绿斑驳病、花叶病这几类病害每年造成大量减产图像分类是自动化诊断的基础方案竞赛里常用的数据也是两万余张标注图像、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, typeint, defaultEPOCHS) parser.add_argument(--batch_size, typeint, defaultBATCH_SIZE) parser.add_argument(--resume, typestr, defaultNone) args parser.parse_args() device setup_device() train_loader, val_loader build_dataloader( batch_sizeargs.batch_size, num_workers4) model create_model(num_classesNUM_CLASSES, pretrainedTrue) train(model, train_loader, val_loader, device, args)逻辑说明build_dataloader 返回训练和验证两个 loadercreate_model 负责构造 ViTtrain 函数在 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/1614共 14x14196 个 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_size224, patch_size16, in_channels3, embed_dim768): super().__init__() self.num_patches (img_size // patch_size) ** 2 self.proj nn.Conv2d(in_channels, embed_dim, kernel_sizepatch_size, stridepatch_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_ratio4, dropout0.1): super().__init__() self.norm1 nn.LayerNorm(embed_dim) self.attn nn.MultiheadAttention(embed_dim, num_heads, dropoutdropout) 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_size224, patch_size16, num_classes5, embed_dim768, num_heads12, num_layers12): 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], dim1) # [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 和 labellabel 是 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, transformNone): self.img_dir img_dir self.label_df label_df.reset_index(dropTrue) 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(p0.5), transforms.RandomRotation(15), transforms.ColorJitter(brightness0.2, contrast0.2, saturation0.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 显卡有的只有 CPUGpu.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_sizeBATCH_SIZE, shuffleTrue, num_workers4, pin_memoryTrue) val_loader DataLoader(val_dataset, batch_sizeBATCH_SIZE, shuffleFalse, num_workers4, pin_memoryTrue)参数说明shuffleTrue 只用于训练集验证集必须 False否则每轮验证的样本顺序都在变不方便对比。num_workers 在 Windows 上经常踩坑如果报 spawn 相关错误说明主程序没包在if __name__ __main__:里或者直接设 0 最省心。pin_memoryTrue 能加速 GPU 训练时的 Host 到 Device 拷贝CPU 环境开了也无害。3.3 run.py 训练循环从 loss 到 checkpointrun.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(), lrLEARNING_RATE, weight_decayWEIGHT_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(fepoch {epoch1} batch {batch_idx1} floss {loss.item():.4f}) train_acc 100.0 * correct / total print(fepoch {epoch1}/{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 学习率与 warmupTransformer 最敏感的两个旋钮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_epochs5, total_epochs30): 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_lambdalr_lambda) for epoch in range(EPOCHS): train_one_epoch(...) scheduler.step()逻辑说明warmup 阶段系数小于 1比如第 0 个 epoch 实际学习率是 1e-4/52e-5到第 5 个 epoch 才达到完整学习率。余弦阶段让学习率平滑降到接近 0避免训练后期学习率过大导致 loss 震荡。warmup_epochs 一般取总 epoch 的 10%~20%30 个 epoch 取 3~5 个就够。4.2 batch size 与显存取舍先定 patch size再定 batchViT 的显存占用是三层叠加的patch 序列长度、batch size、embedding 维度。224 分辨率、patch 16 时序列长度是 19612 层 768 维编码器的显存占用明显高于同规模 CNN。显存不够优先降 BATCH_SIZE其次换成小配置。显存配置 EMBED/NUM_HEADS/NUM_LAYERSbatch sizepatch size6GB384 / 8 / 832166GB768 / 12 / 88168GB768 / 12 / 1281612GB768 / 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(fclass {c} acc: {class_correct[c]/max(class_total[c], 1):.2%})逻辑说明这段必须在 model.eval() 和 torch.no_grad() 下执行否则 dropout 的随机性会让验证结果不准还会白白占用显存记录计算图。类别名建议同时打印出来对照0 对应细菌性疫病 CBB1 对应褐条病 CBSD2 对应绿斑驳病 CGM3 对应花叶病 CMD4 对应健康叶片。如果发现某类 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 指定 encodingutf-8-sig 去 BOM并对 image_id 做 strip()。注意csv 的 label 列如果是从 Excel 导出的有可能被存成了文本格式读进来变成字符串 0 而不是整数 0训练时 CrossEntropyLoss 会直接类型报错顺手统一转 int。6. 模型落地单图推理脚本与混淆矩阵验证6.1 写一个独立的单张图像推理函数模型训完别急着收工。我习惯把 checkpoint 装回模型写一个独立于训练脚本的推理函数先证明它真的「用得上」def predict_single(model, image_path, device, transformval_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(dim1).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_locationdevice)[model_state_dict])map_location 保证在 CPU 机器上也能加载。切换到 eval 模式是必须的否则 dropout 会让同一张图每次预测结果不同。6.2 用混淆矩阵看 acc 背后的真实短板单张预测只能证明脚本能跑真正检验模型的是验证集上的混淆矩阵。如果健康叶片频频被误判成花叶病 CMD说明模型在病征不明显的样本上偏向多数类这种信息整体 acc 根本看不出来。把每个验证样本的 (label, pred) 收集起来统计成 5x5 矩阵后按行归一化看每一类的查全率比单个 acc 数字直观得多。从那以后我每次做完分类项目都会强制走一遍「单图推理、混淆矩阵、每类 acc」三步验证看起来多花十分钟但能拦住大多数「acc 好看、实际不能用」的假成功。这份源码拆下来真正有价值的不是 ViT 本身而是把训练、验证、落地的链路完整走通——你照着第 2 章到第 5 章的顺序检查一遍跑通是大概率事件希望能帮到你。本文还有配套的精品资源点击获取
网站建设高端定制企业官网