尧图网站设计 尧图网站设计YAOTU DESIGN
ARTICLE DETAIL

资讯详情

深耕网站设计与一线实操的经验洞察。

芒果成熟度图像分类实战:PyTorch+ResNet18数据集解析

芒果成熟度图像分类实战:PyTorch+ResNet18数据集解析 简介面向计算机视觉学习与研究的芒果成熟度图像分类数据集包含约9000张已标注图片按成熟、未成熟、损坏三种状态划分可支撑图像识别、目标检测、分类模型训练等方向。资源包为7z压缩格式共2000个文件其中1998张为JPG格式芒果图像另含1个Python脚本与1个JSON标注文件JSON文件储存了每张图片的类别标签及训练集、验证集、测试集划分信息便于直接按标准流程加载Python脚本则用于可视化数据集样本帮助直观检查图像质量与标注情况。压缩包整体约261.44MB体积适中适合在本地完成训练与验证。目前已有62人学习下载使用该数据可快速开展三类别的分类任务也可作为迁移学习或模型改进实验的起点尤其适合初学图像分类、准备芒果分选相关项目的开发者。1. 芒果成熟度数据集9000 张实拍标注图能解决什么做图像分类项目的人最清楚卡住进度的往往不是模型结构而是手里没有一份可靠的标注数据。这份芒果成熟度图像分类数据集我第一次解包时有点意外9000 张图对三分类任务来说谈不上海量但它不是白底渲染图或网络爬虫抓来的风景文件名里的 IMG_20230124 时间戳一眼就能看出是手机实拍背景里有装芒果的塑料筐、有桌面阴影、有不同时段的光线变化复杂程度比实验室数据集高不少。它能直接解决一个具体诉求不用再花几周时间收集芒果照片、请人逐张标注数据已经按成熟、未成熟、损坏分成三类并且划分好了训练集、验证集、测试集还带一个可视化脚本。适合要交第一个图像分类项目的新手快速跑通全流程也适合做农产品分级的从业者先拿这份数据跑出基线再决定是否补充自采样本。2. 拆开数据集结构与标注目录、JSON 和 show.py 可视化的正确姿势拿到数据集之后不要急着写网络、调参第一件事永远是把数据本身看清楚。很多翻车现场的根子就出在这个阶段文件名改过、标签错位、目录和 JSON 内容对不上。下面按我自己的拆包顺序来讲。2.1 目录结构背后的设计为什么按类别分文件夹从摘要信息和文件名特征来看这份数据是典型的图像分类标注布局也就是 ImageFolder 风格训练集、验证集、测试集三大目录底下再按类别建子目录所有成熟芒果放进同一个文件夹不用额外维护一份标签列表。这样的好处是 PyTorch 的datasets.ImageFolder可以直接读类别名就是文件夹名顺序按字典序排列不需要自己写 csv 解析逻辑。dataset ├── train │ ├── ripe │ │ ├── IMG_20230124_104343_jpg.rf.25933f35a47d0170d7ff91c43d249d93.jpg │ │ ├── IMG_20230116_105647_893_jpg.rf.e10e6ff3524599acd610fd87aa38b94a.jpg │ │ └── ... │ ├── unripe │ │ └── ... │ └── damaged │ └── ... ├── valid │ ├── ripe │ ├── unripe │ └── damaged ├── test │ ├── ripe │ ├── unripe │ └── damaged └── show.py注意文件名里那串_rf.25933f35...这是标注平台导出时留下的重命名痕迹rf是 Roboflow 的缩写后面是唯一标识。这类命名对训练没有影响但如果你之后要核对原图得认得出这些随机串来自导出改名不是原始拍摄文件名。为什么强调按类别分文件夹因为分类数据集的标签可靠性完全依赖目录名与图片内容一致。这份数据已经划分好了 train/valid/test比例上大致是常见的 7:2:1 或者 8:1:1具体划分不需要你操心但你要知道 test 集和 valid 集都独立存在这是为了后面的模型评估不污染。2.2 classes.json 怎么看标签名与文件名的对应规则摘要里明确写了具体查看 json 文件解包后要注意的是 JSON 里存的是类别名到索引的映射。常见格式是下面这种三个类别分别对应 0、1、2 三个索引{ ripe: 0, unripe: 1, damaged: 2 }注意如果 JSON 里给的是[ripe, unripe, damaged]这种数组那么类别索引就是数组下标。无论哪种写法你必须自己先读一遍确认它和 train 目录下的子文件夹一一对应。我遇到过不止一次JSON 里顺序是damaged, ripe, unripe而 ImageFolder 按字典序解析成damaged, ripe, unripe如果直接拿 JSON 的索引去匹配文件夹前几个 epoch 看起来没毛病后期调试类别输出时会发现预测结果和标签完全错位。正确的核对方式是写个两行脚本打印映射import json with open(dataset/classes.json) as f: labels json.load(f) print(labels)要么输出{ripe: 0, unripe: 1, damaged: 2}要么输出三个类别的 list。拿到之后先记下来后面写训练代码、画混淆矩阵时classes list(labels.keys())这个顺序必须和模型输出的 logits 索引一致。2.3 show.py 可视化先跑通它再碰训练代码资源里带了 show 脚本这个脚本的存在很有价值它逼着你先看数据。我的习惯是拿到任何数据集先随机抽几张图看真实外观确认标签和图像内容匹配再谈建模。这个 show 脚本的典型实现逻辑是从指定目录下每个类别随机抽 3 到 5 张图用 matplotlib 拼成网格展示。from PIL import Image import matplotlib.pyplot as plt import os import random root dataset/train classes [ripe, unripe, damaged] fig, axes plt.subplots(3, 3, figsize(9, 9)) for i, cls in enumerate(classes): folder os.path.join(root, cls) names random.sample(os.listdir(folder), 3) for j, name in enumerate(names): img Image.open(os.path.join(folder, name)) axes[i, j].imshow(img) axes[i, j].set_title(f{cls}: {name[:26]}) axes[i, j].axis(off) plt.tight_layout() plt.show()random.sample保证三个类别抽的数量一致避免某类图片太多导致抽样不均匀name[:26]截断长文件名因为之前看到的完整文件名可能长达 60 多个字符全打在标题里会把图挤得很小。提示如果你用的是中文类名成熟、未成熟、损坏matplotlib 默认字体可能显示成方块跑脚本前先执行plt.rcParams[font.sans-serif] [SimHei]。跑完 show 脚本你对每个类别的视觉特征就有底了。成熟果通常整体偏黄未成熟果偏绿损坏果表面有黑斑或霉变。但也别太乐观芒果这东西成熟度本身是一个连续变化过程放进三个离散类别里边界样本一定存在后面训练误差分析时还要回来处理它。3. 用 PyTorch 复现图像分类流程ResNet18 训练脚本与关键参数解析数据看明白之后就可以进入模型训练环节。这份数据最适合的模型不是最新最深的网络而是 ResNet18 这种小参数量结构。下面按选型理由、数据加载、训练超参数三个层面展开。3.1 ResNet18 的选型理由3 类 9000 张的算力与精度平衡虽然图像分类算法已经卷到 ViT、EfficientNet 一堆新结构但对 9000 张图像、三分类、还要在单卡上快速出结果的项目来说ResNet18 是最稳妥的起点。理由是泛化误差和算力成本之间的平衡芒果成熟度分类的类别区分特征主要集中在颜色和纹理上不需要特别深的语义抽象能力ResNet18 的残差连接已经足够拟合而 ViT 这类模型在小规模数据上反而容易过拟合除非你有充分的预训练权重和很重的数据增强否则跑起来又慢又不出效果。预训练权重要不要用我的建议是用。ImageNet 预训练的 ResNet18 已经学会底层边缘、颜色斑块、纹理基元我们只需替换最后一层全连接让它适配三个类别。这在真实项目里能明显缩短收敛时间也不容易在验证集上灾难性过拟合。3.2 数据加载与增强ImageFolder 读取和 transform 参数直接用ImageFolder加载 train、valid、test 三个目录无需手工写标签。关键在于 transform 的设置。from torchvision import datasets, transforms train_transform transforms.Compose([ transforms.Resize((224, 224)), transforms.RandomHorizontalFlip(p0.5), transforms.ColorJitter(brightness0.2, contrast0.2, saturation0.2), transforms.ToTensor(), transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]) ]) train_ds datasets.ImageFolder(dataset/train, train_transform) valid_ds datasets.ImageFolder(dataset/valid, transforms.Compose([ transforms.Resize((224, 224)), transforms.ToTensor(), transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]) ]))这里只做了水平翻转和颜色抖动没加随机裁剪、遮挡、旋转这一类重增强原因是芒果成熟度对颜色非常敏感成熟果偏黄、未熟果偏绿如果 ColorJitter 的亮度、对比度系数调太高会把原本偏黄的果变得偏绿等于在制造错误标签。色调饱和度的抖动幅度控制在 0.2 以内比较安全。Resize((224, 224))是 ResNet 系列的标准输入尺寸Normalize用的是 ImageNet 统计值因为用了预训练权重输入分布要和预训练保持一致这也是新手最容易漏掉的细节。3.3 训练脚本与超参数lr、batch_size、epoch 的推荐值模型、损失函数、优化器这一套组合要一起说清楚import torch from torchvision import models from torch import nn, optim model models.resnet18(pretrainedTrue) model.fc nn.Linear(model.fc.in_features, 3) device cuda if torch.cuda.is_available() else cpu model model.to(device) criterion nn.CrossEntropyLoss() optimizer optim.SGD(model.parameters(), lr3e-3, momentum0.9, weight_decay1e-4) scheduler optim.lr_scheduler.StepLR(optimizer, step_size3, gamma0.1)我用的是 SGD 而不是 Adam。不是说 Adam 不行而是在迁移学习场景下SGD 配小学习率对预训练权重的扰动更小最终收敛精度通常比 Adam 高一点。lr3e-3是针对整个网络都参与训练的情况如果你打算冻结前面几层只微调最后一层学习率可以提到 1e-2。StepLR每 3 个 epoch 把学习率降到原来的 0.1防止后期在损失曲面边缘震荡。train_loader torch.utils.data.DataLoader( train_ds, batch_size32, shuffleTrue, num_workers4) for epoch in range(10): model.train() running_loss 0.0 for images, labels in train_loader: images images.to(device) labels labels.to(device) optimizer.zero_grad() outputs model(images) loss criterion(outputs, labels) loss.backward() optimizer.step() running_loss loss.item() * images.size(0) scheduler.step() epoch_loss running_loss / len(train_ds) print(fepoch {epoch 1}: loss{epoch_loss:.4f})batch_size32在颜色差异明显的数据集上够用显存不够就降到 16num_workers4让数据加载不吃主力线程optimizer.zero_grad()必须在每个 batch 前清空梯度否则会在多个 batch 间累加。epoch 数建议从 10 开始观察训练损失下降趋势如果 5 个 epoch 内损失不再下降说明学习率需要调小或者数据量对于这个任务来说已经足够没必要硬跑几十轮。训练结束后把模型保存下来torch.save(model.state_dict(), mango_model.pth)注意只保存state_dict不保存整个模型对象这样换环境加载的时候更灵活也不会因为类路径不同报错。4. 评估和误差分析混淆矩阵暴露出的成熟度分类真实难点很多人在训练集损失降到很低之后就直接宣布项目完成这是不负责任的。分类任务必须看验证集和测试集表现而且要深入误差分布知道模型到底把哪一类认错了。4.1 用 sklearn 生成混淆矩阵哪一类被认错一目了然准确率在类别不平衡时会骗人混淆矩阵不会。我用这种方式快速生成测试集预测结果from sklearn.metrics import confusion_matrix, classification_report import seaborn as sns model.eval() y_true, y_pred [], [] with torch.no_grad(): for images, labels in test_loader: images images.to(device) outputs model(images) y_pred outputs.argmax(dim1).cpu().tolist() y_true labels.tolist() cm confusion_matrix(y_true, y_pred) print(precision/recall:) print(classification_report(y_true, y_pred, target_namesclasses)) sns.heatmap(cm, annotTrue, fmtd, xticklabelsclasses, yticklabelsclasses) plt.show()torch.no_grad()是推理阶段的必需操作关闭梯度计算既省显存又加速outputs.argmax(dim1)取每个样本 logits 最大的索引作为预测类别classification_report直接输出每个类别的精确率、召回率、F1 值。混淆矩阵中行代表真实类别列代表预测类别对角线越亮说明分类越准。4.2 边缘样本分析未成熟与损坏的视觉重叠是主要误差来源运行上面代码最可能看到的误差集中在两处一是未成熟果被预测成成熟二是损坏果被预测成未成熟。原因要回到数据本身去理解。未成熟的芒果如果已经有一点转色或者成熟果因为光照泛白二者在颜色特征空间里相当接近损坏果的早期病变可能只是表面一小块暗斑放在复杂背景的照片里模型很难把局部纹理和全局颜色整合起来判断。另外手机实拍的光线不统一同一颗芒果在阴影里和直射光下拍出来的颜色差异可能大于类别间差异。这类边缘样本不是模型结构能彻底解决的更现实的思路是确认是否符合业务场景。如果这个模型要部署到分拣线上那么分拣线通常是均匀光照下拍摄而数据集里的自然光照片会让模型对阴影过度敏感这时你需要在微调时加入针对性的光照增强或者干脆补充一批分拣线场景的图片。4.3 改进策略类别权重与难样本聚焦如果误差主要出现在某一类最直接的改进是给损失函数加类别权重。比如 damaged 类样本数量明显少或者它的召回率低于其他类class_weights torch.tensor([1.0, 1.0, 2.5]).to(device) criterion nn.CrossEntropyLoss(weightclass_weights)这个权重意味着把损坏类分错的代价提高到其他类的 2.5 倍模型会更偏向于把样本预测为损坏类从而提升该类的召回率。代价是其他类别的精确率可能略有下降要不要承受这个 trade-off取决于业务上更怕哪类错误。如果是分拣线漏掉损坏果比误杀好果更严重所以损坏类的权重应该大于 1。如果调权重还不够下一个手段是难样本挖掘把验证集里预测错的所有图片保存到单独文件夹人工观察是标签问题还是模型问题。这一步不写代码块用简单的 Python 遍历就能完成但价值很大——很多你以为的算法问题最后发现是标注本身就是错的。5. 避坑指南标注错位、数据泄露、类别不平衡的 4 个典型案例这部分是我在拆解类似数据集时反复踩过的坑按现象、原因、解决三段式写每条都可以直接对照。5.1 文件名对不上 JSON标注错位的现象与处理现象用 ImageFolder 训练时 loss 很反常明明 3 分类训练精度却长期在 40% 到 60% 之间波动验证集表现几乎是随机水平。原因数据从标注平台导出时文件名被整体重命名并追加了_rf.xxxxx随机串而 JSON 文件里的索引是按原始类名生成的。如果你的加载逻辑是自己读 JSON 再按索引匹配图片就可能出现标签和图像错位尤其当类别名称有中英文混用的情况。解决统一用ImageFolder的类名文件夹推导标签不手工创建映射如果必须手工读取 JSON先跑一遍核对脚本把误报的样本全部打印出来。还有一条铁律classes sorted(os.listdir(dataset/train))必须和 ImageFolder 的默认行为一致。5.2 训练精度奇高验证精度上不去数据泄露和增强过强现象训练集准确率很快到 98%验证集只有 60% 出头而且训练集准确率还在涨时验证集已经不涨甚至下降。原因最常见的是预处理泄露也就是验证集和测试集误用了训练集独有的增强或者统计信息。另一种是把RandomHorizontalFlip、ColorJitter也应用到验证集上导致验证集输入在每次 eval 时都随机变化结果极不稳定。解决验证集、测试集的 transform 里只保留Resize、ToTensor、Normalize不包含任何随机增强。训练集增强保持轻量ColorJitter 的值不要超过 0.2。5.3 损坏类样本数量少类别不平衡导致的整体性能虚高现象整体准确率看起来有 90% 以上但看 classification_report 才发现 damaged 类的召回率只有 50% 左右大部分损坏果被预测成了成熟或未成熟。原因成熟和未成熟这两类基数大模型只要把这两类学好了整体准确率就能被推高损坏类样本少梯度更新里它对损失的贡献也小模型自然倾向牺牲它。解决先统计每类样本数确认不平衡比例。如果损坏类占比低于 15%用CrossEntropyLoss(weighttorch.tensor([...]))给稀缺类别加权重或者对损坏类做 oversamplingWeightedRandomSampler。5.4 show.py 一运行就报缺包opencv 与 matplotlib 的依赖坑现象在干净环境里跑 show.py报ModuleNotFoundError: No module named cv2或者No module named matplotlib。原因这个脚本依赖 opencv 或 matplotlib但很多 Python 环境默认没有安装这些图像库还有一些情况是 conda 环境装了 opencv-python 但没装 pillow导入时报错信息很具迷惑性。解决按依赖文件安装我的习惯是直接用 requirementspip install opencv-python pillow matplotlib torch torchvision提示如果是在服务器上跑且没有显示器把 show.py 里的plt.show()改成plt.savefig(preview.png)不要硬开图形界面。6. 进阶验证Grad-CAM 把黑匣子变成可视化热力图模型精度达标之后应该再问一个问题它到底是看芒果的颜色判断成熟度还是看背景里的塑料筐颜色瞎蒙这个问题靠准确率回答不了要用 Grad-CAM 看模型的注意力区域。6.1 为什么要做 Grad-CAM精度高不等于看对了地方我在这个数据集上见过一种诡异情况因为训练集里损坏果大多装在深色塑料筐里模型学到的是只要中间有深色区域就是损坏而不是关注芒果表面的病斑。这种模型在测试集上准确率可能仍然不错但一换场景就崩。Grad-CAM 能定位模型做决策时最关注的图像区域帮你判断它学的到底是业务特征还是背景伪影。6.2 Grad-CAM 实现替换全连接层并注册 hookimport torch import torch.nn as nn import torchvision.transforms as transforms from PIL import Image import numpy as np import cv2 model models.resnet18(pretrainedTrue) model.fc nn.Linear(512, 3) model.load_state_dict(torch.load(mango_model.pth, map_locationcpu)) model.eval() class GradCAM: def __init__(self, model, target_layer): self.model model self.gradients None self.activations None target_layer.register_forward_hook(self.save_activation) target_layer.register_backward_hook(self.save_gradient) def save_activation(self, module, inp, out): self.activations out.detach() def save_gradient(self, module, grad_in, grad_out): self.gradients grad_out[0].detach() def generate(self, input_tensor): output self.model(input_tensor) pred output.argmax(dim1).item() self.model.zero_grad() output[0, pred].backward() weights self.gradients.mean(dim(2, 3)) cam (weights[:, :, None, None] * self.activations).sum(dim1) cam torch.relu(cam).squeeze().cpu().numpy() cam cv2.resize(cam, (224, 224)) cam (cam - cam.min()) / (cam.max() - cam.min() 1e-8) return cam, pred cam_layer model.layer4[-1] grad_cam GradCAM(model, cam_layer)register_forward_hook和register_backward_hook是 PyTorch 的钩子机制分别捕获前向传播的特征图和反向传播的梯度weights由梯度做全局平均得到代表每个通道对预测类别的重要程度之后把加权特征图做通道求和、ReLU、归一化就得到了 0 到 1 的热力图。6.3 三类样本的验收标准热力图该落在哪里拿一张成熟果、一张未成熟果、一张损坏果分别跑上面的代码检查注意力分布成熟果热力图应该覆盖果皮颜色最黄的区域背景基本不亮。未成熟果热力图应当集中在果皮偏绿的位置或者整个果体均匀激活。损坏果热力图应该聚集在表面黑斑、霉点附近而不是塑料筐边缘。如果成熟果的热力图大面积落在背景桌面上说明模型已经学到背景特征这组数据就算精度再高上线前都必须要重新筛图或做场景增强。这个验证步骤花的资源很少却能避免把模型部署到现场后才发现它冲着背景做判断。做这个数据集最大的教训是我一开始跳过 show 脚本直接跑训练跑出 92% 的准确率很兴奋结果画混淆矩阵才发现模型把大量浅色背景的图片全部判成成熟热力图一叠加模型根本不是在看芒果。从那以后我每次拿到标注数据都强制走一遍可视化、读 JSON、跑 Grad-CAM 这三步确认标签和注意力都对了才允许自己碰训练循环。希望今天这篇拆解也能帮你少走这一段弯路。本文还有配套的精品资源点击获取
返回列表