资讯详情

水果分类数据集实战指南:从解压到PyTorch训练全流程

📅 2026/10/10 3:06:31 | 华诺云谱 👁 阅读
水果分类数据集实战指南:从解压到PyTorch训练全流程
简介这份压缩包是面向机器学习初学者的水果图像分类数据集包含了苹果、香蕉、葡萄、橙子、梨五种常见水果的带标签图片整体按类别子目录组织文件名也保留了标签信息可直接用于监督学习和计算机视觉模型训练。包内共有1310个文件主体为超过1300张图像文件另附2个列表文件、1个配置文件与1个辅助脚本总大小仅14.07MB加载非常方便。该数据集已有3625人学习使用适合用来练习数据组织、图像预处理、训练集划分和分类指标评估等关键环节。借助这些数据读者可以动手构建并训练卷积神经网络等模型经历从数据读取、特征学习到精度与召回率计算的完整流程扎实掌握图像分类任务的核心思路也为后续数据挖掘项目积累实践经验。1. 水果分类数据集动手做图像分类先别急着写模型水果分类数据集比如常见的 fruits分类数据集.rar 这一类压缩包是很多开发者接触真实图像分类任务的第一站。解压开它你看到的不是 MNIST 那种已经帮你处理成黑白小图的干净数据而是几十个按品种命名的文件夹图片尺寸参差、光照混乱、背景杂乱这份“不整齐”恰好逼着你在训练前把数据这一课补齐。它的核心价值在于类别语义足够清晰你不需要纠结标注口径能把从数据准备、格式转换、模型训练到结果评估的完整链路一次走通。它不会替你把脏活累活干完但能让你把图像分类的基本功练扎实。适合刚学完 CNN、想摆脱玩具数据、或者准备在迁移学习上做实操的开发者。2. 拆开 rar 先看数据结构类别分布、图片格式与清晰度陷阱2.1 解压第一步路径、编码与解压工具的三个注意点拿到 .rar 文件之后先别急着双击解压。常见做法是把它放到一个纯英文的目录里再用手头的解压工具释放。这里第一个注意点是目标路径避免中文和空格原因是后面 PyTorch 的 DataLoader 在 Windows 上读取含中文路径时偶尔会报编码错误而带空格的路径会让某些脚本库的路径拼接解析出问题。第二个注意点是解压时优先选“解压到 fruits分类数据集/”而不是“解压到当前文件夹”因为数据集打包者习惯在最外层套一层同名目录选错了解压方式就多出一层嵌套后续脚本里的相对路径全部要跟着改。第三个注意点是解压完成后先扫一眼文件夹名和类别名是否正常某些工具在 Windows 下解压打包自 Linux 的 rar 时会出现中文乱码水果类别名一旦乱码后面 class_to_idx 映射就全乱了。这一步不需要写代码但值得花两分钟确认。路径问题在训练到一半时才暴露的话你的耐心会被消耗得非常快。我一般会在解压目录旁边建一个 readme.txt顺手记录压缩包来源、解压时间、总大小虽然看似多余但数据集迭代到第二版、第三版时这份记录就是后悔药。提示数据集解压路径做好全程无中文、无空格避免训练到第 10 个 epoch 才被一个奇怪的路径报错打断。2.2 统计类别和数量一条 os.walk 扫描脚本解压之后我习惯先统计数据全貌而不是直接打开文件夹凭肉眼估数量。用 os.walk 把每个子文件夹里的图片数量扫出来一次看清三类问题类别数是否符合预期各类数量是否均衡有没有大小为 0 或几百字节的损坏文件。这条脚本在后续排查训练问题时还会反复用到值得留在工程目录里。import os import collections data_root rF:\datasets\fruits # 改成你自己的解压路径 class_counts collections.Counter() total_images 0 bad_files [] for root, dirs, files in os.walk(data_root): for f in files: if f.lower().endswith((.jpg, .jpeg, .png)): cls os.path.basename(root) class_counts[cls] 1 total_images 1 # 文件低于 512 字节的基本可以判定为损坏 fpath os.path.join(root, f) if os.path.getsize(fpath) 512: bad_files.append(fpath) print(类别数:, len(class_counts)) print(总图片数:, total_images) for cls, cnt in class_counts.most_common(): print(f{cls}: {cnt}) print(疑似损坏文件:, bad_files)核心逻辑是把子文件夹名当作类别标签统计所以解压后的目录结构一旦被打乱统计结果就失去意义了。参数上三个点需要注意data_root 用原始字符串 r 防止反斜杠转义问题后缀名过滤只保留 jpg、jpeg、png 三种常见格式如果数据集里混有 bmp、webp 需要自己补后缀文件大小小于 512 字节的图片绝大多数是下载中断产生的半截文件值得单独列出来人工检查。跑完这条脚本后你会得到一组关键数字。类别数对上了说明压缩包没缺目录各类数量差距在 5 倍以内做均衡训练问题不大如果某个类别只有十几张图后续就要考虑加权采样或者直接放弃这个类别。我在某个图像处理 Demo 里遇到过“梨”目录只有 9 张图训练时准确率怎么都上不去最后发现是冷门类别拖后腿剔除后整体准确率反而涨了三个百分点。2.3 图片尺寸与真实格式PIL 抽样扫描防止解码翻车统计完数量之后下一步是看图片本身的尺寸分布和真实编码格式。很多爬虫收集的数据集保留了原始分辨率有的图是 4000x3000有的是 640x480混在一起训练会让模型被迫适应不同尺度下的特征表达。更隐蔽的问题是“假后缀”——文件叫 .jpg 但实际是 PNG 编码文件叫 .png 但里头是 BMP 数据靠后缀判断格式的加载函数会在训练中段频繁报“解码失败”。from PIL import Image import os, collections data_root rF:\datasets\fruits size_counter collections.Counter() format_counter collections.Counter() checked 0 max_check 500 # 抽样上限避免全量扫描太慢 for root, dirs, files in os.walk(data_root): if checked max_check: break for f in files: if checked max_check: break if not f.lower().endswith((.jpg, .jpeg, .png)): continue fpath os.path.join(root, f) try: with Image.open(fpath) as im: size_counter[im.size] 1 format_counter[im.format] 1 checked 1 except Exception as e: print(打不开的图片:, fpath, e) print(尺寸分布 Top10:, size_counter.most_common(10)) print(真实编码格式分布:, format_counter)这段脚本是抽样遍历每打开一张图就记录尺寸和真实格式把打不开的坏图一并列出。PIL 的 Image.open 是惰性加载只有在访问 .size 和 .format 时才真正解析文件头所以 with 块内不用手动 close。为什么只抽 500 张而不是全量一张高分辨率图片解码约几十毫秒几千张全量扫一遍要几分钟而分布规律在 500 张抽样里已经足够明显如果怀疑某个特定类别有问题再对该类别单独全量扫描。这一步的输出直接指导预处理基线。如果大部分图集中在 500x400 到 2000x1500 之间统一 Resize 到 256 再 CenterCrop 224 是合理的如果发现一批 4000x3000 的全景图和一批 200x200 的缩略图混在一起就更应该考虑用倍数缩放而不是直接拉伸。这个决定影响后面所有实验的输入分布早做比晚做好。2.4 隐藏的图片属性EXIF 旋转与 RGBA 通道怎么处理抽样扫描尺寸和格式之后还有两个容易被忽略的隐藏属性。第一个是 EXIF 旋转信息很多手机拍摄的照片在 EXIF 里记录 Orientation 字段某些解码器忽略它直接按原始像素读图片显示是横的但实际数据是竖的模型训练时同一类别里混着横竖两种方向的图特征提取就多了一个不需要的变量。处理方式是在 Dataset 的getitem里加一步统一处理或者更干脆——在预处理阶段把所有图转存成 RGB 三通道的干净副本。第二是通道数jpg 图片通常是三通道但 PNG 可能是 RGBA 四通道还可能出现单通道灰度图。预训练模型第一层卷积接受三通道输入四通道直接报 mismatch单通道虽然能广播但特征分布不匹配。我在这类数据集的 Dataset 里强制 .convert(RGB)就是为同时消化这两种情况。颜色对水果分类是强判别特征灰度图丢失颜色信息后青苹果和红苹果的区分度大幅下降遇到这种情况最好把灰度图也转成三通道副本再做后续增强。3. 用 PyTorch 加载水果数据集从文件夹到 DataLoader 的标准流程3.1 按文件夹结构划分训练集与验证集随机种子与比例的选择分类数据集最常见的组织方式是每个类别一个文件夹这是 torchvision.datasets.ImageFolder 能直接消费的结构。但直接拿全量数据训练再在同样的数据上评估得到的是虚高准确率属于典型的训练集泄漏。常见做法是在训练前按类别把图片拆出训练集和验证集拆分的粒度直接决定后续实验能不能复现。python -c import os, random, shutil src rF:/datasets/fruits dst rF:/datasets/fruits_split random.seed(42) val_ratio 0.2 for cls in os.listdir(src): cls_path os.path.join(src, cls) if not os.path.isdir(cls_path): continue imgs [f for f in os.listdir(cls_path) if f.lower().endswith((.jpg,.jpeg,.png))] random.shuffle(imgs) val_n int(len(imgs) * val_ratio) for i, f in enumerate(imgs): target val if i val_n else train out_path os.path.join(dst, target, cls) os.makedirs(out_path, exist_okTrue) shutil.copy(os.path.join(cls_path, f), os.path.join(out_path, f)) print(split done:, dst) 用 python -c 直接执行省去临时脚本文件的管理。核心逻辑是按类别遍历类别内部自行打乱取前 20% 进验证集、后面的进训练集。三个参数我要逐一说明random.seed(42) 保证任何人跑同一份脚本得到完全相同的划分结果团队协作时大家的 baseline 才可比val_ratio0.2 是验证集比例总数据量在几千张时 20% 合理数据量越大这个比例可以越小上万张数据集用 10% 就够shutil.copy 用的是复制而非移动调试阶段原始数据不动是最稳妥的确认划分无误再删源文件不迟。为什么不直接 torchvision 的 random_split它在内存里划分索引结果不落地换框架或换机器就得重来而文件层面的划分生成干净的目录结构PyTorch、TensorFlow 这类主流框架都能直接消费肉眼也能检查验证集里有没有异常图。粗看多写了几行代码后续省的时间远超这几分钟。还有个容易忽略的点这个脚本要按类别分别打乱而不是全量文件打乱后再分否则某些类别可能全部落进训练集验证集里直接缺了类别。3.2 自定义 Dataset 类标签映射与在线增强的串联方式目录结构就绪后写一个继承 Dataset 的类来管理样本列表。常见做法是在init里把所有图片路径和对应标签提前存成列表避免训练过程中反复 os.listdir 扫描目录。这样内存换速度样本量在几十万以内都没有问题。import torch from torch.utils.data import Dataset, DataLoader from torchvision import transforms from PIL import Image import os class FruitDataset(Dataset): def __init__(self, root, transformNone): self.transform transform self.samples [] self.classes sorted([d for d in os.listdir(root) if os.path.isdir(os.path.join(root, d))]) self.class_to_idx {c: i for i, c in enumerate(self.classes)} for cls in self.classes: cls_dir os.path.join(root, cls) for f in os.listdir(cls_dir): if f.lower().endswith((.jpg, .jpeg, .png)): self.samples.append((os.path.join(cls_dir, f), self.class_to_idx[cls])) def __len__(self): return len(self.samples) def __getitem__(self, idx): path, label self.samples[idx] img Image.open(path).convert(RGB) if self.transform: img self.transform(img) return img, label关键点在 self.classes 用了 sorted 排序。文件夹的遍历顺序在不同操作系统上不一致显式排序能保证每次运行时类别索引的对应关系稳定否则这次“苹果”是第 0 类下次可能变成第 5 类checkpoint 里的预测层权重就对不上了。Image.open 之后接 .convert(RGB) 的原因从第二章延伸过来RGBA 和灰度图统一转三通道让后续的预训练模型输入形状固定。Dataset 本身不感知增强逻辑增强全部交给 transform 参数这样训练和验证可以复用同一个 Dataset 类只换变换管线。3.3 数据增强变换的参数怎么设训练集做加法验证集做减法数据增强代码写在 Dataset 外面通过 transform 参数传进来。训练集和验证集应该用完全不同的变换这是不少新手翻车的地方。训练集要引入随机性来扩大数据分布验证集要保证每次评估都在同一视角上判断否则结果被随机性盖住你没法判断模型是变好了还是运气好。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.4但别用 ImageNet 默认的 (0.08, 1.0)——那会把水果裁成一堆背景噪声。ColorJitter 的三个 0.2 控制亮度、对比度、饱和度的抖动振幅自然光照数据上效果明显花哨一点的做法是加 RandomRotation 和 RandomAffine但旋转角度超过 20 度会让“苹果”看起来不像“苹果”我一般不这么干。Normalize 的 mean 和 std 用的是 ImageNet 统计值迁移学习场景下这是标配别自己算一套新的替换掉除非你从头训练模型没有加载预训练权重。验证集的 Resize 到 256 再 CenterCrop 224 是评估标准动作Resize 尺寸和 Crop 尺寸相差 32 像素是为了让中心裁剪有一点缩放余量避免直接 224 裁剪丢失边缘信息。3.4 DataLoader 调参batch_size、num_workers 与 pin_memory 的取舍Dataset 写好之后DataLoader 的参数选择直接影响训练速度和稳定性。我常用的配置是 batch_size32、num_workers4、pin_memoryTrue但这不是万金油换机器要按实际调整。train_loader DataLoader(train_ds, batch_size32, shuffleTrue, num_workers4, pin_memoryTrue) val_loader DataLoader(val_ds, batch_size64, shuffleFalse, num_workers4, pin_memoryTrue)batch_size 决定每次梯度更新的近似质量。水果分类图片经 Resize 后输入是 224x224x332 张一批在单卡 8GB 显存上大约占 4GB 左右含激活值如果是 6GB 卡降到 16 或者 8 更稳妥。num_workers 是数据加载子进程数Windows 上建议 4 到 6设太高进程调度开销反而变慢Linux 上可以适度调到 8。pin_memoryTrue 把数据锁页传输到 GPU减少一个拷贝环节代价是系统内存占用上升机器内存小于 16GB 时关掉它。容易踩的坑是 shuffle 参数训练集必须 True验证集必须 False。验证集打乱对结果没有数学影响但会让每次评估的 batch 组成不稳定排查问题时多一个变量。另一个老毛病是 num_workers 在 Windows 上需要配合 ifname main 使用否则多进程会递归创建子进程报错这是 PyTorch 在 Windows 上特有的麻烦Linux 和 mac 不必担心。4. 水果分类数据集避坑指南标注错误、类别不平衡与训练集泄漏数据集本身是静态的但训练过程是动态的很多问题不是一眼能看出来的而是训练到一半才浮出水面。这一章我按踩坑频率从高到低把最常见的问题写成“现象→原因→解决”三段式方便你遇到类似情况时直接对照。4.1 现象验证集准确率奇高测试集却翻车我最初在某次模拟项目X里就遇到过验证集准确率 97%拿模型去识别新拍的水果照片错误率高得离谱。反复排查确认是两个原因叠加。第一个是划分验证集时没有固定随机种子每次跑脚本划分结果不同训练集和验证集高度相似第二个是数据集打包时同一类水果的连续拍摄帧天然形成数据簇随机划分让同一簇的图片同时出现在训练集和验证集模型实际上是在“朝题背答案”表面准确率自然虚高。解决方法是划分前必须固定 random.seed同时查看图片命名规律。如果发现文件名是 IMG_0001 到 IMG_0100 这种连续编号说明同一批照片可能来自同一次拍摄正确的划分方式是按编号区间切片而不是随机打乱更严格的做法是把文件按名称 hash 后分桶确保同一簇图片只落在同一个集合。这一步做扎实验证集准确率才真的有参考意义。4.2 现象某些类别怎么训练都记不住训练时整体 loss 正常下降查看分类报告后发现某品种的准确率长期在五成以下其他类别都在九成以上。统计类别数量才发现这个品种只有 20 余张图热门类别有数百张。水果数据集里的类别不平衡很常见爬虫收集阶段热门水果容易拿得多冷门品种能凑齐 30 张已经不错。优先试三个办法按成本排序。第一是用 WeightedRandomSampler 提高少数类的采样概率改造成本最低几行代码就能让每个 batch 里冷门类别不再缺席第二是在 CrossEntropyLoss 里传 class_weight 给少数类加权训练时梯度更新更偏向少数类第三是对少数类做针对性增强比如固定做左右翻转、随机旋转 15 度、轻微色彩抖动。前两个方法本质上在调整样本权重第三个在增加样本多样性可以叠加使用但别把权重拉太极端否则模型会过度记忆少数类的那几张图。4.3 现象训练中途 loss 变成 NaN训练到第 30 轮左右 loss 突然跳成 NaN后续每个 epoch 都是 NaNcheckpoint 里的准确率节奏全部作废。这种翻车通常有两个来源一是学习率过大导致梯度爆炸二是数据里有全零或极端像素的损坏图片送进网络后产生无效梯度。排查顺序我一般先查数据再查优化器。先回第二章的 PIL 扫描脚本把损坏图片揪出来重点关注大小为 0 和纯黑色、纯白色的图数据确认干净后把优化器换成带梯度裁剪的 Adam或者把初始学习率从 1e-3 降到 3e-4。顺带说一个玄学规律loss 在某个 epoch 突然从个位数跳到几十大概率是数据增强里 ColorJitter 的数值设太大把某张图抖动成了全白或全黑这种极端输入比坏文件更隐蔽肉眼还看不出来。4.4 现象训练集准确率 100%验证集只有 70%这是教科书级的过拟合信号但在水果数据集上它有特殊诱因。如果直接拿 ResNet18 甚至 ResNet50 在几千张图上从头训练模型参数远超样本信息量背下训练集是必然结果。很多人第一反应是加 Dropout 或者换更大的模型方向反了——数据量固定时模型越大越容易记住训练集。我的习惯是先看数据增强够不够泼辣把 RandomResizedCrop 的 scale 从 (0.6, 1.0) 加到 (0.4, 1.0)再考虑模型正则化。如果增强已经很强但验证集还是上不去就只训练分类头特征提取层保持冻结或者换 MobileNetV3 这类小模型参数量少反而更容易学到泛化特征。4.5 现象数据里混入了非水果图片数据集整理时人工标注难免走神有的目录里混入背景墙、人手的特写、甚至其他无关物体的图片。这类脏数据不会引发报错也不会让 loss 跳变但会稳定地压低准确率。检测方法藏在第六章的混淆矩阵里——把验证集里被分错的样本图攒成一张拼图人眼扫一遍就能发现那些“标注是苹果但图里是木头桌子”的硬错误。这类问题调参永远解决不了只能回数据层做清洗把错误的标注修正或删除。5. 用 ResNet18 在这个数据集上跑通完整训练最小可行代码与参数5.1 迁移学习配置为什么先冻结特征层在选择模型时ResNet18 不是唯一选项。我看到有人一上来就用 EfficientNet 甚至 Swin Transformer小数据集上表现反而不如 ResNet18 稳定。原因在于水果分类的难点是品种间的细微差异而不是场景复杂度18 层残差网络搭配 ImageNet 预训练权重特征提取能力足够训练时间、显存占用、调试成本都低一个量级。等 ResNet18 跑通了再换大模型不迟。import torch import torch.nn as nn from torchvision import models num_classes 16 # 改成统计脚本输出的类别数 model models.resnet18(weightsmodels.ResNet18_Weights.IMAGENET1K_V1) model.fc nn.Linear(model.fc.in_features, num_classes) # 冻结特征提取层只训练分类头 for name, param in model.named_parameters(): if not name.startswith(fc): param.requires_grad False optimizer torch.optim.Adam(model.fc.parameters(), lr1e-3, weight_decay1e-4) criterion nn.CrossEntropyLoss()冻结特征层的动机很直接预训练模型的浅层和中层已经学到了边缘、纹理、形状这些通用特征水果品种的区分更多依赖这些特征的高层组合而分类头就是做线性组合的部分。只训练分类头时可训练参数约 1 万多几分钟就能在一个完整 epoch 上过一遍适合先跑通流程之后再考虑解冻部分卷积层提升精度。注意num_classes必须和第三章统计脚本输出的类别数一致差一个数字训练时就会报 shape mismatch。5.2 训练循环与 checkpoint 保存不只存最后一个 epoch训练循环本身没有魔法但 checkpoint 的保存策略常被忽略。只保存最后一个 epoch 的模型如果验证集准确率在第 14 轮达到峰值、第 20 轮回落你手里就只有次优模型。正确做法是每轮验证准确率刷新就覆盖保存这样峰值模型永远在手。device torch.device(cuda if torch.cuda.is_available() else cpu) model.to(device) best_acc 0.0 num_epochs 30 for epoch in range(num_epochs): model.train() running_loss 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() running_loss loss.item() * images.size(0) model.eval() correct 0 total 0 with torch.no_grad(): for images, labels in val_loader: images, labels images.to(device), labels.to(device) outputs model(images) _, preds torch.max(outputs, 1) correct (preds labels).sum().item() total labels.size(0) val_acc correct / total avg_loss running_loss / len(train_ds) print(fEpoch {epoch1}/{num_epochs} loss{avg_loss:.4f} val_acc{val_acc:.4f}) if val_acc best_acc: best_acc val_acc torch.save({ epoch: epoch, model_state_dict: model.state_dict(), optimizer_state_dict: optimizer.state_dict(), best_acc: best_acc, }, best_fruit_model.pth) print(f新最佳模型已保存, acc{best_acc:.4f})这段代码里有两个不能省的细节。model.eval() 切换 BatchNorm 和 Dropout 到推理模式如果漏了验证阶段 BatchNorm 会使用当前 batch 的均值方差验证准确率偏低且不稳定with torch.no_grad() 关闭梯度追踪省显存也避免意外计算。checkpoint 里除了模型权重还存了 optimizer_state_dict这是断点续训的关键训练中断时不用从头再来。观察运行日志时如果 loss 持续下降但 val_acc 连续 5 轮不动不用硬等满 30 轮直接手动中断换下一组参数。迁移学习模式下 20 到 30 轮基本稳定从头训练则要 50 轮以上时间成本差一个量级。5.3 微调与学习率调度解冻最后一层参数后的两个选择冻结训练收敛后如果想追求更高准确率常见做法是解冻 ResNet18 最后几个残差块用很小的学习率做全模型微调。这一步成功与否取决于两个参数的选择哪些层参与更新更新步长多大。for name, param in model.named_parameters(): param.requires_grad True # 全部解冻 optimizer torch.optim.Adam([ {params: model.fc.parameters(), lr: 1e-4}, {params: [p for n, p in model.named_parameters() if not n.startswith(fc)], lr: 1e-5}, ], weight_decay1e-4)这里把不同层的学习率分开设置分类头用 1e-4因为它是随机初始化的需要相对快的更新速度预训练层用 1e-5因为它们的权重已经接近局部最优学习率太大会破坏学到的通用特征。微调阶段最常见的问题是过拟合解冻后可训练参数暴增数据量撑不起来就会在验证集上反弹。如果微调 3 个 epoch 后验证集准确率反而下降回退到只训练分类头的状态或者只解冻最后一个残差块再试。学习率调度上我倾向于用 CosineAnnealingLR相比 StepLR 它能平滑地在训练后期把学习率降下来减少在最优解附近来回震荡的情况。调度器接在优化器之后每个 epoch 结束调用 scheduler.step() 即可。6. 从训练结果反推数据质量混淆矩阵与错误样本分析技巧训练结束只是开始真正值钱的是从验证集错误里反推数据集质量。我用一小段脚本画出混淆矩阵直接找出错误率最高的类别对。准确率只是一个数字混淆矩阵是那个告诉你哪里出错的地图。import numpy as np from sklearn.metrics import confusion_matrix import torch all_preds [] all_labels [] model.eval() with torch.no_grad(): for images, labels in val_loader: 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) for i in range(len(train_ds.classes)): for j in range(len(train_ds.classes)): if i ! j and cm[i][j] 5: print(f{train_ds.classes[i]} 被误判为 {train_ds.classes[j]}: {cm[i][j]} 次)看到“青苹果被误判为绿葡萄”这种组合时不要急着调模型回到数据集里翻这两个类别的原图。常见原因有三类标注时手滑放错文件夹两个品种外观上确实接近比如黄苹果和黄桃在顺光下难分辨图片里水果只占很小区域背景主导了模型判断。第三种情况最隐蔽因为数据增强里的 RandomResizedCrop 已经把图裁过但裁剪范围的随机性可能把边缘的背景也当成特征学进去。我的个人习惯是每次训练结束把 top-20 错误样本连同预测值和真实标签拼成一张图直接面对错误。有一次在识别香蕉这个类别时错误样本里超过一半是弯曲的黄色物体翻数据集才发现混入了没有剥皮的黄皮南瓜图片。这种错误靠调参永远发现不了只能人和图对视。血泪经验总结成一句话模型准确率不够时第一嫌疑人是数据而不是网络结构。把这个流程固化下来每次拿到新的水果分类数据都按解压、扫描、划分、训练、看错误样本的顺序走一遍踩坑概率会低很多。希望这套流程帮你在自己的数据集上少走弯路。本文还有配套的精品资源点击获取
📝

华诺云谱内容团队

资深建站顾问 · 行业研究员

10年+企业数字化服务经验,专注智能建站、SEO优化与品牌营销,持续输出建站技巧、行业洞察与营销干货,已帮助5000+企业实现数字化增长。

你可能需要的服务

订阅华诺云谱资讯周报

每周一封,精选建站技巧、SEO与营销干货,直达邮箱。已有 8,000+ 企业主订阅,助你少走弯路。

↑