四类害虫图像分类实战:数据体检与ResNet18微调全指南
2026/9/23 6:47:45 网站建设 项目流程

简介:面向农业病虫害智能识别与图像分类需求,提供一套开箱即用的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_Store,Windows 会有 Thumbs.db,ImageFolder 默认不认这些文件,但有些数据里还会混进 .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(p=0.5), 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_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", transform=train_tf) val_ds = datasets.ImageFolder("dataset/val", transform=val_tf) train_loader = DataLoader(train_ds, batch_size=32, shuffle=True, num_workers=4, pin_memory=True) val_loader = DataLoader(val_ds, batch_size=32, shuffle=False, num_workers=4, pin_memory=True)

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(weights=models.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(), lr=1e-3, momentum=0.9, weight_decay=1e-4) scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=30) 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(dim=1) correct += (pred == labels).sum().item() val_total += labels.size(0) val_acc = correct / val_total print(f"epoch {epoch+1: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 的写法,旧版本用 pretrained=True 等价替换;二是优化器选 SGD 加 momentum,图像分类微调场景里它通常比 Adam 更稳,Adam 前期收敛快但后期容易在验证集上抖动;三是 CosineAnnealingLR 的 T_max 设成和 epoch 数一致,30 轮里学习率会从 1e-3 平滑降到接近 0,比固定学习率省去了手动调衰减点的麻烦;四是只在验证集提升时保存权重,这是最便宜的“后悔药”,哪怕后面训练崩了也能回滚。

3.3 参数速查与微调顺序

参数推荐起始值说明
lr1e-3SGD 微调常用起点;batch 翻倍时按比例调
momentum0.9SGD 标配
weight_decay1e-4防过拟合,别设到 1e-2
batch_size32显存不够就 16,同时把 lr 降到 5e-4
输入分辨率224ResNet 系列默认
最大 epoch30看 val_acc 提前结束
学习率调度CosineAnnealingLR(T_max=30)或 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(p=0.5) # 常用做法:轻度几何扰动 + 颜色扰动,避免大尺度裁剪 transforms.RandomResizedCrop(224, scale=(0.7, 1.0)) transforms.RandomHorizontalFlip(p=0.5) transforms.ColorJitter(brightness=0.2, contrast=0.2, saturation=0.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, dtype=torch.float32).cuda() criterion = nn.CrossEntropyLoss(weight=weights)

归一化这步很重要。不归一化的话,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(f"epoch {epoch+1}: val_acc 连续 {patience} 轮未提升,提前停止") break

patience 对 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.benchmark=False 会牺牲一点训练速度换取确定性;DataLoader 每个 epoch 的 shuffle 还依赖一个 generator,想完全复现需要同时固定它:

g = torch.Generator() g.manual_seed(0) train_loader = DataLoader(train_ds, batch_size=32, shuffle=True, generator=g, num_workers=4)

固定种子不是为了拿到“最好”的结果,而是为了让你在调参时能区分“这个改动真的有效”还是“随机波动带来的假象”。我见过有人为了追验证集 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(dim=1).cpu().numpy()) all_labels.extend(labels.numpy()) print(classification_report(all_labels, all_preds, target_names=train_ds.classes)) print(confusion_matrix(all_labels, all_preds))

classification_report 里重点看每个类的 recall:谁低谁就是主要矛盾。混淆矩阵的打印结果里,非对角线数字大的格子就是 5.3 说的互踩类别对。

6.2 置信度阈值:宁可说“不确定”,也不要说错

农业场景里,一个错误的害虫判断可能触发错误的打药指令,比“无法判断”代价高得多。经验做法是给预测设一个置信度阈值,低于阈值就返回“人工复核”:

probs = torch.softmax(out, dim=1) conf, pred = probs.max(dim=1) 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), dim=1)[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 节的统计脚本,然后在训练前花两分钟确认训练集和验证集文件的来源。这套流程救过我很多次,希望帮到你。

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

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

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

立即咨询