7类果树水果图像分类实战:ResNet-18微调与数据集详解
简介这是一套面向图像分类入门与算法验证的7类果树水果数据集包含草莓、甜瓜、橙子、苹果等常见品种约500张已标注图片数据经过预处理可直接作为分类网络输入。适合计算机视觉初学者或研究者快速开展分类实验也可用于模型改进与对比测试。压缩包内共587个文件以584张jpg图像为主体附带1个json文件用于类别标签映射、1个Python脚本用于数据可视化图片预览png可供快速查看包体约23.19MB。资源已划分好训练集与测试集同类图片按目录存放配合脚本即可完成数据分布检查与样本可视化。目前已有218人学习下载对于需要现成小型数据集验证网络结构或撰写实验报告的读者可省去采集与整理标注的时间直接聚焦于模型设计与调参环节。1. 把 7 类果树水果做成图像分类这份已标注数据集到底能干什么做图像分类最磨人的不是搭模型而是数据。很多人一上来就拿 ImageNet 预训练权重去微调结果发现自己的业务场景里根本没那么多带标注的图。这份「7 种长在果树上的水果」数据集规模不大约 500 张但它是被预处理过的能直接喂给分类网络。这意味着你不需要自己爬图、去重、裁边、写标注省的是从 0 到 1 最脏的那段活。它的分类个数是 7涵盖草莓、甜瓜、橙子、苹果等常见果树水果。数据已经划好了训练集和测试集同类图片放在同一目录下还带了一个 show 脚本跑一下就能可视化。适合三类人刚入门想跑通一个完整分类流程的学生、需要一份干净数据验证网络改动的在校研究者、以及做农业或生鲜识别想先有个 baseline 的算法工程师。对老手来说这份数据最大的价值在于可以用来快速对比不同网络的收敛行为和 trick 的收益而不是多大、多全。说白了它是一门手艺课的操练素材。下面我按实际拆包的顺序把这套数据从文件结构、训练配置到踩坑记录完整过一遍。2. 数据集的目录结构训练集、测试集与类别映射是怎么组织的2.1 解包后先看什么类别目录与 JSON 映射文件拿到资源后第一步不是急着开训练脚本而是先把目录树列出来。常见做法是解压后在同一级目录下执行 tree 命令或是用 find 命令看一眼层级。这份数据集的组织方式是「按类别分目录」训练集和测试集各自独立各自存放同一类别的图片。也就是说你会在 train 目录下看到 7 个子目录每个目录名对应一个类别test 目录同理。这种组织方式叫「按文件夹分类」PyTorch 的torchvision.datasets.ImageFolder和 Keras 的image_dataset_from_directory都能直接读省掉自己写标注解析的麻烦。如果你看到类别目录名是 0、1、2 这种数字编号别急着改目录先打开配套的 JSON 文件。这份资源里类别名和编号的映射在 JSON 文件里写死了里面会明确写出0: apple、1: strawberry这样的对应关系。# 解压后先看整体结构不要直接开训练 $ find . -maxdepth 2 -type d | sort ./dataset/ ./dataset/test/ ./dataset/train/ ./dataset/train/apple/ ./dataset/train/strawberry/ ./dataset/train/melon/ ./dataset/train/orange/ ./dataset/train/peach/ ./dataset/train/grape/ ./dataset/train/pear/这段命令只列两级目录目的是确认类别数量和目录命名方式。如果类别不是 7 个或者目录里有隐藏文件后面训练会出各种奇怪的错误。接着打开 json 文件确认映射关系{ 0: apple, 1: strawberry, 2: melon, 3: orange, 4: peach, 5: grape, 6: pear }这个映射文件是整个数据集的「字典」。之后训练脚本里如果用了class_to_idx它读到的顺序就从这个 json 来。建议训练前先打印一次这个映射确认与你理解的类别顺序一致否则预测结果张量的第 3 个位置对应哪一类你会毫无头绪。2.2 单类样本数量统计约 500 张是怎么分配到 7 类里的先做一次数量统计。我一般会跑一个小脚本统计每个类别下 train 和 test 的图像数量。这一步能提前暴露类别不均衡问题避免训练完才发现某个类只有 20 张训练图。import os from collections import Counter train_root dataset/train test_root dataset/test train_counter Counter() test_counter Counter() for cls_name in os.listdir(train_root): cls_path os.path.join(train_root, cls_name) if os.path.isdir(cls_path): train_counter[cls_name] len([ f for f in os.listdir(cls_path) if f.lower().endswith((.jpg, .jpeg, .png)) ]) for cls_name in os.listdir(test_root): cls_path os.path.join(test_root, cls_name) if os.path.isdir(cls_path): test_counter[cls_name] len([ f for f in os.listdir(cls_path) if f.lower().endswith((.jpg, .jpeg, .png)) ]) print(Train:, dict(train_counter)) print(Test:, dict(test_counter)) print(Total train:, sum(train_counter.values())) print(Total test:, sum(test_counter.values()))这段脚本只做一件事把每个类别的样本数数清楚。注意我用了endswith做扩展名过滤防止目录里混入Thumbs.db这类系统文件被误统计。.jpg、.jpeg、.png三种常见格式都纳入了。跑完后你会发现总量在 500 张左右每类平均 70 张上下其中测试集会少一些这就是「约 500 张数据」的实际含义。从统计结果能判断这个数据集的上限每类 70 张左右的训练图微调一个 ResNet-18 是够的但想从头训练一个 ViT 就不现实了。这个判断后面会反复用到——很多训练翻车不是因为代码写得差而是数据量和模型容量不匹配。2.3 show 脚本的用途可视化样本而不是只信文件名资源里附带的 show 脚本核心作用是把图片和对应的标签一起展示出来。直接贴一段简化的可视化逻辑import matplotlib.pyplot as plt from torchvision.datasets import ImageFolder dataset ImageFolder(dataset/train) class_names dataset.classes fig, axes plt.subplots(2, 4, figsize(12, 6)) for i, ax in enumerate(axes.flat): if i len(dataset): break img, label dataset[i] ax.imshow(img.permute(1, 2, 0)) ax.set_title(class_names[label]) ax.axis(off) plt.tight_layout() plt.savefig(sample_check.png, dpi150)这个脚本用ImageFolder直接加载训练集目录dataset[i]会返回第 i 张图像张量和它对应的标签索引。class_names[label]把索引翻译成类别名标题栏直接显示。把图存成sample_check.png的原因是在远程服务器上跑plt.show()会直接报错或黑屏保存成文件再查看是最稳妥的。我之所以建议跑这个脚本是因为文件名的迷惑性远比你想象的大。比如n12761284_14775.jpg这种文件名是 ImageNet 的 WordNet ID 风格不打开看根本不知道里面是什么水果。有些类之间长得极其接近像甜瓜和梨单看文件名你完全没法判断数据有没有放错目录。可视化这一步不花多少时间但能避免训练一个标签错乱的模型——那才是真正的灾难。3. 训练一个分类器从 ResNet-18 微调到训练参数的选择3.1 为什么选择 ResNet-18 而不是更大的网络数据量只有 500 张每类几十张这个尺度下盲目上 ResNet-50 或 EfficientNet-B4 不会带来增益反而容易过拟合。ResNet-18 是 18 层残差网络参数约 1100 万在 ImageNet 上预训练过。用它做迁移学习把最后一层全连接改成 7 类输出的结构是这份数据最合理的起点。残差结构在这里的关键作用是缓解梯度消失——小数据集上你往往只能训练少量 epoch如果网络太深前面的层根本学不动那和随机初始化没什么区别。ResNet-18 因为短接路径的存在即使只用 20 个 epoch深层部分也能拿到有效的梯度信号。相比之下VGG-16 没有残差结构在这个数据量下表现明显差一截。另外推理成本也是现实考量。树莓派或 Jetson Nano 这类边缘设备上跑 ResNet-18单张图推理能做到几十毫秒换 ResNet-50 就得翻倍。如果你做的只是一个原型 demoResNet-18 的精度和速度平衡是最舒服的。3.2 完整的训练脚本关键参数与数据增强配置直接给一套能跑的微调脚本基于 PyTorch。假设数据目录结构如上面所述训练前先定义 transforms注意这步直接决定模型最终精度import torch from torch import nn from torch.utils.data import DataLoader from torchvision import datasets, transforms, models train_transform transforms.Compose([ transforms.Resize((224, 224)), transforms.RandomHorizontalFlip(p0.5), transforms.RandomRotation(15), # 轻度旋转模拟采摘角度差异 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_transform transforms.Compose([ transforms.Resize((224, 224)), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]), ]) train_data datasets.ImageFolder(dataset/train, transformtrain_transform) val_data datasets.ImageFolder(dataset/test, transformval_transform) train_loader DataLoader(train_data, batch_size32, shuffleTrue, num_workers4, drop_lastTrue) val_loader DataLoader(val_data, batch_size32, shuffleFalse, num_workers4) model models.resnet18(weightsmodels.ResNet18_Weights.IMAGENET1K_V1) model.fc nn.Linear(model.fc.in_features, 7) device cuda if torch.cuda.is_available() else cpu model model.to(device)这段代码里有几个参数值得展开说。weights参数显式指定了预训练权重版本不传的话 PyTorch 会给出警告而且以后版本可能默认随机初始化别踩这个坑。model.fc原来的输出是 1000 类改成 7 类和这份数据集的类别数对齐。如果 json 里写着 7 类这里就必须是 7多一个少一个都会在训练时报错或静默错位。drop_lastTrue是因为训练集总量不是 32 的整数倍最后一批如果只有十来张BatchNorm 的统计量会异常尤其在迁移学习时容易引起精度抖动。num_workers4在本地机器的 Windows 上如果报错改成 0 最省事。训练循环这里不贴了那是标准的交叉熵加 SGD。优化器我会用带动量的 SGD初始学习率 0.001weight decay 5e-4。对这份数据SGD 比 Adam 的最终精度通常高 1 到 2 个点因为小数据集上 Adam 容易收敛到尖锐极小值泛化差一些。3.3 训练集与测试集划分的意义验证你「真的」学到了特征这个资源把训练集、测试集分开存放在不同目录意图很明确告诉你在固定划分下做实验。你不需要自己搞train_test_split但必须清楚跨类别的图片可能有同源关系——比如同一棵树的照片被拆进了两边的集合那测试集就过于乐观了。但因为这里每类只有几十张且来源分散同源污染的概率不大。这就是固定划分的意义所有用户在同一基准上对比结果而不是各自随机切分导致指标不可比。训练完成后用测试集评估 Top-1 Accuracy。以这份数据的大小和经验微调 ResNet-18 应当能达到 80% 到 90% 之间。如果低于 75%大概率是类别界定不清或训练流程有问题而不是数据本身不够这个精度。把测试集看成「验收标准」不要拿训练集精度当宣传口号。训练精度 99% 只能说明模型记住数据测试集才对泛化能力说了算。4. 推理与可视化把模型输出翻译成「这是苹果」的过程4.1 加载权重、预处理单张图片、输出类别置信度模型训练完了你需要一个能对单张图做预测的脚本。这里最容易出错的地方是推理时的预处理必须和验证时完全一致否则图片像素分布不一样预测结果会凭空打折扣。以下脚本把权重载入、类别映射、推理三个阶段合并在一起import json import torch from PIL import Image from torchvision import transforms, models with open(dataset/class_mapping.json, r) as f: idx_to_name {int(k): v for k, v in json.load(f).items()} model models.resnet18(weightsNone) model.fc torch.nn.Linear(model.fc.in_features, 7) model.load_state_dict(torch.load(best_model.pth, map_locationcpu)) model.eval() infer_transform transforms.Compose([ transforms.Resize((224, 224)), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]), ]) img Image.open(single_test.jpg).convert(RGB) x infer_transform(img).unsqueeze(0) with torch.no_grad(): logits model(x) prob torch.softmax(logits, dim1) conf, idx torch.max(prob, dim1) print(fPredicted: {idx_to_name[idx.item()]}, Confidence: {conf.item():.4f})这段代码的关键点有三个。一是weightsNone意为不从网络下载预训练权重直接用刚训练好的模型文件。习惯了写resnet18(pretrainedTrue)的人在这里容易翻车版本不匹配会直接报错。二是model.eval()切换 BatchNorm 和 Dropout 到推理模式。BatchNorm 在训练时用批内统计量推理时用全局统计量不切的话结果忽上忽下。三是convert(RGB)防止读入带 Alpha 通道的 PNG 图。这份数据里大多是 JPG但保不齐你在验证时会拿到一张带透明通道的图不转 RGB 会导致通道数和预训练输入不匹配。4.2 类别间相似性对置信度的影响输出 0.45 该怎么解读最终输出的置信度不一定都是 0.9 上下。对这份水果数据甜瓜和梨、苹果和橙子这两对类别的视觉特征本身就有重叠真实场景中模型给出 0.45 的置信度是完全可能的。这个时候不能只取最大概率做决定要看第二、第三位的概率分布。比如「甜瓜 0.45、梨 0.32、苹果 0.18」说明模型在甜瓜和梨之间犹豫这在生产环境里应该被设计成交互确认而不是自动判定。另外如果一张朝向、光线极端的测试图落到了低置信区间不要急着加数据先想想这个数据集的边界在哪。约 500 张的规模意味着它覆盖的是「果树水果在常见角度和光照下的形态」不是高温强光、重度遮挡、腐烂病变的环境。认清边界模型在边界内用边界外交给人工这是工程落地和课堂作业最大的差别。5. 避坑与常见问题训练翻车最多的四个真实场景5.1 错误信息Expected 3D or 4D tensor for input或者类别数不匹配现象训练脚本刚跑起来就报Expected 3D or 4D tensor for input或者shape不匹配的错。原因最常见的是Resize和RandomResizedCrop混用。Resize((224, 224))是把图直接拉伸RandomResizedCrop(224)是先随机裁剪再缩放两者的输出尺寸固然都是 224但如果你的输入图是灰度单通道经过ToTensor()得到的是[1, 224, 224]而预训练模型期待[3, 224, 224]。另外model.fc输出的 7 和CrossEntropyLoss的标签最大值不对齐也会报错。解决读图后强制convert(RGB)保证三通道。类别数先用脚本打印len(train_data.classes)确认是 7再改model.fc的out_features。如果换过网络结构别忘了把num_classes参数一并调掉。这是能靠日志排查的通病报错信息里shape后面跟着的数字会明确告诉你哪个维度对不上。5.2 精度上不去训练集 98% 而测试集只有 60%现象训练到最后train accuracy 接近 100% 但测试集始终在 60% 多或者完成 20 个和 50 个 epoch 的精度几乎一样。原因类别不平衡。这份数据「约 500 张」7 类均分和 6 类各 40 张、1 类 150 张是完全不同的情况。如果少数类样本只有 20 张训练图模型对它的泛化会非常差。另一个常见原因是标签噪声——前面提到过某些类之间长得像数据整理时极易放错目录。解决先跑统计脚本看每类数量。如果个别类确实偏少可以对该类多做数据增强不要只依赖水平翻转。再跑一次 show 脚本人工抽查每类样本把明显错放的图剔掉或挪回正确目录。最后如果类别严重不平衡用WeightedRandomSampler做采样加权from torch.utils.data.sampler import WeightedRandomSampler targets train_data.targets class_counts torch.bincount(torch.tensor(targets)) weights 1.0 / class_counts.float() sample_weights weights[targets] sampler WeightedRandomSampler(sample_weights, num_sampleslen(sample_weights)) train_loader DataLoader(train_data, batch_size32, samplersampler)WeightedRandomSampler会让每个 epoch 中少数类被抽到的概率按权重提升。注意sample_weights必须是一个和数据集长度相同的张量每个位置对应一张图。这是处理小数据集不平衡最直接的解法比手工复制图片省事得多。5.3 迁移学习反而比随机初始化差小数据集上的「玄学」问题现象用预训练模型微调后验证精度反而比用一穷二白的模型从头训练还要低或者 loss 在验证集上持续震荡。原因这大概率是学习率太大把预训练权重在前期几个 epoch 里冲掉了。预训练权重是经过大量数据打磨的滤波器微调的核心是只让最后的分类头学习前面的卷积层不做剧烈更新。学习率 0.001 给了全网络卷积层的预训练特征就被破坏了。另一种可能数据增强太弱模型在 500 张图上把背景也记住了。解决将主干部分的学习率调为分类头的十分之一。常见做法是给model.parameters()和model.fc.parameters()分别设置学习率backbone_lr 1e-4 head_lr 1e-3 params [ {params: [p for n, p in model.named_parameters() if fc not in n], lr: backbone_lr}, {params: model.fc.parameters(), lr: head_lr}, ] optimizer torch.optim.SGD(params, momentum0.9, weight_decay5e-4)这样主干只做细微调节分类头从零学起。如果换用 Adam把两个 lr 同步降五倍。这条经验我在多个小数据集上反复验证过是微调参数里最值得优先调的旋钮。5.4 数据集仅 500 张增强过头导致学到的全是噪声现象加了大量数据增强后训练 loss 下降变慢甚至测试精度反而更低了。原因数据增强对小数据集的收益不是单调的。RandomRotation(60)、RandomResizedCrop(scale(0.4, 1.0))会把苹果裁成只剩一个角落模型看到的训练样本和真实分布差距越来越大相当于在学一堆失真图。增强的本质是让模型看到「同一个物体的合理变化」不是制造 wéi 反直觉的样本。解决强度保守为先去跑 20 轮精度够就停。翻转概率 0.3 到 0.5、旋转 10 到 15 度、亮度对比度扰动 0.1 到 0.2是这份水果数据的安全区间。注意每一类别的样本量只有五六十张时「增强即生成」的思路会非常容易把模型往伪影上带。这时候最值得做的是把所有 per-class 的准确率打出来谁低补谁的真实样本而不是整体堆增强强度。6. 把这个数据集的剩余价值榨干类别均衡与混淆矩阵分析处理完训练与基本推理后这份数据还能再做一步「深度学习标准动作」——把混淆矩阵和逐类准确率打印出来。这一步不必写训练代码只写评估代码就够了。它让你看清模型到底在哪两类之间犯错从而决定下一个版本是要收集更多数据还是增加新的特征分支。混淆矩阵的实现如下先取全部测试集预测值再交给 sklearnimport numpy as np import torch from sklearn.metrics import confusion_matrix, classification_report from torch.utils.data import DataLoader model.eval() all_preds, all_labels [], [] with torch.no_grad(): for imgs, labels in val_loader: imgs imgs.to(device) logits model(imgs) preds torch.argmax(logits, dim1).cpu().numpy() all_preds.extend(preds) all_labels.extend(labels.numpy()) cm confusion_matrix(all_labels, all_preds) print(Confusion Matrix:\n, cm) print(classification_report(all_labels, all_preds, target_namestrain_data.classes))confusion_matrix的行是真实类别列是预测类别。对角线的数值是正确分类数。如果一个非对角元特别大比如甜瓜那行梨那列有 5说明甜瓜样本被模型判成梨的数量最多。这个信息直接指导后续策略与其盲目加数据不如针对「甜瓜 vs 梨」这两类的边界样本多采一些或者干脆为这两类单独设计后处理规则。注意preds和labels都是整数索引打印时用target_names做映射。另外我习惯把 top-2 的概率也打印出来因为甜瓜与梨的混淆往往在 top-2 里能看到而 top-1 天生把真实类别排在了后面。这一步做完这个 500 张数据集的潜力才算真正被摸透。从那以后我每接手一个分类数据都强制走一遍数量统计、可视化抽查、混淆矩阵这三板斧把「猜」换成「看数据」。一份数据放久了会「馊」不是数据坏了而是你对它的理解没跟上。希望帮到你。本文还有配套的精品资源点击获取