☰
果蔬分类数据集实战:从4200张标注图到ResNet50与ViT模型落地
2026/10/10 9:31:14 网站建设 项目流程

简介:本资源为常见果蔬多类别图像分类数据集,面向从事图像分类、分割网络改进及计算机视觉项目实践的开发者与学习者,可直接作为分类网络输入使用。数据集共标注36个类别,涵盖香蕉、苹果、梨、葡萄、橙子、黄瓜、胡萝卜、辣椒、洋葱、土豆等常见果蔬,并已划分训练集、测试集与验证集,各类别图片分别存放,便于直接加载训练与评估。压缩包为7z格式,内含约2000个文件,以jpg图像为主,另附1个py脚本与1个json标注文件,整体约364.87MB;运行show脚本可快速可视化数据集分布与样本效果。目前已有119人学习下载。借助该资源,读者可省去繁琐的数据采集与清洗环节,将精力集中于模型结构改进、超参调优与对比实验,同时结合json文件核对类别映射,快速复现分类基线并拓展至分割等下游任务。

1. 果蔬分类数据集怎么选:4,200 张标注图背后的真实门槛

手上有一批约 4,200 张、已经标注好的常见果蔬多类别图像分类数据集,第一反应往往是「直接丢进 ResNet 或 ViT 跑一把」。但我见过太多团队在这一步翻车:类别不均衡、拍摄域单一、标注粒度和任务目标错位,最后模型在验证集上 98%,一上真实货架或分拣线就崩。果蔬图像分类这件事,难点从来不在模型结构,而在数据本身能不能撑住你要落地的那个场景。

这个数据集适合谁?做农产品分拣、智能秤、零售生鲜识别、冷链质检的算法同学,以及想用真实多类别数据练手图像分类、小样本学习、迁移学习的工程师。它能解决的核心问题是:给你一个类别覆盖常见果蔬、规模适中、已标注、可直接切分训练/验证/测试的起点,让你把精力放在预处理、增强、模型选型和部署上,而不是从零爬图标注。接下来我按「先看清数据 → 再跑通基线 → 再调优 → 再避坑 → 最后进阶」的顺序,把这条链路讲透。

2. 先看清 4,200 张果蔬图:类别分布、分辨率与标注格式核查

拿到任何图像分类数据集,别急着写训练脚本。先做三件事:统计类别分布、看分辨率分布、确认标注文件结构。这三步决定了你后面用不用重采样、要不要统一尺寸、标签怎么读。

2.1 用脚本统计类别分布与长尾情况

果蔬类别天然不均衡:苹果、香蕉这类常见品可能几百张,杨桃、秋葵这类可能只有几十张。先跑一段统计脚本,把每个类别的样本数、占比、以及最大/最小类比值打出来。

import os from collections import Counter from pathlib import Path # 假设数据集按类别分文件夹:data/train/苹果/*.jpg root = Path("data/train") counter = Counter() for cls_dir in root.iterdir(): if cls_dir.is_dir(): # 只统计图片文件,过滤掉 .DS_Store 等杂项 n = len([f for f in cls_dir.iterdir() if f.suffix.lower() in (".jpg", ".jpeg", ".png", ".bmp")]) counter[cls_dir.name] = n total = sum(counter.values()) print(f"类别数: {len(counter)}, 总样本: {total}") for cls, n in counter.most_common(): print(f"{cls:12s} {n:5d} {n/total*100:5.2f}%") max_n, min_n = max(counter.values()), min(counter.values()) print(f"最大/最小类比值: {max_n/min_n:.1f}")

逻辑说明:按文件夹名当类别名是最常见的组织方式,脚本遍历一级子目录计数。参数上,suffix.lower()做大小写兼容,避免.JPG漏统计。判读标准:最大/最小类比值超过 10 就要考虑重采样或类别加权;超过 30 基本必须处理,否则模型会偏向头部类。

2.2 分辨率与通道核查:别让统一缩放毁掉细粒度特征

