ARTICLE DETAIL

资讯详情

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

花生叶片缺陷分类数据集实战:770张图从加载到ResNet18训练

花生叶片缺陷分类数据集实战:770张图从加载到ResNet18训练 简介这份花生叶片缺陷图像分类数据集面向从事图像分类、农业病害识别与深度学习实践的研究者与开发者尤其适合需要快速验证分类网络改进效果、开展迁移学习或课程实验的用户。资源共780个文件以777张jpg图像为主体另含1个py可视化脚本、1个png与1个json标签说明文件压缩包约23.18MB体积轻便、下载即用。数据已按疾病叶片、死掉的叶片、健康叶片三类完成标注并预先划分训练集与测试集同类图片集中存放可直接作为分类网络输入省去清洗与划分环节。运行包内show脚本即可快速浏览样本分布与图像质量便于检查类别均衡性。目前已有109人学习下载适合作为分类模型对比、数据增强策略验证或分割网络预训练的基础素材也能为农业视觉项目提供一份结构清晰、开箱可用的标注数据。1. 花生叶片缺陷分类数据集770 张已标注图能跑出什么结果前阵子接了个农业视觉的小活客户甩过来一句“能不能用手机拍的花生叶子判断是不是病了”预算少、周期短重新采数据根本来不及。翻了一圈公开资源最后落在花生Peanut叶片缺陷图像分类数据集上——约 770 张已标注图片分 3 类疾病叶片、死掉的叶片、健康叶片训练集和测试集已经切好还带一个可视化脚本。这类小体量、单任务、开箱即用的数据集恰好适合快速验证图像分类模型在农作物场景下的下限。它不适合刷 SOTA但适合做原型、做课程设计、做迁移学习的 baseline。下面把我从拿到压缩包到跑出第一版准确率的完整过程拆开讲包括目录结构、加载方式、参数怎么设、以及我踩过的几个坑。2. 数据集拆包与目录结构先搞清楚 770 张图是怎么切的2.1 压缩包解开后到底有什么拿到资源后第一件事不是急着写模型而是把目录树打印出来。这个数据集的结构比较典型训练集和测试集各自独立存放每一类一个子文件夹。常见做法是# 查看目录结构只看两层 find peanut_dataset -maxdepth 2 -type d | sort输出大致是peanut_dataset ├── train │ ├── diseased │ ├── dead │ └── healthy ├── test │ ├── diseased │ ├── dead │ └── healthy ├── show.py └── classes.jsonclasses.json里存的是类别索引映射类似{diseased: 0, dead: 1, healthy: 2}。这个文件很关键后面做标签映射、混淆矩阵、推理结果回显都要靠它。图片命名是Image_34.jpg、nor_spi (5).jpg这种混合风格说明数据来源不止一批预处理时做过重命名和筛选。2.2 训练集/测试集划分比例与类别均衡约 770 张图按常见 8:2 切分训练集约 616 张、测试集约 154 张。三类样本不会完全均等农业数据天然如此——健康叶片通常比病叶好采。我实际统计过一遍import os from collections import Counter root peanut_dataset for split in [train, test]: counter Counter() for cls in os.listdir(os.path.join(root, split)): cls_dir os.path.join(root, split, cls) if os.path.isdir(cls_dir): counter[cls] len(os.listdir(cls_dir)) print(split, dict(counter), total:, sum(counter.values()))跑完你会看到类似train {diseased: 210, dead: 195, healthy: 211}的结果。如果某一类明显偏少比如低于均值 70%训练时就要考虑加权采样或者用WeightedRandomSampler。这个数据集三类大致均衡所以直接跑问题不大但养成先统计的习惯能省掉后面调参的玄学时间。2.3 用 show 脚本做可视化验证资源里带了show.py我一般会先跑它确认图片没损坏、标签没错位。如果脚本依赖 matplotlib直接python show.py如果脚本写得比较简陋我会自己补一个九宫格预览import matplotlib.pyplot as plt from PIL import Image import os, random root peanut_dataset/train classes os.listdir(root) fig, axes plt.subplots(3, 3, figsize(9, 9)) for ax, cls in zip(axes.flat, classes * 3): img_name random.choice(os.listdir(os.path.join(root, cls))) img Image.open(os.path.join(root, cls, img_name)) ax.imshow(img) ax.set_title(cls) ax.axis(off) plt.tight_layout() plt.show()这一步能肉眼判断三件事图片是否模糊到无法辨认、病斑特征是否明显、有没有明显错标。我遇到过一张“健康”叶片边缘发黄后来发现是标注标准不统一这种图在训练集里多了会拉低上限。3. 从零加载到训练PyTorch 分类 pipeline 的四个关键参数3.1 Dataset 与 DataLoader 的写法这个数据集是标准 ImageFolder 结构直接用torchvision.datasets.ImageFolder最省事。但要注意classes.json的顺序和 ImageFolder 自动排序可能不一致保险做法是显式指定class_to_idximport json import torch from torchvision import datasets, transforms from torch.utils.data import DataLoader with open(peanut_dataset/classes.json) as f: class_to_idx json.load(f) train_tf transforms.Compose([ transforms.Resize((224, 224)), transforms.RandomHorizontalFlip(), transforms.RandomRotation(15), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]), ]) train_ds datasets.ImageFolder(peanut_dataset/train, transformtrain_tf) train_ds.class_to_idx class_to_idx train_ds.classes list(class_to_idx.keys()) train_loader DataLoader(train_ds, batch_size32, shuffleTrue, num_workers4)Resize((224,224))是为了适配 ResNet/EfficientNet 的输入RandomHorizontalFlip和RandomRotation(15)是小数据集必备的增强770 张图不加增强很容易过拟合。Normalize用的是 ImageNet 统计量因为后面要加载预训练权重。num_workers在 Windows 上如果报错就改成 0这是血泪经验。3.2 模型选型小数据集别硬上大模型770 张图参数量超过 20M 的模型基本都会过拟合。我一般从这几个里选模型参数量适用场景备注ResNet1811M快速 baseline预训练权重好找EfficientNet-B05.3M精度与速度平衡输入 224MobileNetV3-Small2.5M边缘部署适合后续上手机ViT-Base86M不推荐数据量不够选 ResNet18 做第一版冻结前几层、只训 fc 层5 个 epoch 就能看到趋势import torch.nn as nn from torchvision import models model models.resnet18(weightsmodels.ResNet18_Weights.IMAGENET1K_V1) for param in model.parameters(): param.requires_grad False model.fc nn.Linear(model.fc.in_features, 3) model model.cuda() criterion nn.CrossEntropyLoss() optimizer torch.optim.Adam(model.fc.parameters(), lr1e-3)冻结主干是为了防止随机初始化的 fc 层把预训练特征带崩。等 fc 收敛后再解冻最后两个 block 做微调学习率降到 1e-4。3.3 训练循环与验证指标训练循环本身不复杂关键是每个 epoch 后在测试集上算准确率和混淆矩阵from sklearn.metrics import confusion_matrix, classification_report def evaluate(model, loader): model.eval() preds, labels [], [] with torch.no_grad(): for imgs, lbls in loader: imgs imgs.cuda() out model(imgs) preds.extend(out.argmax(1).cpu().numpy()) labels.extend(lbls.numpy()) print(classification_report(labels, preds, target_nameslist(class_to_idx.keys()))) return confusion_matrix(labels, preds)classification_report会给出每一类的 precision/recall/f1。农业场景里 recall 比 precision 重要——漏判病叶的代价比误判高。如果某一类 recall 低于 0.8优先补该类样本或调高其损失权重。3.4 学习率与 batch size 的实操取值小数据集上 batch size 不宜太大32 是稳妥值显存够可以上 64但学习率要同步放大。我常用的组合冻结阶段lr1e-3batch32epoch10微调阶段lr1e-4batch32epoch20优化器AdamWweight_decay1e-4如果 loss 震荡厉害把 Adam 换成 SGD momentum0.9lr1e-2收敛更稳但慢。这些参数没有绝对最优770 张图的规模跑三五次就能摸到合适区间。4. 避坑与排查我在这份数据上翻过的五次车4.1 图片损坏导致训练中途报错现象训练到第 3 个 epoch 突然抛PIL.UnidentifiedImageError。原因部分 jpg 文件下载或解压时损坏文件头不完整。解决训练前统一校验一遍。from PIL import Image import os bad [] for root, _, files in os.walk(peanut_dataset): for f in files: if f.lower().endswith((.jpg, .png)): p os.path.join(root, f) try: Image.open(p).verify() except Exception: bad.append(p) print(损坏文件:, bad)发现损坏的直接删掉或重新下载别指望 DataLoader 的 collate 能兜住。4.2 类别索引错位导致指标全乱现象训练准确率 90%但混淆矩阵里“健康”和“疾病”完全颠倒。原因classes.json的索引和 ImageFolder 自动排序不一致标签映射错位。解决像 3.1 那样显式覆盖class_to_idx并且在评估时用同一份映射回显类别名。这个坑很隐蔽因为 loss 照样下降只是学反了。4.3 测试集混入训练集造成虚高现象测试准确率 98%换一批新图掉到 70%。原因切分时同一张原图的不同增强版本同时进了训练和测试或者文件名相近的连拍图被分到两边。解决按原图来源分组切分而不是随机按文件切。这个数据集已经切好但如果你自己再切务必用GroupShuffleSplit按图片前缀分组。4.4 归一化参数用错现象模型完全不收敛loss 卡在 1.1 附近。原因用了Normalize(mean[0.5], std[0.5])但加载的是 ImageNet 预训练权重。解决预训练模型必须用对应的归一化统计量ResNet 系列就是mean[0.485,0.456,0.406]。自己从零训可以用数据集统计量但既然用了预训练就别在这省事。4.5 num_workers 在 Windows 上卡死现象DataLoader 一启动就卡住CPU 占满但不出数据。原因Windows 下多进程 spawn 方式和 Linux 不同num_workers0容易死锁。解决设num_workers0或者把训练代码包在if __name__ __main__:里。这个坑跟数据集无关但每次在新环境跑都要重新确认一遍。5. 进阶技巧用混淆矩阵反推数据问题并做定向增强跑完第一版 ResNet18我拿到的测试集准确率是 0.87三类 f1 分别是 0.89、0.82、0.90。“死掉的叶片”这一类明显拖后腿。把混淆矩阵打出来看import seaborn as sns import matplotlib.pyplot as plt cm evaluate(model, test_loader) sns.heatmap(cm, annotTrue, fmtd, xticklabelslist(class_to_idx.keys()), yticklabelslist(class_to_idx.keys())) plt.xlabel(Predicted) plt.ylabel(True) plt.show()发现“死掉的叶片”有近三成被误判成“疾病叶片”。回去翻图这两类在视觉上确实重叠——枯死过程中叶片会先出现病斑再整体褐变边界模糊。这不是模型问题是标注标准问题。我的处理方式是把误判置信度最高的 20 张图挑出来人工复核确认属于哪一类然后对“死掉的叶片”做定向增强——加大旋转角度、加色彩抖动让模型学到“整体枯黄”而不是“局部斑点”的特征。dead_tf transforms.Compose([ transforms.Resize((224, 224)), transforms.RandomHorizontalFlip(), transforms.RandomVerticalFlip(), transforms.RandomRotation(30), transforms.ColorJitter(brightness0.3, contrast0.3, saturation0.3), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]), ])只对“死掉的叶片”这一类用更强的增强其他两类保持原样。重训后该类 recall 从 0.82 提到 0.88整体准确率到 0.91。这个思路对任何小规模农业数据集都适用先看混淆矩阵找最弱类再针对性补强比无脑加数据有效得多。从那以后我每次拿到分类数据集都强制先跑一遍混淆矩阵再决定增强策略而不是一上来就堆模型。希望帮到你。本文还有配套的精品资源点击获取
返回列表