ARTICLE DETAIL

资讯详情

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

PyTorch图像识别实战:11种水果分类数据集训练与部署指南

PyTorch图像识别实战:11种水果分类数据集训练与部署指南 简介这是一份面向深度学习入门与图像分类实战的11种水果分类数据集涵盖苹果、鳄梨、蓝莓、辣椒、樱桃、猕猴桃、芒果、橙子、岩瓜、草莓、小麦共11个类别。数据已按类别分文件夹存放训练集含2562张图片测试集含636张图片无需额外标注即可直接用于卷积神经网络等模型的训练与评估。压缩包共2000个文件以jpeg图片为主同时包含png、webp及少量bmp格式方便不同场景下的图像加载另附classes.json类别字典和可视化脚本py文件便于快速查看样本分布与验证分类效果。压缩包整体约855MB目录层级清晰适合正在学习深度学习图像分类、需要现成数据完成课程设计或模型验证的开发者使用。目前已有1667人学习下载是一份结构规范、上手门槛低的水果图像分类训练资源。1. 为什么 11 种水果分类是深度学习图像识别最合适的练手题「深度学习图像识别数据集11种水果分类数据集」这类资源在行业内流传很广常见形态是一个包含 11 个类别子文件夹的图片集苹果、香蕉、橙子、柠檬、猕猴桃、葡萄、桃、梨、西瓜、石榴、草莓各占一个目录总量在几百到一万多张之间。它解决的问题很直接让一个刚接触视觉任务的人用最短路径跑通「数据加载 → 模型训练 → 评估 → 推理」全流程而且 11 分类的难度恰好卡在「不能靠猜」和「不至于练不动」之间。适合三类人准备做课程设计或毕设的学生、想验证自己 PyTorch 环境是否配对的初学者、需要快速建立图像分类 baseline 的工程师。别小看这个小数据集它能把深度学习里最核心的过拟合、类别不均衡、数据增强、迁移学习这几个问题全部暴露一遍。2. 打开数据集目录结构、加载代码与张量形状的全套读法拿到「11种水果分类数据集」之后第一件事不是立刻训练而是先把这个数据集「读对」。分类任务的数据集组织形式直接决定后续代码怎么写也决定你踩不踩 label 错位的坑。2.1 目录结构为什么 train / val / test 三层文件夹就是标准答案绝大多数水果分类数据集的压缩包解压之后内部是这样一个树形结构fruits-11/ ├── train/ │ ├── apple/ │ │ ├── apple_001.jpg │ │ ├── apple_002.jpg │ │ └── ... │ ├── banana/ │ ├── orange/ │ ├── lemon/ │ ├── kiwi/ │ ├── grape/ │ ├── peach/ │ ├── pear/ │ ├── watermelon/ │ ├── pomegranate/ │ └── strawberry/ ├── val/ │ └── 同样按类别分子文件夹 └── test/ └── 有的版本没有 test只有 val这个结构之所以是标准答案是因为 PyTorch 的torchvision.datasets.ImageFolder就是为它设计的。它按子目录名自动生成类别索引文件夹名即标签不需要额外的 CSV 标注文件。我一般会先跑一遍find命令确认每个类别的图片数量防止某些类别图片被漏拷。find fruits-11/train -type d -exec sh -c echo $1: $(ls $1 | wc -l) _ {} \; | sort上面这条命令会列出 train 下每个子文件夹的图片数量。正常情况下 11 个类别数量应大致均衡如果发现某个类别只有几十张后面就要做类别加权或者放弃该类别这是后话。数量分布直接决定你后面用不用WeightedRandomSampler这一步花 30 秒值得。2.2 用 PyTorch 的 ImageFolder 一次性读入全部数据确认目录没问题之后加载数据只需一段十几行的代码。这里我直接给出一份「跑通版」的 PyTorch 数据加载片段适用于 torchvision 0.13 及以上版本。import torch from torchvision import datasets, transforms # 训练集与验证集使用不同的预处理策略 train_transform transforms.Compose([ transforms.Resize((256, 256)), # 先缩放到统一尺寸 transforms.RandomResizedCrop(224), # 随机裁剪数据增强 transforms.RandomHorizontalFlip(p0.5), # 随机水平翻转 transforms.ToTensor(), # PIL Image - Tensor像素值归一化到 [0,1] transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]) ]) val_transform transforms.Compose([ transforms.Resize((256, 256)), transforms.CenterCrop(224), # 验证集不做随机增强只中心裁剪 transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]) ]) train_data datasets.ImageFolder(fruits-11/train, transformtrain_transform) val_data datasets.ImageFolder(fruits-11/val, transformval_transform) print(train_data.classes) # 按字母序排列的 11 个类名 print(train_data.class_to_idx) # {apple: 0, banana: 1, ...} print(len(train_data), len(val_data))这段代码的逻辑是ImageFolder在初始化时会遍历fruits-11/train下的所有子文件夹按字母序生成class_to_idx映射并把每张图片的路径和对应整数标签存进内部列表。训练集和验证集必须用同一个数据源目录结构否则类别索引会错位。Normalize里的均值方差用的是 ImageNet 统计值因为我们后面要做迁移学习预训练模型就是在这些数值上训练的输入分布不一致会让微调效果打折扣。2.3 输入张量到底是什么从 JPEG 像素到归一化数组很多新手在这里犯迷糊图片明明是.jpg文件怎么ToTensor()之后就变成了三维数组一张 224×224 的 RGB 彩色图片经过ToTensor()后得到的 Tensor 形状是(3, 224, 224)三个维度分别是通道数、高度、宽度数值范围从 0255 缩放到 01。再经过Normalize每个通道减去均值除以标准差数值变成近似标准正态分布这是为了让模型训练更稳定。你可以在加载之后单独检查一下数据形状sample, label train_data[0] print(sample.shape) # torch.Size([3, 224, 224]) print(label) # 整数比如 7这里的7不是类别名字而是class_to_idx里对应的整数。后面模型输出的也是 11 个概率值取 argmax 后得到整数索引再通过train_data.classes[idx]反查回水果名字。这个「整数索引 ↔ 文件夹名」的双向映射是整个训练过程中最容易错位的环节建议在加载之后立即打印一遍确认。3. 训练一个 11 分类水果识别模型迁移学习与三个关键超参数数据集读进来了接下来进入核心环节训练模型。11 分类水果识别的训练本质上是一个标准的图像分类问题选对模型和超参数每个人都能稳定跑到 94% 以上的验证准确率。3.1 选型理由为什么这个量级的数据不配用大模型常见做法是用 ImageNet 预训练权重做迁移学习。水果图片和 ImageNet 里已有的苹果、香蕉、草莓等类别高度相关预训练模型已经学会了纹理、边缘、颜色这些底层特征我们只需要替换最后一层分类头让它输出 11 个类别。这是个「站在巨人肩膀上」的策略比从零开始训练省下大量时间准确率还高得多。基础版本我推荐 ResNet18而不是 ResNet50 或 Vision Transformer。原因很现实水果数据集通常只有几千到一万张图ResNet18 参数量约 1100 万足够拟合这个规模的数据ResNet50 参数量翻了几倍在小数据上更容易过拟合训练时间也成倍增加。如果你打算部署到边缘设备或手机端那mobilenet_v3_small更合适精度比 ResNet18 低 12 个点但体积小一个数量级。教学演示和 baseline 场景老老实实用 ResNet18 就行。3.2 预处理与数据增强让模型看见更多「角度的苹果」数据增强是整个训练里性价比最高的一环。水果拍摄角度、光照、遮挡情况各异想让模型在验证集和真实场景里都稳定工作就得在训练时人为制造「多样性」。前面 2.2 节代码里已经写进了三个增强策略随机裁剪、随机翻转、随机旋转。我再补一个完整的版本包含颜色抖动import torchvision.transforms as T train_transform T.Compose([ T.Resize(256), T.RandomResizedCrop(224, scale(0.6, 1.0)), T.RandomHorizontalFlip(), T.RandomRotation(15), T.ColorJitter(brightness0.2, contrast0.2, saturation0.2), T.ToTensor(), T.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]) ])注意RandomResizedCrop的scale(0.6, 1.0)意思是每次随机裁剪原图 60%100% 的区域再缩放到 224×224。这个参数决定了模型能看到「多局部的物体」scale 下限设得太低比如 0.08模型会经常只看到水果的半个切面训练难度骤增设得太高又起不到增强作用。水果数据集里 0.6 起步比较稳。RandomRotation(15)控制在 15 度以内超过 30 度会让真实世界里的「水果朝上摆放」这个语义失真。颜色抖动幅度也别太大亮度对比度各 0.2 足够因为在真实场景里过强的颜色扰动会让红苹果和红柿子难以区分。3.3 训练循环损失、优化器与学习率的三个关键旋钮训练代码的核心只有三个旋钮损失函数、优化器、学习率外加一个 batch size。分类任务损失函数固定用交叉熵优化器选 SGD 带动量这是迁移学习里最稳的组合。学习率是重中之重微调预训练模型时新加的分类头需要较大的学习率而前面几层卷积已经学会通用特征学习率过大会把预训练权重冲坏。常见做法是给全连接层设 0.001特征提取层设 0.0001。import torch.nn as nn from torchvision import models # 加载 ImageNet 预训练权重torchvision 0.13 推荐写法 model models.resnet18(weightsmodels.ResNet18_Weights.IMAGENET1K_V1) # 替换最后一层全连接输出 11 类 num_features model.fc.in_features # 512 model.fc nn.Linear(num_features, 11) # 全连接层用较大学习率特征层用较小学习率 fc_params [p for name, p in model.named_parameters() if fc in name] base_params [p for name, p in model.named_parameters() if fc not in name] optimizer torch.optim.SGD([ {params: base_params, lr: 1e-4}, {params: fc_params, lr: 1e-3} ], momentum0.9, weight_decay5e-4) criterion nn.CrossEntropyLoss()训练循环本身不复杂每轮迭代做四件事取一个 batch 的数据和标签、前向传播计算损失、反向传播计算梯度、优化器更新参数。如果数据集总量在 8000 张左右、batch size 为 64一个 epoch 大约 125 个 stepResNet18 在消费级显卡上跑一个 epoch 只要 3060 秒。通常训练 1525 个 epoch 就能收敛到 92%96% 的验证准确率。设置一个学习率衰减策略lr_scheduler.StepLR(optimizer, step_size7, gamma0.1)每 7 个 epoch 学习率降为原来的十分之一让收敛更平滑。整个训练过程的核心观察指标只有一个验证集准确率。每跑完一个 epoch 打印一次 val acc如果连续 5 个 epoch 不再上升就提前终止训练这就是早停策略。4. 验证做扎实混淆矩阵、分类报告与每类召回率的排查价值训练完模型之后准确率数字只说明「整体还行」。11 分类任务里真正能暴露问题的是混淆矩阵和每一类的精确率、召回率。一个 98% 整体准确率的模型可能对石榴这类深色水果的召回率只有 60%这种偏科在测试阶段必须揪出来。4.1 混淆矩阵才是分类任务的照妖镜写一个评估函数跑完整个验证集后生成混淆矩阵import numpy as np import torch from sklearn.metrics import confusion_matrix, classification_report def evaluate(model, dataloader, device): model.eval() all_preds [] all_labels [] with torch.no_grad(): for images, labels in dataloader: images, labels images.to(device), labels.to(device) outputs model(images) _, preds torch.max(outputs, 1) # 取概率最大的类别索引 all_preds.extend(preds.cpu().numpy()) all_labels.extend(labels.cpu().numpy()) cm confusion_matrix(all_labels, all_preds) print(classification_report(all_labels, all_preds, target_namestrain_data.classes)) return cm, np.array(all_labels), np.array(all_preds)这段代码的逻辑是把验证集所有图片过一遍模型收集预测结果和真实标签然后用 scikit-learn 的confusion_matrix和classification_report输出详细指标。运行后会看到一张 11×11 的矩阵和一份包含每类 precision、recall、f1-score 的报告。我重点看两处一是对角线上的数字是否明显大于非对角线二是 recall 最低的那一两类是什么。4.2 从 PyTorch 模型到 ONNX推理脚本怎么写验证通过后模型还是要落地的。新手常犯的错误是训练完直接拿着.pth文件到处跑但.pth只存了参数没有模型结构定义。一个更实用的做法是把模型导出成 ONNX 格式这样部署时不需要重新定义模型类也方便后续用 ONNX Runtime 做推理加速。dummy_input torch.randn(1, 3, 224, 224) torch.onnx.export( model, dummy_input, fruits11_model.onnx, input_names[input], output_names[output], opset_version11, dynamic_axes{input: {0: batch_size}, output: {0: batch_size}} )导出前必须把模型切到model.eval()模式否则 BatchNorm 层的统计值会跟着改变。dynamic_axes参数设置了动态 batch这样导出后的模型既支持单张图片推理也支持批量预测。导出完成后用一行代码验证 ONNX 输出和 PyTorch 原模型是否一致import onnxruntime as ort sess ort.InferenceSession(fruits11_model.onnx) onnx_out sess.run(None, {input: dummy_input.numpy()})两者的输出差异如果小于 1e-4说明模型转换没有引入数值偏差。这个数值一致性检查是 ONNX 部署步骤里最容易省掉的省掉之后大概率在某个环境里遇到玄学 bug。5. 避坑指南水果数据集训练里最常见的 5 个翻车现场这个数据集看起来简单实际训练起来坑并不少。下面五条全是真实训练中会反复遇到的现象每条都按「现象 → 原因 → 解决」的顺序拆开讲。5.1 训练集准确率 99%验证集只有 71%这是最典型的过拟合信号。11 分类水果数据集总量不大而 ResNet18 有 1100 万参数模型完全有能力把训练集里的每一张图都背下来。对比一下两个准确率的差距如果训练集 99% 而验证集 71%差距超过 20 个点基本可以判定过拟合了而不是验证集分布问题。解决方案按性价比排序第一加大数据增强强度把RandomRotation从 15 度提到 20 度把ColorJitter的饱和度扰动调到 0.3让模型看到更多变体第二把weight_decay从 5e-4 提到 1e-3加强对大权重的惩罚第三在验证集准确率连续 3 个 epoch 不上升时触发早停。做了这三步之后就算个别类别仍然过拟合整体差距也会缩到 5 个点以内。5.2 Label 与路径映射错位ImageFolder 的类名顺序按字母排有一次训练出来的模型把苹果全部识别成香蕉准确率只有 9%看混淆矩阵发现对角线全偏了一位。排查后发现问题出在类别索引ImageFolder生成的class_to_idx是按文件夹名字母序排列的我的数据集里apple排第 0 位但之前某份 CSV 格式的数据是人工打标的类别编号跟字母序对不上。两边一旦混用标签就整体错位。解决方法是每次加载数据后强制打印print(train_data.class_to_idx)并且用同一个ImageFolder实例去获取类名映射绝不用手写的硬编码列表。训练前花 10 秒看一眼这行输出能省掉一下午的排查。5.3 单张图推理速度 200ms预处理反而比模型更慢训练完模型丢进推理脚本发现单张图片要 200ms 才能出结果其中模型推理只占 3ms。问题出在我的预处理里调用了 PIL 的resize到 256 再center_crop到 224而resize传入的是普通 Python 函数每张图都走一遍 PIL 全量插值。图片分辨率越大这个瓶颈越明显。解决思路是推理阶段用固定尺寸输入把Resize(256)换成Resize((224, 224))省掉CenterCrop模型输入直接是 224×224。再配合 ONNX Runtime 的 CPU 推理单张耗时能压到 20ms 以内。训练阶段保持随机裁剪没问题但推理阶段追求的是确定性一切多余的预处理都是浪费。5.4 类别不均匀某个类只有 30 张另一个类有 400 张水果数据集下载来源五花八门有的版本石榴照片很少草莓特别多。直接训练的结果是模型对石榴的召回率可能低到 40%。横纵坐标看一眼数据分布就能确认数量少于中位数一半的类别基本就是问题类别。解决方法是两选一要么在DataLoader里配WeightedRandomSampler让每个类别每个 epoch 被抽到的概率大致相同要么干脆把少量类别的图片做针对性增强比如对石榴多做几次旋转和裁剪变体。我一般倾向后者因为WeightedRandomSampler会让模型反复看同几张石榴图对这类样本过拟合的风险更高。5.5 验证 Loss 不降反升梯度爆炸把准确率打到 10%从头训练模型时最容易遇到前几个 epoch loss 一直不降突然一个 epoch 后准确率掉到 10% 以下loss 数值变成几百上千。原因基本是初始学习率设太大梯度一步跨越了最优点。我在调试时把学习率从 0.001 改成 0.1结果第一个 epoch 之后整个模型的权重全乱了。解决方法是给优化器加一个 warmup前 5 个 epoch 让学习率从 0 线性上升到目标值之后再做衰减。PyTorch 自带的torch.optim.lr_scheduler.LambdaLR可以实现warmup_epochs 5 def lr_lambda(epoch): if epoch warmup_epochs: return (epoch 1) / warmup_epochs return 0.1 ** ((epoch - warmup_epochs) // 7) scheduler torch.optim.lr_scheduler.LambdaLR(optimizer, lr_lambdalr_lambda)这个调度器在前 5 个 epoch 做线性预热之后每 7 个 epoch 学习率降一个数量级。配上这个策略即使初始学习率设到 0.01训练也能稳定起步。6. 上线前最后一公里写一个人人可用的单张图片推理脚本训练和评估都通过了最后给一个能直接用的推理脚本不依赖训练时的 DataLoader单张图片输入、类别文字输出。这个脚本可以直接丢给同事用不需要他们懂 PyTorch。from PIL import Image import torch from torchvision import models, transforms device torch.device(cuda if torch.cuda.is_available() else cpu) model models.resnet18(weightsNone) model.fc torch.nn.Linear(512, 11) model.load_state_dict(torch.load(fruits11_best.pth, map_locationdevice)) model.to(device) model.eval() # 类名列表必须与训练时 class_to_idx 保持一致 classes [apple, banana, orange, lemon, kiwi, grape, peach, pear, watermelon, pomegranate, strawberry] predict_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]) ]) def predict_image(img_path, topk3): img Image.open(img_path).convert(RGB) # 统一转 RGB防止灰度图翻车 tensor predict_transform(img).unsqueeze(0).to(device) # (1, 3, 224, 224) with torch.no_grad(): probs torch.softmax(model(tensor), dim1)[0] topk_probs, topk_idx torch.topk(probs, topk) for p, i in zip(topk_probs.tolist(), topk_idx.tolist()): print(f{classes[i]}: {p:.4f}) if __name__ __main__: predict_image(test_apple.jpg)这个脚本的关键细节在最后几行topk默认输出置信度最高的前三个类别而不是只给一个答案。实际使用中这个设计很有用因为模型对某些外观接近的水果青苹果和绿色猕猴桃本身就会混淆输出前三个结果交给业务方判断比硬给一个答案靠谱得多。另外注意Image.open(...).convert(RGB)有些手机拍的图片是 RGBA 模式或灰度模式不转换的话会在ToTensor()阶段维度报错。我的经验是凡是给非算法同事用的推理脚本一定要保证「输入一张图输出一眼能看懂的文字」不要输出整数索引更不要把模型结构定义留在训练脚本里。这个脚本已经用在我的好几个图像分类项目里了现在遇到类似需求我都是先把数据集目录结构确认好再从这个推理脚本反向搭训练流程反而比从头搭更少踩坑。希望帮到你。本文还有配套的精品资源点击获取
返回列表