果蔬分类里,同色系不同品类(比如青苹果和青柠)靠的是纹理和形状细节。如果原图分辨率差异大,直接resize到 224 会把小图拉糊、大图压丢细节。先统计宽高分布。

from PIL import Image import numpy as np sizes, modes = [], [] for cls_dir in Path("data/train").iterdir(): if not cls_dir.is_dir(): continue for img_path in cls_dir.iterdir(): if img_path.suffix.lower() not in (".jpg", ".jpeg", ".png"): continue with Image.open(img_path) as im: sizes.append(im.size) # (w, h) modes.append(im.mode) # RGB / L / RGBA ws = np.array([s[0] for s in sizes]) hs = np.array([s[1] for s in sizes]) print(f"宽: min={ws.min()} p50={np.percentile(ws,50):.0f} max={ws.max()}") print(f"高: min={hs.min()} p50={np.percentile(hs,50):.0f} max={hs.max()}") print("颜色模式分布:", Counter(modes))

参数说明:p50是中位数,比均值更能反映典型尺寸。如果中位数在 500 以上,用 224 输入会损失较多细节,可考虑 320 或 384;如果大量图片是灰度(L模式),说明采集设备或场景有限,训练时要统一转 RGB,否则三通道模型读不进去。

2.3 标注格式确认:分类任务也要防标签错位

图像分类的「标注」通常就是文件夹名或 CSV 里的标签列。常见坑是文件夹名带空格、中文编码不一致、或者 CSV 里路径和标签列错位。核查时把标签集合和文件夹集合做一次交集比对。

import pandas as pd # 若标签在 CSV:columns = [filepath, label] df = pd.read_csv("labels.csv", encoding="utf-8") print("标签列唯一值:", sorted(df["label"].unique())) print("是否有空标签:", df["label"].isna().sum()) # 检查路径是否存在,防止标注与文件脱节 missing = [p for p in df["filepath"] if not Path(p).exists()] print(f"缺失文件数: {len(missing)}")

逻辑说明:isna().sum()抓空标签,missing抓标注指向了不存在的文件。这两类问题在合并多个来源的数据时特别常见,不提前清掉,训练时会在DataLoader里随机报错,很难定位。

3. 跑通第一个果蔬分类基线:从 ResNet50 到 ViT 的取舍

数据核查完,先要一个能跑通、能复现的基线。果蔬分类的基线选择,绕不开 ResNet50 和 ViT 这两条路线。我的建议是:数据量在几千张、类别几十个这个量级,先用 ResNet50 打底,再决定要不要上 ViT。

3.1 ResNet50 迁移学习基线:冻结与解冻的两段式训练

4,200 张图从头训 ResNet50 必然过拟合,标准做法是加载预训练权重,先冻结主干只训分类头,再小学习率解冻微调。这是最稳的起点。

