蘑菇图像识别分类实战:PyTorch与YOLOv5全流程解析
发布时间:2026/9/28 14:32:12来源:尧图网络
简介一份面向实际项目需求的蘑菇图像分类数据集覆盖姬松茸、阿曼妮塔、牛肝菌、Cortinarius等多个常见菌类目标用户是想训练YOLOv5分类模型或自建CNN图像分类网络的开发者与研究者。数据已经按照训练集和测试集分文件夹存放并附带JSON类别字典文件清楚标注每个类别的名称下载后可直接套用到分类模型流程中。压缩包采用7z格式共包含2000个文件其中1998张JPG图片、1个Python可视化脚本、1个JSON字典整体大小约97.67MB目录结构清晰明了。配套的Python脚本可一键批量展示样本图像方便快速核查各类别图片数量与质量省去手动整理标签和路径的麻烦。目前已有1046人学习使用对需要获取分类基准数据、快速开展训练试验或进行迁移学习的入门和进阶读者都非常适合。1. 12种蘑菇图像识别数据集下载前先搞清三件事标题写着12种蘑菇解压后打开类别字典 json你大概率只数得到 8 个类。别急着退货这个 12 是原始采集批次的命名口径真正喂给图像分类网络的是 8 个类别训练集 9600 张、测试集 2400 张目录按「train/test 每类一个文件夹」组织类别字典文件 json 也一并给好了。做图像分类、森林图像识别或者深度学习图像识别相关实验的从业者最烦的不是没数据而是拿到原始图后还要自己洗数据、划训练测试集、写标签文件一下午就耗在预处理上。这份资源的价值就在这里目录已经划分好、标签已对齐直接用 torchvision 的 ImageFolder 或 yolov5 classify 都能吃。它适合正在练手分类网络的初学者也适合需要快速跑一轮分类基线的工程团队。下载量不大但你省下的是半天到一天的脏活。2. 拆目录结构train/test 划分、类别字典 json 与一次跑通的可视化脚本2.1 data 目录下到底长什么样先看文件组织再谈训练拿到资源后的第一件事不是急着写训练脚本而是先完整把目录结构列一遍。常见做法是直接用tree /fWindows或find命令看一遍find data -maxdepth 2 -type d | sort正常你会看到类似下面的组织结构data/ ├── train/ │ ├── Agaricus/ # 姬松茸 │ ├── Amanita/ # 阿曼妮塔 │ ├── Boletus/ # 牛肝菌 │ ├── Cortinarius/ │ └── ... # 其余类别以 json 实际内容为准 ├── test/ │ ├── Agaricus/ │ ├── Amanita/ │ └── ... └── classes.json # 类别字典文件这份资源的 train 和 test 是按「每类一个文件夹」组织的不是那种影像分类里常见的images/labels.txt平铺结构。前者对 PyTorch 的torchvision.datasets.ImageFolder和 yolov5 的 classify 任务都是开箱即用的不需要再写额外的路径解析逻辑。这一点是选型时最省事的地方。打开classes.json里面是一个类别名到编号的映射。以摘要里提到的类别为例能看到姬松茸Agaricus、阿曼妮塔Amanita、牛肝菌Boletus、Cortinarius丝膜菌属等实际完整类别清单以你下载到的 json 为准。也可以用一段非常短的 python 把类名和数量打出来import json with open(data/classes.json, r, encodingutf-8) as f: class_idx json.load(f) print(f类别数量: {len(class_idx)}) for name, idx in class_idx.items(): print(f {idx}: {name})这里encodingutf-8是我习惯性加上的很多 json 文件在 Windows 下用默认编码读会直接抛UnicodeDecodeError加一个显式编码参数能避开一半的路径类事故。看到类别数量是 8并且打印出来的类名里没有乱码说明数据源是完整的可以进入下一步。2.2 用 show 脚本做数据集可视化9 宫格抽样看真实图像质量资源里带了一个 show 脚本用于可视化数据集。它的作用就是随机从每个类里抽图拼成网格让你一眼看出这批蘑菇图像的质量有没有黑边、有没有水印、有没有模糊到没法用的图。如果你下载的脚本能直接跑那就直接跑如果环境里缺依赖跑不起来下面这个等效脚本我用得很频繁直接抄走就行import json import glob import random import matplotlib.pyplot as plt from PIL import Image # 读取类别字典 with open(data/classes.json, r, encodingutf-8) as f: class_idx json.load(f) classes list(class_idx.keys()) # 3x3 网格每次随机抽 9 张 rows, cols 3, 3 fig, axes plt.subplots(rows, cols, figsize(9, 9)) for i in range(rows * cols): c random.choice(classes) # 注意有些图可能是 .png 或 .jpeg多匹配几种后缀更稳 candidates glob.glob(fdata/train/{c}/*.jpg) \ glob.glob(fdata/train/{c}/*.jpeg) \ glob.glob(fdata/train/{c}/*.png) img_path random.choice(candidates) ax axes[i // cols][i % cols] ax.imshow(Image.open(img_path)) ax.set_title(c, fontsize10) ax.axis(off) plt.tight_layout() plt.show()逻辑说明先从classes.json读类别清单然后用glob按类在 train 目录下找图片后缀同时兼容 jpg、jpeg、png 三种常见格式。random.choice负责随机抽样所以每次运行看到的 9 张图都不一样适合快速把整个数据集的图像质量摸个大概。参数上figsize(9, 9)控制整体画布大小3x3 网格下刚好fontsize10是标题字号。如果你想把全部 8 类都看一遍就把rows, cols改成 2x4或者循环里固定c classes[i % len(classes)]而不是随机抽类。看图的重点就三个图像是否带多余水印、是否有多张子图拼在一张里的原始采集图、是否存在明显失焦的图。看到问题图后建议在训练前手动清理而不是指望网络自己学出来。3. 用 PyTorch 训练分类基线ImageFolder 加载、验证集划分与超参表3.1 用 ImageFolder 加载数据并先切出验证集torchvision 的ImageFolder是这类「文件夹即标签」数据集最省事的加载方式。它会自动按文件夹名的字母序生成类别索引并把每张图解析成(image, label)对。这里有一个值得注意的点classes.json里的编号顺序和ImageFolder按字母序生成的编号不一定一致所以加载后先打印dataset.classes做一次对齐别急着训练。另外这份数据只给了 train 和 test没有 val。很多新手会拿 test 边训练边验证这是个大坑后面避坑章会细说。正确做法是从 train 里固定切 10% 出来当验证集import torch from torch.utils.data import DataLoader, random_split from torchvision import datasets, transforms # 数据增强Resize 到 256再随机裁剪 224符合 ImageNet 预训练模型输入 train_transform transforms.Compose([ transforms.Resize(256), transforms.RandomResizedCrop(224), transforms.RandomHorizontalFlip(), transforms.ColorJitter(0.3, 0.3, 0.3), transforms.ToTensor(), transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]), ]) full_train datasets.ImageFolder(data/train, transformtrain_transform) print(full_train.classes) # 确认类别顺序 total len(full_train) n_val int(total * 0.1) train_ds, val_ds random_split(full_train, [total - n_val, n_val]) train_loader DataLoader(train_ds, batch_size32, shuffleTrue, num_workers4) val_loader DataLoader(val_ds, batch_size64, shuffleFalse, num_workers4)逻辑说明ImageFolder的classes属性就是它实际使用的类别顺序打印出来跟classes.json对比能提前暴露标签错位。random_split按固定比例把训练集切成 90% 训练 10% 验证验证集只用来看收敛情况不参与权重更新。Normalize用的是 ImageNet 的均值和标准差因为后面加载的是在 ImageNet 上预训练过的 ResNet输入分布必须对齐。参数说明Resize(256)RandomResizedCrop(224)是分类任务最常见的尺度策略给裁剪留了随机空间相当于做了一次尺度扰动。ColorJitter(0.3, 0.3, 0.3)对亮度、对比度、饱和度各做 0.3 幅度的随机扰动对蘑菇这类颜色敏感的类别不建议把幅度调得更大否则会破坏菌盖颜色的判别信息。batch_size32在 ResNet34 8GB 显存下比较稳num_workers4是 Linux 下的常用值Windows 下如果报 DataLoader worker 相关的错把它改成 0 是最快的解决办法。3.2 训练脚本与超参设置ResNet34 在 9600 张上的基线预处理做完就该上模型了。蘑菇图像类间差异小、类内差异大菌盖形状和颜色是主要判别特征所以我一般先用 ResNet34 而不是更大更深的模型跑基线参数量适中在 9600 张的训练规模下不容易过拟合训练一轮的时间也够短方便反复试错。import torch import torch.nn as nn from torchvision import models # 类别数以 json 为准而不是以标题为准 num_classes len(class_idx) model models.resnet34(weightsmodels.ResNet34_Weights.IMAGENET1K_V1) model.fc nn.Linear(model.fc.in_features, num_classes) device torch.device(cuda if torch.cuda.is_available() else cpu) model model.to(device) criterion nn.CrossEntropyLoss() optimizer torch.optim.Adam(model.parameters(), lr1e-3) scheduler torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max30) for epoch in range(30): 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) # 每个 epoch 结束在验证集上评估一次 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) _, predicted torch.max(outputs, 1) total labels.size(0) correct (predicted labels).sum().item() acc 100.0 * correct / total print(fEpoch {epoch1:02d} | Loss: {running_loss/len(train_ds):.4f} | Val Acc: {acc:.2f}%) scheduler.step()逻辑说明加载在 ImageNet 上预训练过的 ResNet34只把最后一层全连接换成 8 类输出这是迁移学习里标准的 finetune 做法。预训练权重提供了丰富的底层特征蘑菇图像虽然在 ImageNet 里不多但菌盖的纹理、边缘、颜色分布依然能复用底层卷积特征所以收敛速度远快于从零训练。参数说明优化器用 Adam 而非 SGD图的是前期收敛快lr1e-3对 finetune 来说比较温和CosineAnnealingLR把学习率在 30 个 epoch 内按余弦曲线降到接近 0后半程相当于做精细微调。CrossEntropyLoss自带 Softmax所以模型输出层不需要额外加激活。如果你的显存只有 4GB把batch_size降到 16学习率同步降到 5e-4效果差别不大。这里给一张我跑这类蘑菇分类任务常用的超参表直接照着填问题不大参数推荐值说明输入分辨率224x224兼顾细节与显存Batch Size328GB 显存可跑Epochs30看验证集是否 plateau优化器Adam / SGD(momentum0.9)Adam 快SGD 稳初始学习率1e-3finetune从零训练则用 1e-2学习率调度CosineAnnealing后半程微调更细腻权重初始化ImageNet 预训练强烈建议别从零训跑完 30 个 epoch验证集准确率一般能到 90% 上下。如果没到先别急着换模型回到 2.2 节的可视化脚本看看是不是有脏图或者去读一下避坑章里相似类的问题。4. 接到 yolov5 / yolov8 分类从 train 切 val 到一条命令跑起来4.1 yolov5 classify 的数据组织与类别映射yolov5 自带的分类任务classify/train.py和 yolov8 的yolov8n-cls.pt都遵循同一个约定数据集根目录下必须有两个子目录train/和val/每个子目录内按「每类一个文件夹」组织。注意它找的是val而不是test。这份资源给的是train/和test/所以不能直接开训要先做一步目录对齐。最省事但不推荐的做法是把test/改名为val/直接训练。问题是这样最终评估就只能用同一个 val指标虚高。我习惯的做法是从train/里每类抽 10% 出来独立成val/把原始test/留作训练结束后的最终盲测。import os import glob import random import shutil src data/train dst_val data/val_for_yolo os.makedirs(dst_val, exist_okTrue) val_ratio 0.1 random.seed(42) for class_dir in os.listdir(src): class_path os.path.join(src, class_dir) if not os.path.isdir(class_path): continue images glob.glob(os.path.join(class_path, *.jpg)) \ glob.glob(os.path.join(class_path, *.jpeg)) random.shuffle(images) n_val max(1, int(len(images) * val_ratio)) os.makedirs(os.path.join(dst_val, class_dir), exist_okTrue) for img in images[:n_val]: shutil.move(img, os.path.join(dst_val, class_dir, os.path.basename(img)))逻辑说明对 train 下每个类别目录先 glob 出所有图片按固定比例抽 10%用shutil.move从 train 挪到新的 val 目录。random.seed(42)保证每次运行切分结果一致这样不同轮次的实验对比才公平。切分后 train 少了 10% 的图总量从 9600 变成约 8640对训练影响不大。如果你用的是 yolov8 的ultralytics包数据组织方式完全一样也是根目录下train/和val/两个子目录类别由文件夹名自动推断。所以做完这步切分两份框架都能直接吃这份数据。4.2 训练命令与本地推理yolov5 / yolov8 训练自己的蘑菇分类模型目录切好后训练命令很短。yolov5 的 classify 分支是这样跑的# 进入 yolov5 仓库根目录 python classify/train.py \ --model resnet34 \ --data /path/to/mushroom_dataset \ --epochs 50 \ --img 224 \ --batch 32 \ --device 0--model指定 backbone除了 resnet34 还可以换resnet18跑更快、efficientnet_b0跑更省内存--data指向包含train/和val/的根目录--img 224和--batch 32跟 PyTorch 那套一致。跑完后权重存在runs/train-cls/exp/weights/best.pt。yolov8 的写法更简洁用的是命令行直接指定数据集路径yolo classify train \ modelyolov8n-cls.pt \ data/path/to/mushroom_dataset \ epochs50 \ imgsz224 \ batch32 \ device0modelyolov8n-cls.pt会从官方源自动下载预训练分类权重data指向同样的根目录。yolov8 的 n 模型非常轻显存占用不到 2GBCPU 也能跑但精度上限比 resnet34 低一些适合先验证流程通不通。推理也简单yolov5 用classify/predict.pypython classify/predict.py \ --weights runs/train-cls/exp/weights/best.pt \ --source /path/to/test/Amanita/xxx.jpg输出会打印这张图属于每个类别的置信度。批量验证整个 test 目录也可以把--source直接指向data/test它会递归遍历所有子目录并输出每张图的预测结果。yolov8 则是yolo classify predict modelruns/train-cls/exp/weights/best.pt source/path/to/test到这里你已经拿到了一个能跑通全流程的分类模型。剩下的问题不是模型架构行不行而是有没有踩中数据本身的暗坑。5. 避坑8 类还是 12 类、标签错位与类别不均衡5.1 类别数量标题写 12json 里只有 8现象资源标题写着「12种蘑菇图像识别数据集」但classes.json打开只有 8 个类别训练脚本里num_classes填几都要犹豫半天。原因12 是原始采集图片的批次口径。项目正文的文件名里能看到 Lactarius、Pluteus、Entoloma、Cortinarius 等多个属名采集阶段可能按 12 个批次或 12 个来源归档但整理成发布版时合并成了 8 个可判别类别。摘要里明确写了「分类个数8」所以 12 和 8 并不矛盾只是统计口径不同。解决一切以classes.json为准。写训练脚本前先用len(class_idx)确认类别数再回去核对ImageFolder.classes的长度是否一致。任何地方出现num_classes 12都是错的跑完训练再去排查标签错位就晚了。5.2 标签顺序ImageFolder 的字母序和 json 编号对不上现象训练时 loss 正常下降但验证集的 top-1 准确率始终在 30% 上下震荡怎么看都像随机猜。原因ImageFolder按文件夹名的字母序自动生成标签比如Agaricus排 0、Amanita排 1、Boletus排 2而classes.json里的编号可能是按采集顺序或人工指定顺序写的。两份顺序一旦不一致加载进来的 label 含义就错位了模型学到的映射关系全是乱的。解决训练前做一个硬性校验不通过就不开跑import json from torchvision import datasets with open(data/classes.json, r, encodingutf-8) as f: class_idx json.load(f) dataset datasets.ImageFolder(data/train) # 方式一直接对比文件夹顺序和 json 顺序 print(ImageFolder 顺序:, dataset.classes) print(json 顺序:, list(class_idx.keys())) # 方式二断言两者必须一致不一致就抛异常 assert list(dataset.classes) list(class_idx.keys()), \ 类别顺序不一致需要统一后重新生成 json如果断言失败解决办法是把classes.json重新按dataset.classes的顺序生成一遍而不是去改文件夹名。文件夹名是给训练框架看的json 只是给人看的辅助文件以文件夹名为准重建 json两边就对齐了。5.3 类别不均衡9600 张摊到 8 类未必均匀现象训练集总共 9600 张但训练时发现某些类别的 loss 一直下不去混淆矩阵里个别类 recall 特别低。原因9600 是总数不代表每类正好 1200 张。蘑菇采集本身受季节、地域影响很大某些常见种可能占了 3000 张稀有类别可能只有 500 张。类别不均衡会让模型偏向样本多的类。解决训练前先做一个快速统计用numpy数一下每类样本数import numpy as np from torch.utils.data import Dataset # 用 ImageFolder 的 targets 属性直接统计 counts np.bincount(dataset.targets) for cls, count in zip(dataset.classes, counts): print(f{cls}: {count} 张)如果发现差距超过 2 倍就用WeightedRandomSampler做样本加权采样给少样本的类更高的采样概率from torch.utils.data import WeightedRandomSampler class_counts np.bincount(dataset.targets) weights 1.0 / class_counts[dataset.targets] sampler WeightedRandomSampler(weights, num_sampleslen(dataset), replacementTrue) train_loader DataLoader(train_ds, batch_size32, samplersampler)采样器会让每轮 epoch 里少样本类别被抽到的次数显著增加缓解偏向。代价是每个 epoch 实际上看到重复样本训练轮数不用加太多30 轮以内足够。5.4 没有验证集反复用 test 调参会虚高现象训练时一直拿data/test当验证集看准确率调了几轮超参后测试集准确率到了 95%但换成真实场景的新图准确率掉到 80%完全没法解释。原因test 集被反复用于超参选择和 early stopping它的信息已经在训练过程中泄漏给模型了。测试集的有效性在于「只用一次」反复用它调参它就不再是独立评估集结果虚高是必然的。这是分类实验里最常见的翻车点。解决严格三集分离。train 训练、val从 train 里切 10%调参、test 只做最终评估。第 4 章给的切分脚本就是专门干这个的训完所有实验后用test/跑最后一次推理这个数字才是能写进报告的结果。5.5 相似类难分Amanita 和 Entoloma 分不清是正常的现象混淆矩阵里Amanita和Entoloma两个类互认错训练初期 loss 下降比其他类慢很多换大模型效果也没明显改善。原因这两个属的蘑菇都是典型伞菌形态菌盖颜色、菌褶走向在视觉上非常接近人眼都容易认错。类间相似度高是蘑菇分类数据集的固有难点不是模型或者代码的 bug。文件名里出现的Entoloma_original、Lactarius_original这类字样也能看出来原始图里大量是野外同一生长环境拍的背景干扰大。解决先接受基线结果再针对性地做三件事。一是输入分辨率从 224 提到 320让模型看到更多菌盖细节二是对易混类做定向数据增强比如更强的随机旋转和光照扰动三是在验证时看 top-2 准确率对这类相似类任务top-2 比 top-1 更有实际参考价值。别指望把这两个类的区分做到 100%能做到 90% 已经超过多数人工标注水平。6. 进阶用混淆矩阵和 Top-2 准确率给模型做一次体检训练完的模型光看一个总准确率是不够的。蘑菇分类这种类间相似度高的任务真正能说明问题的是混淆矩阵哪两个类容易被搞混、每个类的召回率是多少都在这一张图里。下面这段代码用 sklearn 生成混淆矩阵和每类的 precision / recall是我每次跑完分类任务都会执行的固定动作import torch import numpy as np from sklearn.metrics import confusion_matrix, classification_report import seaborn as sns import matplotlib.pyplot as plt model.eval() y_true, y_pred [], [] with torch.no_grad(): for images, labels in val_loader: images images.to(device) outputs model(images) _, predicted torch.max(outputs, 1) y_true labels.tolist() y_pred predicted.cpu().tolist() # 每类的 precision / recall / f1 print(classification_report(y_true, y_pred, target_namesval_ds.dataset.classes)) # 混淆矩阵可视化 cm confusion_matrix(y_true, y_pred) plt.figure(figsize(8, 8)) sns.heatmap(cm, annotTrue, fmtd, cmapBlues, xticklabelsval_ds.dataset.classes, yticklabelsval_ds.dataset.classes) plt.xlabel(Predicted) plt.ylabel(True) plt.show()classification_report输出每个类别的 precision、recall、f1-score哪个类是短板一眼就能看出来热力图上对角线以外的亮点就是模型最容易翻车的类别对。看到Amanita和Entoloma互串再去针对性调数据增强比盲试网络结构有效得多。Top-2 准确率在蘑菇分类里比 top-1 更贴近实际使用场景。eDNA 调查或者野外采集辅助识别这类应用通常给出「最可能的两种蘑菇」让用户确认比硬要模型赌一个答案更合理。统计 top-2 的代码也不复杂correct_top2 0 total 0 with torch.no_grad(): for images, labels in val_loader: images images.to(device) outputs model(images) top2 torch.topk(outputs, 2, dim1).indices.cpu() # 取前两个预测 for i, label in enumerate(labels): if label.item() in top2[i].tolist(): correct_top2 1 total 1 print(fTop-2 Accuracy: {100.0 * correct_top2 / total:.2f}%)这段代码的要点是torch.topk(outputs, 2, dim1)取的是每个样本 logits 最高的两个类别索引只要真实标签落在里面就算对。在相似类多的数据上top-2 一般会比 top-1 高 5 到 10 个百分点这个差距本身就是数据难度的直观体现。从那以后我拿到任何分类数据集第一件事永远是三件套数 json 长度、跑 bincount、打印 9 宫格图这三步走完后面基本不会翻车。希望帮到你。本文还有配套的精品资源点击获取
网站建设高端定制企业官网