ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

7类果树水果小样本图像分类实战:500张图玩转PyTorch

7类果树水果小样本图像分类实战:500张图玩转PyTorch 简介这套7种常见果树水果图像分类数据集面向计算机视觉初学者与研究者解决训练分类网络时缺少已标注、可直接输入的高质量数据的问题。数据包含草莓、甜瓜、橙子、苹果等7个类别共约500张经过预处理的图片已划分训练集和测试集各自目录下按类别存放训练集与测试集保持独立便于评估模型泛化能力类别对应关系见JSON文件。压缩包共587个文件主要由584张JPG图像、1个Python可视化脚本、1张PNG示例图及1个JSON标签文件组成总大小约23.19MB目录结构清晰可直接加载训练。附带的show脚本可一键可视化数据集方便快速检查各类别样本分布JSON文件提供类别序号映射训练加载时无需手动整理。目前已有218人学习下载适合希望快速上手图像分类实战的开发者也可用于模型改进实验前的基准数据准备。1. 7种果树水果图像分类数据集500张已标注图片能做的事比你想的多很多人第一次跑图像分类用的是CIFAR-10或者MNIST数据均衡、背景干净、标签规范跟真实项目完全是两回事。而这个“7种常见长在果树上的水果图像分类数据集”只有约500张数据但每张都带标注恰好卡在“公开数据集太大、自己采集没标签”的尴尬区间。这个量级逼着你认真对待数据划分、数据增强和迁移学习而不是无脑堆模型。无论你是要做果园巡检、农产品分拣还是单纯想验证图像分类算法在受限数据下能不能稳住效果这份数据集都值得花一个晚上跑通。它最大的价值不是那500张图而是用这500张图逼你把小样本分类的完整流程走一遍。2. 7类果树水果数据集解构目录结构、标注格式与数据分布2.1 打开压缩包先看什么目录结构与标签文件拿到数据集之后的第一件事不是写模型而是先搞清楚它的目录组织方式和标签载体。常见做法有两种一种是按类别建文件夹比如apple、pear、peach各一个目录图片直接放在里面另一种是图片统一放在images目录再用一个CSV或JSON文件记录文件名和类别之间的对应关系。这两种方式决定了你后续写DataLoader时的路径不同也决定了你在做数据校验时该盯着哪儿查。我一般会先跑一段三行代码把目录树、每类图片数量和图片尺寸一次性确认掉。这一步可以帮你省下后面大量排错时间不然训练到一半突然报错你根本不知道是标签错了还是图片本身打不开。import os from PIL import Image data_root fruit_dataset categories sorted([d for d in os.listdir(data_root) if os.path.isdir(os.path.join(data_root, d))]) print(类别列表, categories) for cat in categories: cat_path os.path.join(data_root, cat) imgs [f for f in os.listdir(cat_path) if f.lower().endswith((.jpg, .jpeg, .png))] print(f{cat}: {len(imgs)} 张) if imgs: img Image.open(os.path.join(cat_path, imgs[0])) print(f 示例尺寸: {img.size}, 模式: {img.mode})这段脚本做了三件事第一列出根目录下所有子文件夹名称这就是类别标签的来源第二统计每个文件夹里的图片数量判断数据是否均衡第三打开每类第一张图片确认尺寸和颜色通道避免后面被灰度图或损坏图片干扰训练。参数说明里最需要注意的是后缀过滤我只统计jpg、jpeg、png三种格式因为图像分类数据集的常见打包方式几乎不会混入BMP或TIFF如果这里统计到的张数和你在下载页看到的数量对不上优先怀疑是不是图片文件名带了空格或者隐藏后缀。如果这个数据集用的是CSV标注而不是文件夹命名你还得额外确认CSV里的label字段和图片文件名能不能一一对上。很多开源数据集的CSV是用Excel编辑过的行尾多一个空格或者文件名大小写不一致经常让训练脚本在加载数据时静默跳过一部分样本。后面第5章我再详细说这个坑。2.2 7类水果的类别设置与标注粒度的两个极端这类果树水果数据集的类别标签通常落在“果实品种”和“果实状态”这两个粒度之间。举一个常见的组合苹果、梨、桃、葡萄、柑橘、香蕉、芒果七类。这里有两个需要你自己拿捏的问题第一标签是按品种分还是按成熟度分同一棵树上青苹果和红苹果如果被归成同一个类模型的分类难度会明显上升第二数据集里的“已标注”具体标到了什么程度是只有类别标签还是连果实边框、遮挡状态也标了如果只是图像分类任务那每个文件名对应一个类别就够了如果你本来想拿它做目标检测那这个数据集就不完全够用还得重新标注边框。在实际处理时我会先把类别名统一转成小写并去空格再做一次标签到索引的映射顺序固定下来以后不要再改。这个映射表要保存成JSON文件因为后续训练、验证、推理三个环节都要用到同一份映射任何一处不一致都会导致结果无法对齐。import json classes sorted(os.listdir(data_root)) class_to_idx {c.strip().lower(): i for i, c in enumerate(classes)} with open(class_to_idx.json, w, encodingutf-8) as f: json.dump(class_to_idx, f, indent2, ensure_asciiFalse) print(class_to_idx)这段代码把文件夹名转成小写去空格后映射到从0开始的索引然后保存成JSON。这里的核心是一旦映射关系写进文件后续所有脚本都只能从这份JSON读取不再从文件夹名重新推断。因为一旦你在训练脚本里临时用os.listdir排序顺序受文件系统影响可能不稳定同一个数据集两次跑出来标签互相错位这种错误最难排查。2.3 约500张数据在小样本场景下的真实分布500张数据、7个类别平均每类70张左右但实际拿到手很少有这么均匀的分布。有的类可能拍了120张有的只有30张这个不均衡直接影响你后续的准确率口径。我拿到数据集后会做一次分布统计然后把结果画成柱状图或者直接打印出来这样心里有数哪些类容易学哪些类天生样本少。样本量少的类别通常在训练时损失波动更大收敛也更慢。解决办法不外乎三类对少数类做过采样复制、用更强的数据增强、调整损失函数里的类别权重。这个数据集因为总量只有约500张我一般不建议先删样本做下采样那样会让有效信息雪上加霜。更务实的做法是让少数类在每次epoch里被反复看到配合随机裁剪、亮度扰动、水平翻转这类增强手段让模型看到的不是同一张图的简单复制。另外要额外确认一下各类别图片的背景复杂程度。果树场景下苹果可能带叶子、葡萄可能带藤蔓、桃可能带树枝这些背景信息会让模型学会“看到叶子就猜苹果”而不是“看到苹果形状才猜苹果”。这种问题在验证集上很难暴露因为验证集和训练集来自同一种拍摄习惯真正部署到果园新环境时才现原形。后面第5章的避坑部分我会再展开。3. 在PyTorch里跑通7类水果分类模型选型与最小复现流程3.1 为什么不用Transformer硬刚500张图图像分类模型这几年迭代很快ViT、Swin Transformer这些名字在热搜里出现的频率越来越高。但在一个约500张数据的小样本数据集上我几乎不会首选纯Transformer结构。原因很简单Transformer靠注意力机制建模全局关系需要大量数据来学习位置编码和注意力分布500张图连让ViT-Tiny吃饱都不够。强行上Transformer的结果通常是训练集loss一路下降验证集准确率卡在50%上下波动典型的过拟合。我的选型逻辑很直接参数量从小到大排序ResNet-18、MobileNetV3-Small、EfficientNet-B0这三个足够覆盖这个量级。ResNet-18是结构最透明的基线出了问题好查好修MobileNetV3在CPU上推理更快适合后面要部署到边缘设备的场景EfficientNet-B0在理论FLOPs上最省但训练时对输入分辨率更敏感。如果非要用Transformer唯一合理的方式是加载一个在ImageNet上预训练过的ViT冻结前几层只微调后段把它当特征提取器用而不是从头训。这里还有一个容易被忽视的点加载预训练权重时最后全连接层必须从1000类改成7类。很多人直接改成7就完事但忘了把fc层的权重重新随机初始化。PyTorch的torchvision模型在替换分类头时会自动初始化新层但如果你用自定义方式修改模型结构就得手动调用reset_parameters()否则模型会带着和类别数不匹配的旧权重跑前向传播第一轮就报维度错误。3.2 最小训练流程一个能跑通的PyTorch脚本下面这段训练脚本是我在500张规模的数据集上经常用的最小骨架省略了TensorBoard日志和早停便于你理解主干逻辑。import torch import torch.nn as nn from torch.utils.data import DataLoader from torchvision import datasets, transforms, models # 数据增强与归一化 train_tf transforms.Compose([ transforms.RandomResizedCrop(224, scale(0.6, 1.0)), transforms.RandomHorizontalFlip(), transforms.ColorJitter(0.3, 0.3, 0.3), 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]), ]) dataset datasets.ImageFolder(fruit_dataset, transformtrain_tf) train_size int(len(dataset) * 0.8) val_size len(dataset) - train_size train_ds, val_ds torch.utils.data.random_split(dataset, [train_size, val_size]) train_loader DataLoader(train_ds, batch_size16, shuffleTrue, num_workers2) val_loader DataLoader(val_ds, batch_size16, shuffleFalse, num_workers2) model models.resnet18(weightsmodels.ResNet18_Weights.IMAGENET1K_V1) model.fc nn.Linear(model.fc.in_features, 7) device torch.device(cuda if torch.cuda.is_available() else cpu) model model.to(device) criterion nn.CrossEntropyLoss() optimizer torch.optim.Adam(model.parameters(), lr1e-4) for epoch in range(15): model.train() total_loss, correct, total 0, 0, 0 for images, labels in train_loader: images, labels images.to(device), labels.to(device) optimizer.zero_grad() outputs model(images) loss criterion(outputs, labels) loss.backward() optimizer.step() total_loss loss.item() * images.size(0) _, preds outputs.max(1) correct preds.eq(labels).sum().item() total labels.size(0) train_acc 100.0 * correct / total model.eval() correct, total 0, 0 with torch.no_grad(): for images, labels in val_loader: images, labels images.to(device), labels.to(device) outputs model(images) _, preds outputs.max(1) correct preds.eq(labels).sum().item() total labels.size(0) val_acc 100.0 * correct / total print(fepoch {epoch1:02d} | train acc {train_acc:.2f}% | val acc {val_acc:.2f}%)逻辑说明这个脚本先用RandomResizedCrop做随机裁剪scale设为0.6到1.0相当于每张图在每次epoch里以不同的比例和位置被裁出来这是小数据集上最有效的增强手段之一。然后用ImageFolder直接按子文件夹名生成标签省去了手写Dataset子类的工作。random_split按80/20比例把数据划分成训练集和验证集这里有一个隐含风险random_split是在整体样本上随机抽取如果某个类别的图片集中在目录末尾可能出现验证集中缺了某一类。后面我会给你更稳妥的划分方法。参数说明输入尺寸224对应ImageNet预训练权重的标准分辨率batch_size取16是为了让一个batch里尽量包含更多类别学习率1e-4是我在迁移学习场景下的常用起点比从头训练小一个量级因为预训练权重已经处在一个比较好的局部最优附近学习率太大会把权重带偏。优化器选Adam而不是SGD原因是小数据集上Adam收敛更稳定对学习率不那么敏感适合快速验证模型和流程正确性。3.3 三个必调参数输入尺寸、batch_size、学习率这三个参数是图像分类项目里改动频率最高、影响也最直接的而且三个参数互相耦合不是孤立调参。输入尺寸方面ResNet-18用224是默认值但如果你发现原图里果实占比很小缩放到224之后细节全丢可以把输入调成128或者保留300左右的长边再中心裁剪。需要同步调整的是RandomResizedCrop的scale下限和模型的输入通道结构。更小的输入尺寸意味着训练更快也意味着模型更难捕捉细粒度差异比如区分两种颜色接近的梨和苹果。batch_size在小数据集上取16到32比较合适。取太大一个batch可能包含某一类过多梯度方向偏向该类取太小比如4每次参数更新的噪声太大损失曲线会抖得很厉害。如果你的显存充足可以试试32同时把学习率象征性地调大一点因为更大的batch对梯度估计更准允许更大的步长。学习率是三个参数里最敏感的一个。我见过的翻车现场十有八九是学习率直接用了默认的0.01甚至0.1预训练权重被几轮迭代就冲坏验证集准确率从80%掉到30%再也回不来。在迁移学习场景下1e-4是一个安全起点如果损失在头几个epoch里完全不动再降到3e-5如果训练集准确率上升很快但验证集原地踏步说明学习率偏大模型在记忆训练集。还有一种做法是先用一个小脚本做学习率扫描从1e-5到1e-2分成五个档每档跑两三个epoch看验证集趋势不玄学很直接。4. 让500张数据发挥更大价值迁移学习、数据增强与评估4.1 迁移学习的正确姿势冻结哪几层取决于你的数据量迁移学习在小样本水果分类里是决定性因素。用ImageNet预训练权重初始化ResNet-18然后在这种7类水果数据上微调通常比从头训练高出20到30个百分点的验证集准确率。原因不难理解ImageNet已经让网络学会了边缘、纹理、果实形状这些通用视觉特征而你的任务只是在这些特征之上重新组合分类边界。冻结策略有讲究。数据量越少越应该冻结靠前的层。ResNet-18分四个stage前两个stage学到的是边缘、颜色块、基础纹理这些特征在水果分类里同样适用不用微调。我一般的做法是当每类样本少于100张时冻结layer1和layer2只微调layer3、layer4和最后的全连接层如果每类样本只有三五十张那就连layer3也冻结只让layer4和fc层去学。判断依据是验证集准确率如果冻结后验证集表现比全部微调好说明数据量撑不起那么多可学习参数。for name, param in model.named_parameters(): if name.startswith(layer1) or name.startswith(layer2): param.requires_grad False这段代码把所有layer1和layer2的权重冻结反向传播不会更新它们只更新layer3、layer4和fc层。注意这里没有冻结bn层的running_mean和running_stat如果数据分布和ImageNet差异很大bn层的统计量默认是按ImageNet分布算的会让验证集出现诡异波动。稳妥做法是把bn层也设成eval模式或者保留requires_grad具体要看训练loss是否稳定。4.2 数据增强与归一化别让你的增强策略毁掉模型数据增强是500张数据集的第二根救命稻草。但增强不是越多越好有些增强对水果分类是有害的。比如RandomErasing随机擦除一块区域如果擦掉的是果实中心模型会被迫学习不完整的特征在样本本来就少的情况下反而学歪。再比如旋转角度设得太大苹果被旋转90度仍然像苹果但葡萄串旋转90度之后的形状分布变化很大模型可能会混淆。我的增强组合是RandomResizedCrop配合RandomHorizontalFlip再加轻微ColorJitter和RandomRotation。水平翻转对大多数水果没有方向性影响安全ColorJitter的亮度、对比度扰动能模拟果园里不同时段的光照对于提高泛化能力很有帮助。RandomRotation的度数我控制在15度以内避免过度改变果实朝向。归一化必须用ImageNet的均值方差这点不能改。因为预训练权重是在ImageNet的归一化条件下学出来的你如果用了自定义的mean和std相当于输入分布直接偏移预训练权重就废了一半。这里有个常见的低级错误在展示图片时忘记反归一化图片看起来泛白但那不是模型的问题是显示问题。训练代码里的归一化一直保留ImageNet参数只在可视化时反向操作回来。4.3 评估不看单一准确率混淆矩阵和单类召回500张数据的验证集大约100张平均每类14张左右。这种情况下一个准确率数字很不靠谱某类图片只有8张模型全错也只影响整体准确率几个百分点但单独看那一类的召回率是0%。我每次训练完都会打印一份混淆矩阵看看哪些类别互相被认错。常见的水果分类混淆模式有两种颜色相近的比如青苹果和梨形状相近的比如桃和苹果。混淆矩阵能直观地告诉你模型到底是在“看颜色”还是“看形状”。如果青苹果大量被预测成梨说明模型更依赖颜色特征而不是轮廓特征这时可以在增强里加入灰度化选项强迫模型学习形状。from sklearn.metrics import confusion_matrix import numpy as np all_preds, all_labels [], [] model.eval() with torch.no_grad(): for images, labels in val_loader: images images.to(device) outputs model(images) _, preds outputs.max(1) all_preds.extend(preds.cpu().tolist()) all_labels.extend(labels.tolist()) cm confusion_matrix(all_labels, all_preds) print(混淆矩阵) print(cm)这段代码把验证集所有预测结果收集起来用sklearn的confusion_matrix生成矩阵。行是真实标签列是预测标签对角线越大的类别学得越好。参数说明里有一个容易忽略的点收集预测时要用.cpu().tolist()否则GPU上的张量会一直占着显存验证集图片较多时可能导致显存溢出尤其在后面叠加多个测试集时特别常见。5. 避坑指南小样本水果分类最常见的5个事故5.1 训练准确率99%、验证准确率60%这不是BUG是过拟合现象每个epoch打印结果显示训练准确率一路飙到接近100%验证准确率却徘徊在60%。这不是代码写错了也不是数据集有问题而是模型把你的训练集背下来了。500张数据太少ResNet-18的参数量足够记住所有训练样本它不需要泛化就能把训练集准确率刷满。原因模型容量远大于任务复杂度同时训练集和验证集来自同一个拍摄批次背景、光线、构图高度相似模型稍微记住一点背景特征就能在训练集上得分。解决先用第4章的冻结策略把大部分参数固定住只微调最后一两个stage。如果冻结后还过拟合下一招是把训练数据增强强度调大比如把RandomResizedCrop的scale下限降到0.4让模型每次看到的都是图像的不同局部。最后一道防线是加Dropout或者权重衰减ResNet-18本身没有Dropout层你可以在全连接层前插入一个dropout0.3的层很多人忘了这一步。5.2 某一类只有30张图另一类有120张类别不均衡现象少数类的验证召回率特别低准确率整体看着还不错但看每个类别的F1分数时某个水果类直接掉到接近0。原因CrossEntropyLoss默认对所有类别一视同仁样本多的类贡献更大的梯度模型倾向于把不确定样本预测到多数类里。解决最简单的做法是在损失函数里传入类别权重。用每类样本数的倒数做归一化让少数类的每个样本在计算loss时被放大。class_counts torch.tensor([30, 45, 120, 60, 80, 95, 70], dtypetorch.float) weights 1.0 / class_counts weights weights / weights.sum() * len(class_counts) criterion nn.CrossEntropyLoss(weightweights.to(device))逻辑说明weights先取每类的倒数再归一化到平均值为1。这样样本少的类别在loss中的占比会上升模型被迫更重视这些类。注意权重要按类别索引顺序放对位置不然等于把错误的加权逻辑加到了错误的类上效果还不如不加权。另外类别不均衡情况下验证集划分也要改成分层抽样确保每类在验证集里的比例和全集一致用PyTorch的SubsetRandomSampler按标签分层取索引而不是简单random_split。5.3 已标注不等于标注无误如何清洗标签噪声现象训练过程中某一张图反复让loss产生一个尖峰或者验证集里出现一张标签和内容明显不符的图片。原因数据集的“已标注”是人工标注的人眼在快速标注几百张水果图时看走眼很正常尤其是青苹果和梨这种长相接近的类别标错一两张一点都不稀奇。解决把训练集里被模型高置信度预测为错误类别的样本单独拎出来看一眼。这里有个实用的循环清洗脚本用当前训练的模型对训练集做一次预测筛选出预测类别与原标签不一致、且置信度超过0.9的样本打印文件名和预测类别。这些样本大概率是标注错误而非模型误判因为真实困难的样本置信度通常不会这么高。确认后要么删除要么手动改标签。一个小数据集里清洗掉三到五张错标图对验证集准确率的提升可能比调一整天参数都明显。5.4 图片尺寸不一致同一批图里混着320×240和4032×3024现象DataLoader在加载某些图片时抛异常报错信息是“Expected 3D tensor, got 2D或者tile cannot extend outside image”。原因数据集里部分图片是手机拍的竖图部分是从网页爬下来的缩略图宽高比完全不同有的还是灰度图。虽然ImageFolder能读进来但RandomResizedCrop在裁剪时会因为尺寸太小抛异常。解决在数据预处理阶段统一做Resize到256×256然后再做CenterCrop到224。更稳妥的做法是写一个预检查脚本把所有尺寸异常或通道异常的图片路径列出来手动决定删除还是转换。重点提醒一下如果你用torchvision的transforms直接链式处理即使某张图只有32×32RandomResizedCrop也会把它强行放大到224图片模糊不清模型等于在学会“模糊的图片是苹果”这种虚假相关性。5.5 模型学的是背景果实小、叶子多、天空亮现象训练和验证准确率都很高但把模型放到一片新果园拍的照片上测试准确率立刻崩塌。原因数据集里的图片拍摄风格太一致——苹果总是带叶子出现葡萄总是挂在藤蔓上模型很可能学到的是“叶子纹理红色区域苹果”而不是苹果本身。解决训练时用RandomResizedCrop把scale下限调低让每次裁剪只露出果实局部同时增加随机擦除或者随机遮挡强迫模型利用多个特征做判断。更极端一点可以用简单的颜色分割预处理把背景压暗但那样会破坏预训练模型的输入分布需要重新归一化不太推荐。你真正要做的是在验证之外再留一组“跨场景测试图”比如从网上找几张不同拍摄环境的水果照片测试模型泛化能力不要被同分布验证集的漂亮数字骗过去。6. 把7类水果分类精度从70%提到90%最后三个操作6.1 先用学习率扫描找到甜点区间不要拍脑袋设学习率。先用一个3到5个epoch的扫描脚本从1e-2到1e-5按十倍步长试一圈记录每个学习率下验证集的最高准确率。这一步花不了多长时间却能避免你在错误的学习率上跑十几个epoch。通常你会看到一个明显的学习率甜点区间比如3e-4到1e-4之间验证集准确率最高低于1e-5则几乎不收敛高于1e-3则训练震荡。把最终训练固定在这个区间内。6.2 余弦退火配合五个epoch的预热在确定学习率区间后我用CosineAnnealingLR做调度前两个epoch用线性预热从0升到目标学习率后二十个epoch按余弦曲线缓慢衰减到接近0。这个做法的好处是让模型在训练后期用极小的步长精细调整边界对小样本数据集尤其友好。相比固定学习率和StepLR余弦退火在最终验证集准确率上经常能再涨一到两个百分点这个提升不是玄学是因为后期小步长让模型避开了损失曲面里的尖锐局部最优。6.3 用置信度阈值做最后一道闸训练结束后还有一个容易忽略的部署技巧不要直接取softmax最大值作为最终输出给分类结果加一个置信度阈值。比如设定0.75低于阈值的图片全部判为“不确定”而不是硬归到某一类水果。果园实际场景里经常拍到树叶遮挡一半的果实、远处模糊的未成熟果实这些样本硬分类必然出错输出“不确定”可以让后续的质检逻辑把人叫回来复核而不是让模型擅自做决定。这个阈值的取值可以从验证集上画置信度分布曲线来确定选一个能覆盖90%正确预测的最低点。这个小改动在真实场景里带给你的体验改善往往比模型结构升级还明显。这一套流程是我在这些小样本分类数据上反复跑出来的习惯从目录检查、迁移学习、分层抽样到清洗标签每一步都踩过坑。尤其是那条“已标注”数据不能全信的经验帮我避开了不少无效实验。希望这些操作细节能让你少走一圈弯路也希望帮到你。本文还有配套的精品资源点击获取
返回列表