import torch import torch.nn as nn from torchvision import models, transforms from torch.utils.data import DataLoader from torchvision.datasets import ImageFolder train_tf = transforms.Compose([ transforms.Resize(256), transforms.RandomResizedCrop(224, scale=(0.7, 1.0)), # 果蔬形状差异大,裁剪范围放宽 transforms.RandomHorizontalFlip(), transforms.ColorJitter(0.2, 0.2, 0.2), # 光照变化常见,加颜色抖动 transforms.ToTensor(), transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]), ]) train_ds = ImageFolder("data/train", transform=train_tf) train_loader = DataLoader(train_ds, batch_size=32, shuffle=True, num_workers=4, pin_memory=True) num_classes = len(train_ds.classes) model = models.resnet50(weights=models.ResNet50_Weights.IMAGENET1K_V2) for p in model.parameters(): # 第一阶段:冻结主干 p.requires_grad = False model.fc = nn.Linear(model.fc.in_features, num_classes) criterion = nn.CrossEntropyLoss() optimizer = torch.optim.AdamW(model.fc.parameters(), lr=1e-3, weight_decay=1e-4)

逻辑说明:RandomResizedCrop的scale=(0.7,1.0)比默认的(0.08,1.0)更保守,因为果蔬主体通常占画面大部分,裁太狠会切掉判别性部位。ColorJitter模拟不同光照,这对生鲜场景很关键。第一阶段只训fc,学习率可以给到 1e-3;第二阶段解冻全部参数时,学习率要降到 1e-4 甚至 1e-5,否则预训练特征会被冲掉。

3.2 ViT 评估时分类头到底调不调

热词里「用 ViT 评估时分类头用调整吗」是个高频疑问。结论是:要调,而且必须换掉。ViT 预训练权重自带的分类头是 ImageNet 1000 类,你的果蔬类别数不是 1000,形状对不上,不换直接报错。换完之后,微调策略有两种。

策略可训练参数适用场景学习率建议
只训分类头仅heads数据极少、算力紧张1e-3
全量微调全部数据几千张以上1e-5 ~ 5e-5
LoRA 微调低秩旁路显存有限、多任务1e-4
from torchvision.models import vit_b_16, ViT_B_16_Weights vit = vit_b_16(weights=ViT_B_16_Weights.IMAGENET1K_V1) for p in vit.parameters(): p.requires_grad = False # 替换分类头,ViT 的分类头挂在 heads.head vit.heads.head = nn.Linear(vit.heads.head.in_features, num_classes) # 只放开新分类头 for p in vit.heads.head.parameters(): p.requires_grad = True

参数说明:ViT 对输入尺寸敏感,vit_b_16要求 224。全量微调时学习率一定要小,ViT 没有卷积的局部归纳偏置,大学习率容易训崩。如果显存吃紧,优先考虑冻结前若干层 Transformer block,只微调后几层加分类头。

3.3 训练循环与验证:早停和混淆矩阵一个都不能少

基线跑起来后,验证阶段别只看准确率。果蔬类别多,准确率会被头部类拉高,必须看每类召回和混淆矩阵。

from sklearn.metrics import confusion_matrix, classification_report def evaluate(model, loader, device): model.eval() preds, labels = [], [] with torch.no_grad(): for x, y in loader: x = x.to(device) out = model(x) preds.extend(out.argmax(1).cpu().numpy()) labels.extend(y.numpy()) print(classification_report(labels, preds, target_names=train_ds.classes, digits=3)) return confusion_matrix(labels, preds) # 早停:验证 loss 连续 5 轮不降就停 best_loss, patience, wait = float("inf"), 5, 0

逻辑说明:classification_report直接给出每类 precision/recall/f1,能立刻看出哪些果蔬被系统性混淆。混淆矩阵进一步告诉你「苹果被认成梨」还是「青柠被认成青苹果」。早停的patience=5是经验值,数据小可以设 3,数据大设 8。

4. 小样本与类别不均衡:1-shot、5-shot 在果蔬分类里怎么落地

果蔬数据集里总有几类样本特别少,这时候常规分类会失效。热词里 1-shot、5-shot、小样本图像分类反复出现,说明这是真实痛点。这一章讲清楚在 4,200 张这个规模下,小样本策略怎么用。

4.1 什么时候该上小样本,什么时候不该

先明确边界:如果每个类别都有 100 张以上,老老实实做常规分类,别碰小样本,收益不明显还增加复杂度。只有当某些类别样本低于 20 张、且你无法补数据时,小样本或度量学习才有意义。判断标准用上一章的类别分布统计,看尾部类有多少。

4.2 用 1-shot / 5-shot 做快速验证的 episode 构造

小样本的核心是 episode 训练:每个 episode 采样 N 个类、每类 K 个支持样本、若干查询样本。下面是一个最小的 episode 采样器。

import random def sample_episode(dataset, n_way=5, k_shot=1, q_query=5): """从 ImageFolder 里采样一个 N-way K-shot episode""" classes = random.sample(dataset.classes, n_way) support, query = [], [] for c in classes: idx = dataset.class_to_idx[c] # 找出该类的所有样本索引 samples = [i for i, (_, y) in enumerate(dataset.samples) if y == idx] random.shuffle(samples) support += [(dataset[i][0], c) for i in samples[:k_shot]] query += [(dataset[i][0], c) for i in samples[k_shot:k_shot+q_query]] return support, query

逻辑说明:n_way是每个 episode 的类别数,k_shot是支持集每类样本数,q_query是查询集每类样本数。1-shot 时k_shot=1,模型只能看一张图就要分类,难度大但能快速暴露特征质量。参数上,训练时n_way可以设 5~10,测试时按你实际要区分的类别数设。

4.3 原型网络做果蔬小样本分类的最小实现

原型网络(Prototypical Network)是小样本里最容易复现的。思路:把每类支持样本过编码器取均值当「原型」,查询样本离哪个原型近就归哪类。

import torch.nn.functional as F def proto_loss(encoder, support, query, n_way, k_shot): # support/query 已是 tensor: [N*K, C, H, W] z_s = encoder(support) # 支持集特征 z_q = encoder(query) # 查询集特征 z_s = z_s.view(n_way, k_shot, -1).mean(1) # 每类原型 # 余弦距离,果蔬纹理差异用余弦比欧氏更稳 logits = F.cosine_similarity( z_q.unsqueeze(1), z_s.unsqueeze(0), dim=2) * 10 target = torch.arange(n_way).repeat_interleave( query.size(0) // n_way) return F.cross_entropy(logits, target)

参数说明:* 10是温度缩放,让余弦相似度进 softmax 前拉开差距,这个系数在 5~20 之间调。mean(1)求原型时,1-shot 下就是单样本特征。编码器可以用前面冻结的 ResNet50 主干,也可以换成更轻的 backbone。

4.4 类别不均衡的加权与重采样

小样本之外,常规训练里的不均衡也要处理。两种手段:损失加权和重采样。损失加权更简单,直接按类别频率的倒数给权重。

counts = torch.tensor([counter[c] for c in train_ds.classes], dtype=torch.float) weights = 1.0 / counts weights = weights / weights.sum() * len(weights) # 归一化 criterion = nn.CrossEntropyLoss(weight=weights)

逻辑说明:weight让尾部类的损失被放大,模型不敢忽略它们。注意权重别拉太极端,否则头部类欠拟合。重采样则用WeightedRandomSampler,让每个 batch 里各类比例接近,但会重复采样尾部类,有过拟合风险。我的习惯是:尾部类样本大于 10 张用加权,小于 10 张才考虑重采样加小样本。

5. 果蔬分类避坑清单:5 个真实踩坑记录

这一章全是血泪经验,每条按「现象 → 原因 → 解决」写,都是我在果蔬类项目里真遇到过的。

5.1 验证集准确率虚高,上线就崩

现象:本地验证 97%,部署到分拣线后错分率飙升。原因:训练和验证图片来自同一批采集,背景、光照、角度高度相似,模型学到了背景捷径而不是果蔬本身。解决:按采集批次或场景切分数据集,确保验证集包含不同光照和背景;加RandomResizedCrop和颜色抖动;必要时做背景替换增强。

5.2 同色系果蔬互相误判

现象:青苹果、青柠、青椒三类互相混淆,召回都上不去。原因:模型主要依赖颜色,而这三类颜色接近,细粒度纹理特征没学到。解决:提高输入分辨率到 320 以上;在增强里减少颜色抖动幅度,避免把颜色线索彻底打乱;引入注意力模块或改用 ViT 捕捉长程纹理。

5.3 DataLoader 随机报「文件损坏」

现象:训练跑几轮后突然报UnidentifiedImageError,重启又能跑。原因:数据集中混入了截断的 JPEG 或非图片文件,num_workers>0时随机命中。解决:训练前用 PIL 全量校验一遍,把坏图移走;或写自定义__getitem__做 try/except 兜底。

from PIL import Image bad = [] for p in Path("data").rglob("*"): if p.suffix.lower() in (".jpg", ".jpeg", ".png"): try: Image.open(p).verify() # verify 只查头部,快 except Exception: bad.append(str(p)) print("坏图:", bad)

5.4 类别名中文导致编码错乱

现象:ImageFolder读出来的classes是乱码,标签对不上。原因:文件夹名是中文,系统默认编码和 Python 读取编码不一致。解决:统一用 UTF-8;或把文件夹名映射成英文/数字 ID,用一份id2name.json维护映射,训练全程用 ID。

5.5 微调学习率设太大,预训练特征被冲毁

现象:解冻主干后第一轮 loss 暴涨,准确率断崖。原因:全量微调时学习率和从头训练一样大,把 ImageNet 学到的特征直接打乱。解决:解冻阶段学习率降到 1e-5 ~ 1e-4,用 warmup 逐步升;或者分层设学习率,主干小、分类头大。

6. 把 4,200 张果蔬图用出上限:分层采样与 TTA 的组合技巧

基线跑通、坑也避了,最后讲一个能把这份数据集榨到极限的具体技巧:分层采样切分加测试时增强(TTA)。很多人把数据一股脑random_split,结果验证集类别分布和训练集不一致,指标忽高忽低。分层切分保证每个类在训练/验证/测试里的比例一致,TTA 则在推理阶段用多视图投票稳住结果。

先做分层切分:

from sklearn.model_selection import train_test_split from collections import defaultdict # 收集 (path, label) samples = [(p, train_ds.class_to_idx[p.parent.name]) for p in Path("data/all").rglob("*.jpg")] paths = [s[0] for s in samples] labels = [s[1] for s in samples] # 先切 train / temp,再切 val / test,stratify 保证分层 X_tr, X_tmp, y_tr, y_tmp = train_test_split( paths, labels, test_size=0.3, stratify=labels, random_state=42) X_val, X_te, y_val, y_te = train_test_split( X_tmp, y_tmp, test_size=0.5, stratify=y_tmp, random_state=42) print(len(X_tr), len(X_val), len(X_te))

逻辑说明:stratify=labels是关键,它让切分后每类的比例和原始一致。random_state固定保证可复现。切完可以再跑一次类别分布统计,确认三个子集的最大/最小类比值接近。

再做 TTA 推理。核心是对同一张图做多种变换,把预测概率平均:

import torch.nn.functional as F tta_tf = [ transforms.Compose([transforms.Resize(256), transforms.CenterCrop(224), transforms.ToTensor(), transforms.Normalize(mean, std)]), transforms.Compose([transforms.Resize(256), transforms.RandomHorizontalFlip(p=1.0), transforms.CenterCrop(224), transforms.ToTensor(), transforms.Normalize(mean, std)]), transforms.Compose([transforms.Resize(288), transforms.CenterCrop(224), transforms.ToTensor(), transforms.Normalize(mean, std)]), ] def predict_tta(model, img_path, device): model.eval() probs = [] with torch.no_grad(): for tf in tta_tf: x = tf(Image.open(img_path).convert("RGB")).unsqueeze(0).to(device) probs.append(F.softmax(model(x), dim=1)) return torch.stack(probs).mean(0) # 多视图概率平均

参数说明:三个视图分别是原尺度中心裁剪、水平翻转、放大后裁剪。果蔬左右翻转通常语义不变,所以翻转视图安全;但上下翻转对某些有方向性的品类(比如香蕉)可能不合理,要按品类决定。mean(0)是概率平均,比投票更平滑。实测在果蔬分类上,TTA 一般能带来 1~3 个点的提升,代价是推理耗时翻三倍,线上要权衡。

我自己的习惯是:切分脚本和数据校验脚本写成一套,每次换数据集先跑校验再跑切分,把类别分布和坏图清单存成日志。这样后面无论换 ResNet 还是 ViT、上不上小样本,数据这一层始终是干净的。果蔬分类没有玄学,把数据看透、把基线跑稳、把坑记牢,4,200 张图足够撑起一个能落地的分类器。希望帮到你。

本文还有配套的精品资源,点击获取

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询