资讯动态

36类果蔬图像分类实战:基于PyTorch与ResNet-18的完整流程

发布时间:2026/9/15 14:38:56 来源:尧图企业网站定制
简介面向图像分类任务的果蔬数据集共包含36个常见类别的已标注图像约3400张覆盖香蕉、苹果、梨、葡萄、橙子、猕猴桃、西瓜、石榴、菠萝、芒果等常见水果以及黄瓜、胡萝卜、辣椒、洋葱、马铃薯、番茄、萝卜等常见蔬菜。数据已经过预处理可直接作为分类网络的输入并划分为训练集与验证集各类别图像分目录存放便于加载与评估。压缩包内共2000个文件以jpg图像为主辅以1个py可视化脚本和1个json类别配置文件整体包体大小94.47MB结构简洁。运行show脚本可快速完成数据集可视化浏览。目前已吸引212人学习下载适合深度学习初学者或需要快速获得标准果蔬分类数据的研究者也可配合作者博客中图像分类网络改进与计算机视觉项目笔记进行拓展实践。1. 36类果蔬图像分类数据集先确认数据再谈模型拿到一个 36 类果蔬图像分类数据集时我最先做的不是写模型而是确认这批数据能不能直接交给网络。约 3400 张图、每类平均不到 100 张划分了训练集和验证集还带一份 JSON 标注——这个体量刚好够做一次图像分类基线验证。对于刚接触分类任务的人它是理解“数据 → 加载器 → 模型 → 评估”全流程的最小闭环对于有经验的工程师它能用来快速测试数据增强、类别均衡和迁移学习策略。本文从数据目录讲起逐步到 DataLoader、ResNet-18 微调和最后的错误分析。2. 数据组织与标注解析先用脚本摸清 3400 张图的真实分布2.1 训练集/验证集的目录约定解压后数据集的主目录大概长这样dataset/ ├── train/ │ ├── 香蕉/ │ │ ├── Image_7.jpg │ │ ├── Image_1.jpg │ │ └── ... │ ├── 苹果/ │ ├── ... ├── val/ │ ├── 香蕉/ │ └── ... └── labels.jsontrain 和 val 内部按类别名建子目录图片是统一的Image_id.jpg命名。这个结构可以直接用torchvision.datasets.ImageFolder读取但我不建议上来就训练原因有两个。第一ImageFolder的类别顺序按目录名字排序如果 JSON 里也是字符串类别名很容易出现重复项直接映射会丢数据原类别列表里“辣椒”和“萝卜”就出现过重复实际要以 JSON 内容为准。第二小数据集需要仔细看每个类到底有多少图我拆过的项目里最多的一类可能比最少的一类多 3 倍以上。所以先把目录和 JSON 对齐再进加载器。2.2 标注 JSON 的结构与解析labels.json 常见结构是 dict里面同时包含 categories、train、val 三段。我在本地把它读过一遍用下面这段脚本可以兼容几种常见格式import json with open(labels.json, r, encodingutf-8) as f: ann json.load(f) if isinstance(ann, dict) and train in ann and val in ann: train_map ann[train] val_map ann[val] categories ann.get(categories, list(train_map.keys())) elif isinstance(ann, dict) and records in ann: records ann[records] train_map {} val_map {} categories [] for r in records: split r.get(split, train) label r[label] fname r[file_name] target train_map if split train else val_map target.setdefault(label, []).append(fname) if label not in categories: categories.append(label) else: raise ValueError(Unknown JSON schema) print(categories num:, len(categories))这段代码的重点是用setdefault把同一个类别的图片名聚合到列表里避免手动判断 key 是否存在。如果你是别的项目拿到的 JSON通常只需要调整label和file_name两个字段名。输出 categories num 时如果少于 36基本可以断定 JSON 或目录里存在重名类别需要人工去重后再映射成整数索引。2.3 类别数量统计与异常项检查拿到 train_map、val_map 后我一般会先跑一个数量统计def split_stats(split_map): return {k: len(v) for k, v in split_map.items()} train_stat split_stats(train_map) val_stat split_stats(val_map) total sum(train_stat.values()) sum(val_stat.values()) print(total images:, total) print(train images:, sum(train_stat.values())) print(val images:, sum(val_stat.values())) # 找出图片数量明显偏少的类别 bind {k: v for k, v in train_stat.items() if v 50} print(classes with 50 train images:, bind)输出结果可能如下表所示数值是我本地一次统计得到类别训练集数量验证集数量香蕉9224苹果8822梨8420葡萄9023橙子8521猕猴桃7819西瓜7418石榴7017菠萝7619芒果8020剩下 26 类也基本在 7095 张之间总量约 3400。如果某个类只有 30 张那么训练时这个类在 DataLoader 里的采样概率会偏低需要在后面用WeightedRandomSampler补回来。这一步统计还有个隐藏作用验证集分布最好和训练集分布一致如果某一类在验证集里只有 5 张Top-1 准确率波动会非常大至少要有 15 张以上才比较可信。3. 数据加载与预处理把 3400 张图变成可直接训练的张量3.1 自定义 Dataset 的写法虽然ImageFolder够用但对于只有 3400 张图的小数据集我更喜欢写一个自定义 Dataset因为要在__getitem__里同时控制 label 映射、图片读取和错误处理。这样后面做难例挖掘、按类别抽样都方便。from torch.utils.data import Dataset from PIL import Image import os class FruitVegDataset(Dataset): def __init__(self, split_map, root_dir, categories, transformNone): self.samples [] self.categories categories self.transform transform for label_name, fnames in split_map.items(): label_idx categories.index(label_name) for fname in fnames: self.samples.append((os.path.join(root_dir, label_name, fname), label_idx)) def __len__(self): return len(self.samples) def __getitem__(self, idx): path, label self.samples[idx] image Image.open(path).convert(RGB) if self.transform: image self.transform(image) return image, label这里categories.index(label_name)会把中文类别名变成整数索引convert(RGB)是为了兼容灰度图或带透明通道的 PNG。如果某些图片读取失败可以在__getitem__里加一个try/except直接把损坏样本换成同类的下一张但要在异常时打印路径不要静默处理。root_dir需要传 train 或 val 的根目录拼接时用os.path.join避免手写字符串拼路径导致 Windows/Linux 分隔符问题。3.2 训练集和验证集的 transform小数据集最容易踩的坑就是验证集做了训练增强。训练集为了泛化要做随机裁剪和颜色扰动验证集必须固定尺寸否则同一张图每次评估的结果都不一样无法对比模型好坏。from torchvision import transforms train_transform transforms.Compose([ transforms.RandomResizedCrop(224, scale(0.6, 1.0)), transforms.RandomHorizontalFlip(p0.5), 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(256), transforms.CenterCrop(224), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]) ])注意RandomResizedCrop的scale从 0.6 到 1.0对果蔬这种主体占比较大的图片比较合适。如果拉到 0.08会和 ImageNet 默认值一样过强的裁剪可能导致番茄这类小块辨别困难。ColorJitter的 brightness、contrast、saturation 都设成 0.2 也是折中值太大会让绿色蔬菜的纹理失真。3.3 类别均衡WeightedRandomSampler统计分布后如果确认某些类别图片偏少可以用WeightedRandomSampler让每个 epoch 的样本权重平均。from torch.utils.data import WeightedRandomSampler def make_weights(split_map, categories): weights [] for label_name, fnames in split_map.items(): idx categories.index(label_name) weights.append(1.0 / len(fnames)) sample_weights [] for label_name, fnames in split_map.items(): for fname in fnames: sample_weights.append(weights[categories.index(label_name)]) return sample_weights weights make_weights(train_map, categories) sampler WeightedRandomSampler(weights, num_sampleslen(weights), replacementTrue)这个权重的思路是反比于类别样本数而不是直接给一个固定常数。num_samples保持和训练集总数一致每个 epoch 能采到全部样本的等效量级。使用 sampler 之后DataLoader 就不能再设shuffleTrue否则会冲突。实际操作中如果某个类图片质量差即使权重提上去了准确率也上不去这时候优先查数据而不是继续调采样策略。4. 从零训练基线与迁移学习先跑通 ResNet-18再谈 Transformer4.1 为什么先选 ResNet-18而不是 ViT检索“最新图像分类模型”时经常看到 ViT、Swin Transformer 这类模型但小数据集上直接训练 Transformer 并不会比 CNN 好。3400 张图只能支撑微调而 ResNet-18 参数少、结构简单适合做第一版基线。先把准确率跑到 90% 左右再换 MobileNetV3 或 EfficientNet 对比才是稳健做法。果蔬分类的类别间差异比较明显不像 ImageNet 那样需要很强的长距离建模能力CNN 在这个场景下完全够用。4.2 加载预训练权重并替换分类头import torchvision.models as models import torch.nn as nn model models.resnet18(weightsmodels.ResNet18_Weights.IMAGENET1K_V1) num_classes len(categories) model.fc nn.Linear(model.fc.in_features, num_classes)pretrained参数在最新 torchvision 里已经标记为废弃推荐用weights...方式。替换fc时注意in_features是 512直接写 512 也可以但用model.fc.in_features能避免换骨干网络时记错。这里只替换了最后一层全连接前面的卷积层全部保留 ImageNet 的预训练特征。对于果蔬数据前几层提取的纹理、边缘特征非常通用不需要重新学习。4.3 训练循环、学习率与早停import torch import torch.nn as nn from torch.optim import AdamW from torch.optim.lr_scheduler import CosineAnnealingLR device cuda if torch.cuda.is_available() else cpu model.to(device) criterion nn.CrossEntropyLoss() optimizer AdamW(model.parameters(), lr1e-4, weight_decay1e-4) scheduler CosineAnnealingLR(optimizer, T_max20) epochs 30 best_acc 0 for epoch in range(epochs): model.train() train_loss 0 for images, labels in train_loader: images, labels images.to(device), labels.to(device) optimizer.zero_grad() logits model(images) loss criterion(logits, labels) loss.backward() optimizer.step() train_loss loss.item() scheduler.step() model.eval() correct 0 total 0 with torch.no_grad(): for images, labels in val_loader: images, labels images.to(device), labels.to(device) logits model(images) pred logits.argmax(dim1) correct (pred labels).sum().item() total labels.size(0) acc correct / total print(fepoch {epoch1}: train_loss{train_loss/len(train_loader):.4f}, val_acc{acc:.4f}) if acc best_acc: best_acc acc torch.save(model.state_dict(), best_resnet18.pth)这段代码里AdamW的weight_decay1e-4对 36 类小数据集是常用配置太大导致欠拟合太小容易过拟合。CosineAnnealingLR的T_max20表示学习率在 20 个 epoch 内完成一个余弦周期配合 30 个 epoch 会在后半段自动降到接近 0。保存模型用的是验证集最高准确率那一版而不是最后一个 epoch这是避免验证集波动丢精度的关键。如果你的显卡显存不够可以把 batch size 从 32 降到 16同时学习率按比例降为 5e-5。5. 可视化与错误分析show 脚本和混淆矩阵才是提分关键5.1 用 show 脚本检查图片质量资源自带的 show 脚本可以把每个类拼成一张网格图运行方式一般类似python show.py --data train --category 香蕉。我跑通后最先检查三样东西裁切是否切到主体、有没有不属于该类的图片、以及标注顺序和目录名是否一致。3000 多张图看一遍不现实调成每类随机抽 8 张网格排列缩略图扫一眼就够了。重点看那些训练集数量偏少的类比如“石榴”和“甜菜根”这两个类别容易混入相似背景的图片会在后面明显拉低准确率。5.2 混淆矩阵定位高频错误import numpy as np from sklearn.metrics import confusion_matrix cm confusion_matrix(all_labels, all_preds) np.fill_diagonal(cm, 0) max_idx np.argwhere(cm cm.max()) for i, j in max_idx: print(categories[i], -, categories[j], cm[i][j])all_labels和all_preds是验证集所有样本的标签和预测结果需要提前在验证循环里收集。np.fill_diagonal(cm, 0)把对角线置零后剩下的最大值就是最常见的混淆对。果蔬分类里比较典型的是“辣椒”和“甜椒”、“萝卜”和“甜菜根”这两对在颜色和形状上都有交叠。5.3 Focal Loss、标签平滑和难例挖掘如果混淆对集中在少数类别可以给损失函数加上 Focal Loss公式上就是在交叉熵基础上乘(1 - p_t)^gamma让模型更关注难样本。对于果蔬分类gamma 设 1 或 2 就够别设太大。另一个更轻量的做法是标签平滑把 CrossEntropyLoss 的label_smoothing参数设为 0.05能减少模型对训练集标注噪声的过拟合。如果某个类别始终和另一个类混淆比较有效的方法是去掉该类增强里的 ColorJitter并在采样时对该类提高重复概率先把“能分开”的特征稳住再谈泛化。本文还有配套的精品资源点击获取

读完文章,也想定制专属网站?

尧图设计师 24 小时内与您沟通定制方案

免费获取报价