ARTICLE DETAIL

资讯详情

深耕网站建设、视觉设计与SEO优化的一线实战洞察。

四类害虫图像分类实战:数据体检与ResNet18微调全指南

四类害虫图像分类实战:数据体检与ResNet18微调全指南 简介面向农业病虫害智能识别与图像分类需求提供一套开箱即用的4分类庄稼害虫数据集涵盖蛀虫、健康无虫、螨虫等类别数据已按文件夹整理可直接用ImageFolder加载训练也可作为YOLOv5分类项目的数据输入。训练集train共620张、验证集test共53张目录结构清晰且经过可用性验证。资源总计676个文件以673张jpg图片为主体另含1份json分类字典、1个可视化py脚本和1张示例预览图整个压缩包约53.89MB。json字典记录了4种分类的映射关系便于训练时读取标签可视化脚本无需修改参数即可运行随机抽取4张图片展示并保存方便快速核查样本与标注效果。目前已有277人学习适合刚接触图像分类或需要现成数据集的开发者也适用于植保无人机虫害监测、农作物病虫害识别等实践场景。1. 4种庄稼害虫图像分类数据集能直接训练但别急着点按钮做农业植保或虫情监测的工程任务时你经常会收到这样一个压缩包里面是一份图像分类数据集4 种庄稼害虫训练集和验证集已经按文件夹分好。很多人拿到手就开训结果验证集精度上不去还以为是模型不行。这份数据真正值钱的地方在于训练集、验证集划分已经就位省掉了最脏最累的采集和标注环节但也正因为是别人分的数据泄漏、类别不均衡、标签噪声这类问题全都得自己做一遍体检。这篇笔记适合刚上手图像分类的算法工程师和做农业识别的从业者——讲的是怎么把这份四类害虫数据从“能跑”做到“可信”。2. 拿到数据先别训练四类害虫分类的难点与训练集验证集结构核对一个反直觉的事实四类害虫分类这种任务模型选型几乎不构成风险数据问题才是。我在多个小样本分类任务上得到过同一个结论——训练脚本 10 分钟能跑通数据问题能让你连续返工一两天。所以拿到数据集的第一件事不是把训练代码敲出来而是先把这两份目录训练集、验证集从头到尾看一眼。2.1 四类害虫图像分类真正的难点是什么图像分类任务里类别越少通常越简单但害虫是个例外。第一害虫体型小田间照片里往往只占画面的一小块模型很容易被叶片纹理、土壤颜色带偏第二同类害虫不同龄期的外观差异可能大于不同类别之间的差异幼虫和成虫放在一起形态变化大这又比 ImageNet 那种“一个类别一个稳定外观”要难第三训练集和验证集经常来自不同拍摄环境——室内白底、田间自然光、手机和单反混着来。这种时候最新图像分类模型的精度差距反而不是首要矛盾。ResNet18 和 EfficientNet-B2 在这类小数据上的差距通常小于数据清洗带来的差距。先把类别搞清楚、把脏图排掉再谈模型。2.2 核对目录结构训练集验证集用什么姿势组织拿到压缩包我一般先解压然后立刻确认目录层级。图像分类最常用的组织方式是 ImageFolder 风格根目录下每个类别一个文件夹文件夹里放这个类别的所有图片。这份数据大概率长这样dataset/ ├── train/ │ ├── class0/ │ ├── class1/ │ ├── class2/ │ └── class3/ └── val/ ├── class0/ ├── class1/ ├── class2/ └── class3/但“大概率”不等于“一定”。先跑一条命令确认训练集和验证集的类别目录是否对齐find dataset/train -maxdepth 2 -type d | sort | head -20 find dataset/val -maxdepth 2 -type d | sort | head -20这里有几个点要盯住。一是训练集和验证集的类别数、类别名必须完全一致别出现训练集 4 类、验证集 3 类这种低级错误二是文件夹名字决定了后续所有输出里的标签显示如果压缩包里给的是 0/1/2/3 这种编号建议先做一个 label_map.json 把编号映射到害虫中文名不然后面看混淆矩阵时全靠猜三是隐藏文件macOS 解压会带出 .DS_StoreWindows 会有 Thumbs.dbImageFolder 默认不认这些文件但有些数据里还会混进 .txt 说明文件最好在统计脚本里直接把非图片后缀过滤掉。常见做法是再跑一条命令快速看每类数量是否均匀for cls in train/*/; do echo $cls: $(ls -1 $cls | wc -l); done这条命令只打印数量更完整的统计放到下一小节。我自己踩过的坑是某次数据里 class0 的文件夹下还嵌了一层子目录ls 统计的数量全都不对好在这种问题在 ImageFolder 训练时会直接报错不至于悄悄带病运行。2.3 训练前必做的样本统计类别分布、图像尺寸、损坏文件这个脚本值得成为你每次拿到图像分类数据的固定动作。我第一次跑害虫数据时发现一件事某个类别的文件夹里混了 20 张从 PDF 截图导出的灰度图尺寸只有 96x96不检查根本发现不了最后这些低质量图全部让验证集精度掉了两个点。import os from collections import Counter from PIL import Image root dataset/train # 换成实际训练集路径 counts Counter() sizes [] broken [] for class_name in sorted(os.listdir(root)): cls_dir os.path.join(root, class_name) if not os.path.isdir(cls_dir): continue for img_name in os.listdir(cls_dir): if not img_name.lower().endswith((.jpg, .jpeg, .png)): continue img_path os.path.join(cls_dir, img_name) try: with Image.open(img_path) as im: im.verify() # 只校验文件完整性不真正解码 w, h im.size sizes.append((w, h)) except Exception: broken.append(img_path) counts[class_name] 1 print(类别分布) for c, n in counts.most_common(): print(f {c}: {n}) print(f图片总数{sum(counts.values())}) print(f损坏图片数{len(broken)}) for p in broken[:5]: print( , p) if sizes: ws [s[0] for s in sizes] hs [s[1] for s in sizes] print(f最小尺寸: {min(ws)}x{min(hs)} 最大尺寸: {max(ws)}x{max(hs)})脚本里 im.verify() 只检查文件头和解码完整性不解码全图速度快很多后缀过滤是为了防止文本文件混进来。日志里输出的“类别分布”能直接暴露类别不均衡问题比如某个类只有另外一类的三分之一时就要提前想好加权还是补充样本而不是等到训练完再去解释为什么偏向多数类。对训练集和验证集我是用同一套脚本分别跑一遍然后对比两个集合的类别比例是否接近。如果训练集里 class2 占 40%、验证集里 class2 只占 10%说明划分的时候没按类别比例分层抽样直接训练会导致 val 精度波动剧烈。分层抽样的做法是先按类别分组再从每类里按比例随机抽而不是对整个文件夹做一次 shuffle 再切成两半。2.4 验证集的“体检”不能省训练集干净了验证集同样不能跳过。验证集图片数量一般比训练集少更容易出现某些类只有十几张的情况。类别数量少的验证集还有个副作用精度指标的置信区间很宽某类的 recall 从 80% 掉到 60%可能只是因为它一共只有 10 张图、其中 2 张预测反了。遇到这种情况别急着调模型先确认验证集是不是太薄。3. 用预训练 ResNet18 训练四类害虫分类器完整脚本与参数调优3.1 迁移学习为什么用 ImageNet 预训练权重而不是从零训练四类害虫分类数据集通常只有几百到两三千张图。这个量级从零训练一个 CNN 是非常危险的模型会把训练集背下来验证集精度停在 50% 到 60% 甚至更差。而 ImageNet 预训练权重里已经包含大量通用视觉特征比如边缘、纹理、颜色分布、局部形状这些特征对害虫和叶片同样有效。迁移学习的常见做法是把预训练模型最后的分类型头拆掉换成一个新分类头去拟合你的 4 类微调时前几层基本不动后面几层和分类头重点调整。如果你追求更高精度EfficientNetV2、ConvNeXt 这类最新图像分类模型完全可以套同一个流程只需要改一行模型构造代码。但 ResNet18 是最稳的起点参数少、显存友好、微调收敛快用这块 4 类数据先跑通流程再换大模型比一开始就上 ViT 少踩很多坑。还有一个方向要单独说明。市面上大量教程在讲“处理数据集用于 YOLO 训练自己的数据集”那是目标检测路线需要类别框坐标作为标签。这份数据给的是整图分类标签没有框直接用 YOLO 等于白扔了标签信息如果后续确实需要从“有没有害虫”升级到“害虫在哪里”可以把这个分类模型当作预筛再单独采集带框数据两套东西分工而不是互相替代。3.2 训练脚本数据加载、增强与训练循环下面这个脚本基于 PyTorch torchvision一份数据集、两段代码可以完整跑通四类害虫分类。第一段数据变换与 DataLoader。# train.py 片段数据变换与 DataLoader from torch.utils.data import DataLoader from torchvision import datasets, transforms # 训练增强尺寸扰动 水平翻转 轻度颜色抖动 train_tf transforms.Compose([ transforms.RandomResizedCrop(224, scale(0.7, 1.0)), transforms.RandomHorizontalFlip(p0.5), 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_tf 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]), ]) # ImageFolder 要求 train/ 下每个类别一个子文件夹 train_ds datasets.ImageFolder(dataset/train, transformtrain_tf) val_ds datasets.ImageFolder(dataset/val, transformval_tf) train_loader DataLoader(train_ds, batch_size32, shuffleTrue, num_workers4, pin_memoryTrue) val_loader DataLoader(val_ds, batch_size32, shuffleFalse, num_workers4, pin_memoryTrue)RandomResizedCrop 的 scale 下限设 0.7是因为害虫在画面里本来就小缩放太狠会把虫子裁成几个像素模型学不到有效纹理。验证集用 Resize(256) CenterCrop(224) 是 torchvision 官方评估 ImageNet 模型的标准做法这个惯例保持住模型代码里的预训练归一化参数才能对得上。第二段模型、优化器与训练循环。# train.py 片段模型构建与训练循环 import torch import torch.nn as nn from torchvision import models # 用 ImageNet 预训练的 ResNet18只换最后一层分类头 model models.resnet18(weightsmodels.ResNet18_Weights.IMAGENET1K_V1) model.fc nn.Linear(model.fc.in_features, 4) # 4 类 model model.cuda() criterion nn.CrossEntropyLoss() optimizer torch.optim.SGD(model.parameters(), lr1e-3, momentum0.9, weight_decay1e-4) scheduler torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max30) best_val_acc 0.0 for epoch in range(30): model.train() train_loss 0.0 total 0 for images, labels in train_loader: images, labels images.cuda(), labels.cuda() optimizer.zero_grad() out model(images) loss criterion(out, labels) loss.backward() optimizer.step() train_loss loss.item() * images.size(0) total labels.size(0) # 每轮结束后在验证集上打分 model.eval() correct 0 val_total 0 with torch.no_grad(): for images, labels in val_loader: images, labels images.cuda(), labels.cuda() out model(images) pred out.argmax(dim1) correct (pred labels).sum().item() val_total labels.size(0) val_acc correct / val_total print(fepoch {epoch1:02d} | loss {train_loss/total:.4f} | val_acc {val_acc:.4f}) if val_acc best_val_acc: best_val_acc val_acc torch.save(model.state_dict(), best_pest_classifier.pth) print( - saved best)这段代码有四个关键点一是 weights 是新版 torchvision 的写法旧版本用 pretrainedTrue 等价替换二是优化器选 SGD 加 momentum图像分类微调场景里它通常比 Adam 更稳Adam 前期收敛快但后期容易在验证集上抖动三是 CosineAnnealingLR 的 T_max 设成和 epoch 数一致30 轮里学习率会从 1e-3 平滑降到接近 0比固定学习率省去了手动调衰减点的麻烦四是只在验证集提升时保存权重这是最便宜的“后悔药”哪怕后面训练崩了也能回滚。3.3 参数速查与微调顺序参数推荐起始值说明lr1e-3SGD 微调常用起点batch 翻倍时按比例调momentum0.9SGD 标配weight_decay1e-4防过拟合别设到 1e-2batch_size32显存不够就 16同时把 lr 降到 5e-4输入分辨率224ResNet 系列默认最大 epoch30看 val_acc 提前结束学习率调度CosineAnnealingLR(T_max30)或 step 每 10 轮乘 0.1常见的微调顺序是先固定 backbone 只训 fc 层跑 5 轮再解冻全网络微调。这个两阶段做法在小数据集上很实用能减少前期 loss 乱跳。我自己通常直接全网络微调配合低学习率效果也够这套顺序更省事。另外 Windows 下跑这段脚本要把整个训练逻辑包进if __name__ __main__:再调用否则 num_workers 多进程会报错。4. 让验证集精度稳住学习率、增强、加权损失的 5 个调节点第一轮训练跑完验证集精度可能卡在某个值上或者训练集精度一路涨、验证集纹丝不动。这时候别急着换模型从这 5 个调节点挨个排查大多数情况下问题出在这里面。4.1 学习率1e-3 起步配合余弦退火微调场景下学习率是最容易出问题的参数。lr 给高了新分类头会在最优解附近震荡验证集精度忽高忽低lr 给低了backbone 的特征基本没被调整验证集从 70% 爬到 75% 就不动了。判断方法很直接如果 loss 从一开始就很大并且过几个 epoch 还在振荡说明 lr 太高如果 loss 降得很慢验证集指标纹丝不动说明 lr 太低。4 类小数据集中SGD 的 lr 从 1e-3 起步是安全的配合 CosineAnnealingLR 衰减到 0。# 线性缩放规则从零训练时用 0.1 * batch / 256 # 迁移学习微调时不要套这个公式直接给定值更稳 base_lr 1e-3 batch_size 32 # 如果 batch 减半到 16通常把 lr 也减半 lr base_lr if batch_size 32 else base_lr * (batch_size / 32)代码里这个缩放逻辑只用作参考batch 变小意味着梯度估计的噪声变大学习率跟着降一点能稳定训练曲线。4.2 数据增强别过度害虫是小目标害虫在画面中的占比小这是选择增强策略时最需要注意的一点。通用的图像分类增强里流行加 RandomErasing 或 CutOut但在害虫数据上要非常谨慎——随机抹掉一块区域很可能正好把虫子本身抹掉模型被迫靠背景去猜类别验证集反而变差。# 容易踩坑的增强随机擦除可能抹掉害虫本体 transforms.RandomErasing(p0.5) # 常用做法轻度几何扰动 颜色扰动避免大尺度裁剪 transforms.RandomResizedCrop(224, scale(0.7, 1.0)) transforms.RandomHorizontalFlip(p0.5) transforms.ColorJitter(brightness0.2, contrast0.2, saturation0.2)RandomResizedCrop 的 scale 下限设 0.7是给害虫这类小目标留余地如果画面里虫子通常占比较大可以放宽到 0.5。颜色抖动用来模拟不同光照条件但 hue 参数别开太大绿色叶片变成紫色模型学到的就是假特征。旋转增强建议只用水平翻转垂直翻转会让虫子姿态违反自然状态不是不能用是收益不明确。4.3 类别不均衡加权损失优于无脑过采样4 类害虫数据的类别数量很少均匀。少数类只有几十张、多数类几百张时CrossEntropyLoss 会对多数类更友好验证集上少数类召回率明显偏低。加权损失的做法是按类别样本数的反比给 loss 加权import numpy as np import torch import torch.nn as nn # 每类样本数按 2.3 节统计脚本的输出填 samples_per_cls np.array([120, 300, 210, 460]) # 反比加权并归一化到均值 1保持 loss 量级不爆炸 weights 1.0 / samples_per_cls weights weights / weights.sum() * len(weights) weights torch.tensor(weights, dtypetorch.float32).cuda() criterion nn.CrossEntropyLoss(weightweights)归一化这步很重要。不归一化的话loss 的绝对值和原来差一个数量级学习率需要重新调很容易和 4.1 的问题混在一起。过采样把少数类图片重复复制也能用但每轮都要把少数类多看几遍训练时间变长而且重复样本容易让模型对少数类过拟合。加权 loss 代码改动最小优先试。4.4 训练轮数与早停别傻跑 50 轮小数据集上训练集精度会很快逼近 100%但验证集通常在第 10 到第 20 轮之间到顶。继续跑下去就是过拟合。用“验证集连续 N 轮不提升就停”的方式比固定跑满 50 轮更稳patience 5 wait 0 best_val_acc 0.0 for epoch in range(50): # ... 训练和验证代码同 3.2 ... if val_acc best_val_acc: best_val_acc val_acc wait 0 torch.save(model.state_dict(), best_pest_classifier.pth) else: wait 1 if wait patience: print(fepoch {epoch1}: val_acc 连续 {patience} 轮未提升提前停止) breakpatience 对 4 类小任务一般取 5 到 8太小容易在验证集精度正常波动时误停太大浪费时间。配合保存 best 权重的逻辑训练结束后拿到的永远是最优验证集状态而不是最后一轮的状态。4.5 验证集打分惯例单尺度、不开增强、关梯度验证集评估必须是一条固定管道否则 val_acc 本身就在随机跳动你没法判断调参到底有没有效果。我见过有人在验证集里也用了 RandomHorizontalFlip结果每次跑测试准确率都不一样最后查了半天才发现是评估代码的问题。三个细节验证集只用 Resize(256) CenterCrop(224)不做任何随机增强推理必须包在 torch.no_grad() 里否则 PyTorch 会为评估图额外建计算图显存开销大而且慢验证集 DataLoader 的 shuffle 设成 False不是为了准确率而是为了保持输出顺序稳定方便后续把预测结果和图片文件名一一对齐排查错误样本时会省很多事。5. 避坑清单训练集验证集上最常翻车的 5 个数据问题模型和损失函数调试是科学的但数据问题经常是玄学——翻车了查半天大概率不是模型结构的问题而是数据自己在搞鬼。以下 5 条按“现象 → 原因 → 解决”记录下来是我在多个图像分类数据集上反复遇到的真实情况。5.1 训练集精度 99%验证集精度只有 60%现象训练 loss 降到很低训练集准确率接近满分验证集 top-1 卡在 60% 到 70% 上不去。原因过拟合。数据量小、模型容量大、增强不够模型开始背训练集的图而不是学害虫的通用特征。害虫和作物背景高度耦合更容易出现这种情况。解决先检查训练集和验证集的图片数量比值超过 8:1 就要警惕然后按 4.2 适度加强数据增强把 weight_decay 从 1e-4 提到 5e-4或者把 ResNet18 换成更小的模型如 ResNet18 只解冻最后两个 block。还有一个隐蔽原因冻结 backbone 时 lr 给太高新分类头在震荡旧特征被破坏验证集同样上不去这时候先把 lr 降到 3e-4 再试。5.2 验证集精度比训练集还高别高兴先查泄漏现象验证集精度 95%训练集精度只有 90%明显反常数据划分肯定有问题。原因数据泄漏。最常见的是同一批次拍摄的连拍图片被同时分进了训练集和验证集模型等于提前见过了“考卷”另一种情况是训练集和验证集来自同一个视频的连续帧前后帧背景几乎一样。分类模型的泛化能力被高估了换个场景就现原形。解决先用 MD5 查完全重复的图片import hashlib from collections import defaultdict from pathlib import Path def md5(path: Path) - str: h hashlib.md5() with open(path, rb) as f: for chunk in iter(lambda: f.read(4096), b): h.update(chunk) return h.hexdigest() hash_map defaultdict(list) for p in Path(dataset).rglob(*): if p.suffix.lower() in (.jpg, .jpeg, .png): hash_map[md5(p)].append(str(p)) for h, paths in hash_map.items(): if len(paths) 1: print(h) for p in paths: print( , p)MD5 只能抓“完全相同”的复制做过缩放压缩的相近图抓不到最终判定还得靠人工抽查从训练集和验证集里各随机抽 20 张肉眼对比背景和环境。如果确实存在连拍泄漏正确做法是按拍摄批次重新划分数据集而不是重新随机抽一次。5.3 模型总把某个类判成另一个类现象整体准确率还行但混淆矩阵里某两个类互相踩。比如 A 类召回率 90%B 类召回率只有 55%而且 B 类的误判集中到了 A 类。原因两种可能。第一A、B 外观本就接近标注人区分它们也有分歧标签里存在噪声第二B 类样本太少模型没见过足够多的姿态变化和光照条件。解决把混淆矩阵里错误最多的图片对拉出来看具体输出方法在第 6 章。如果是标签噪声人工重标注 20 到 30 张就可以看到明显改善如果是样本少优先收集 B 类在不同光照、不同虫龄下的图片而不是盲目改模型。千万别在没看过错图之前就换损失函数那是白费力气。5.4 模型在“看”叶片和土壤而不是害虫现象验证集准确率不低但把模型拿到新的拍摄场地一测精度立刻崩。单独看预测结果发现模型的高置信预测和背景颜色高度相关。原因背景泄漏。统计规律上叶片纹理和害虫类别存在相关性模型学到了“看叶片颜色就能分类”没学到“看害虫本体”。这在田间数据里非常普遍因为不同类别害虫的拍摄场地和作物类型往往不同。解决用 Grad-CAM 或最直观的遮挡实验——把图片中间切一块黑色色块再送进模型看预测结果是否剧变。如果模型只靠背景判断遮挡害虫本体后预测可能不变。根治办法是让训练集覆盖不同背景或者至少保证验证集来自不同场地。如果暂时做不到公布精度时一定要注明数据场景别把一个场地训出的模型拿到另一个场地当通用模型用。5.5 换台机器重新训练验证集精度差两个点现象完全一样的代码、一样的数据两次训练结果不一致验证集精度有时差 1 到 2 个点。原因随机性。权重初始化、数据 shuffle、数据增强、cuDNN 的算法选择都引入了随机性。小数据集上这一点特别明显个别样本被分到哪个 batch 都可能改变收敛结果。解决固定随机种子import random import numpy as np import torch def set_seed(seed: int 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 set_seed()注意 cudnn.benchmarkFalse 会牺牲一点训练速度换取确定性DataLoader 每个 epoch 的 shuffle 还依赖一个 generator想完全复现需要同时固定它g torch.Generator() g.manual_seed(0) train_loader DataLoader(train_ds, batch_size32, shuffleTrue, generatorg, num_workers4)固定种子不是为了拿到“最好”的结果而是为了让你在调参时能区分“这个改动真的有效”还是“随机波动带来的假象”。我见过有人为了追验证集 0.5 个点的提升反复重训同一个配置浪费的算力足够把数据集再统计一遍。6. 验证集的正确用法混淆矩阵、阈值与单图推理验证集不是只用来输出一个 top-1 精度数字的。四类害虫任务里整体准确率会掩盖单类问题尤其是 5.3 说的类别混淆只有看混淆矩阵才能定位。6.1 输出混淆矩阵和每类指标import numpy as np from sklearn.metrics import confusion_matrix, classification_report # 收集验证集推理结果 all_preds, all_labels [], [] model.eval() with torch.no_grad(): for images, labels in val_loader: images images.cuda() out model(images) all_preds.extend(out.argmax(dim1).cpu().numpy()) all_labels.extend(labels.numpy()) print(classification_report(all_labels, all_preds, target_namestrain_ds.classes)) print(confusion_matrix(all_labels, all_preds))classification_report 里重点看每个类的 recall谁低谁就是主要矛盾。混淆矩阵的打印结果里非对角线数字大的格子就是 5.3 说的互踩类别对。6.2 置信度阈值宁可说“不确定”也不要说错农业场景里一个错误的害虫判断可能触发错误的打药指令比“无法判断”代价高得多。经验做法是给预测设一个置信度阈值低于阈值就返回“人工复核”probs torch.softmax(out, dim1) conf, pred probs.max(dim1) uncertain conf 0.6阈值取多少要在验证集上统计画出所有正确样本和错误样本的置信度分布找一个能保留 95% 正确样本的阈值。我实际用过 0.6 和 0.7最终选哪个要看你的业务对漏报和误报的容忍度。6.3 单图验证把模型从训练脚本里“抠”出来训练脚本里的 DataLoader 经历了完整的数据管道但线下演示或部署时只有一张裸图。单独写一个推理函数比每次改训练脚本省心得多def predict_one(path: str, model, class_names, val_tf): img Image.open(path).convert(RGB) img val_tf(img).unsqueeze(0).cuda() with torch.no_grad(): prob torch.softmax(model(img), dim1)[0] idx int(prob.argmax()) print(f{path}: {class_names[idx]} ({float(prob[idx]):.2%})) print(各概率:, dict(zip(class_names, prob.tolist())))这里最容易被忽略的是 val_tf 必须和第 3 章训练时的验证集 transform 完全一致Resize 尺寸、CenterCrop 尺寸、归一化的 mean/std 都不能改否则输入分布漂移精度会掉。我现在拿到任何一份图像分类数据第一件事永远是跑 2.3 节的统计脚本然后在训练前花两分钟确认训练集和验证集文件的来源。这套流程救过我很多次希望帮到你。本文还有配套的精品资源点击获取
返回列表