ARTICLE DETAIL

资讯详情

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

Python蘑菇识别系统:从数据清洗到迁移学习部署的完整实践

Python蘑菇识别系统:从数据清洗到迁移学习部署的完整实践 简介一份可运行的Python蘑菇识别系统源码包面向图像识别初学者、生物爱好者及机器学习开发者涵盖蘑菇图片预处理、特征提取、分类模型训练、模型评估和界面实时预测的完整开发流程。压缩包共54个文件、30.98MB其中包括9个py源文件、20个pyc编译文件、23张png界面与素材图片另附readme文档与说明文本各模块按mogu.py的模型调用、gui_util.py的桌面交互、utils的工具函数划分便于逐段研读和二次改造。系统以“好菇毒”蘑菇识别为主题内置图片选择、按钮交互、识别结果展示等完整桌面应用逻辑读者既能直接运行体验完整识别流程也可自行准备数据集迁移训练。已有382人学习下载适合希望快速掌握Python图像分类应用、理解机器学习项目工程化组织方式并在此基础上做功能扩展的开发者。1. 蘑菇识别不是拍脑袋Python图像分类系统到底在解决什么山路边拍到一朵蘑菇想知道能不能吃这不该靠老话“颜色鲜艳有毒长得朴素可食”去赌。Python蘑菇识别系统做的事就是把“这朵蘑菇长得像什么”变成一个四分类问题可食、有毒、无法判断、非蘑菇并给每个结论附上置信度。源码以zip压缩包分发解压后能看到数据清洗、训练、评估、推理脚本跑通一条完整链路。这个项目真正难的不是模型——迁移学习随便拿个预训练网络都能跑出像样的准确率难的是数据分布、标签可信度、阈值设计这三件事。它们决定模型在陌生场景下敢不敢说“不确定”。适合想从通用图像分类走向垂类识别的Python开发者也适合植物科普类小工具的原型验证。2. 蘑菇数据集从哪来样本采集、按毒性与可食性分桶、数据清洗这一步决定模型上限蘑菇识别的上限由数据决定。模型结构只能决定把上限逼近多少。先别急着写训练代码把数据这一关过完后面几乎不会翻车。2.1 类别体系先定好按“可食/有毒/无法判断”分桶而不是按菌种名硬分不少第一次做蘑菇识别的人会想把类别设成“松茸、牛肝菌、红菇……”几十个类。这个愿望很美好但落地会撞上三堵墙相近菌种外观高度相似细分类别之间边界模糊单类样本数量肯定不够几十个类摊下去每类可能只有几十张图最致命的是细分类一旦错分用户按“可食”吃了后果不是一句“模型不准”能兜住的。我一般把类别体系压缩成四桶可食edible常见且无毒、能放心吃的类群。有毒poisonous文献明确标记有神经毒性或胃肠毒性的类群。无法判断uncertain标签来源不可靠、形态介于两者之间、样本质量差。非蘑菇non_mushroom提供拒识出口后面会说明它的必要性。为什么要留“uncertain”这一类因为训练集里总有标签存疑的图硬归到可食或毒只会污染模型不如让模型专门学一个“我拿不准”的分桶。推理端它会被翻译成“无法判断请找专业人士”比瞎给结论安全得多。数据目录按类别建文件夹对应一份 label_map.json{ 0: edible, 1: poisonous, 2: uncertain, 3: non_mushroom }这份 JSON 是训练脚本和推理脚本共用的“字典”顺序一旦定下来就不要改。最容易踩的坑是训练时按字母序自动分配标签推理时又手写了一份不同顺序的 ID2LABEL两个脚本对不上跑出来全错还不报错。目录建议长这样data/ train/ edible/xxx1.jpg poisonous/xxx2.jpg uncertain/xxx3.jpg non_mushroom/xxx4.jpg val/ test/2.2 采集与清洗去重、裁掉水印、剔除模糊图图片来源一般有三个公共数据集、网络图片、自己实拍。实拍最可靠但数量有限网络图片量大却要人工过滤水印、拼图、带科普文字标注的截图。水印和文字会让模型去学“文字区域”而不是菌盖特征这类图前期可以直接删掉或者裁掉边缘。图片收完先做一遍基础清洗。检查损坏文件和过小图片的脚本import os from PIL import Image data_root data/train min_side 128 def check_and_remove(path: str) - None: try: with Image.open(path) as img: img.verify() # 只校验文件头不完整解码 w, h img.size if min(w, h) min_side: os.remove(path) print(f过小图片已移除: {path}) except Exception as exc: os.remove(path) print(f损坏图片已移除: {path}: {exc}) for root, _, files in os.walk(data_root): for f in files: if f.lower().endswith((.jpg, .png, .jpeg)): check_and_remove(os.path.join(root, f))这里有两个容易忽略的点。img.verify()之后不能再继续load()否则会报 os error所以只用来检查文件是否完好的场景足够。min_side 128是我常用的下限阈值小于这个尺寸的图菌盖边缘和菌褶纹理已经糊掉了模型学到的是噪声删掉比留着增强更划算。如果数据总量太少可以把阈值放宽到 64同时做好心理预期这部分样本会把验证集指标拉低一点。网络下载图还有一个高发问题同一张图被不同网站转载、改尺寸、加滤镜最后看起来是几十个文件实际内容几乎一样。重复样本进训练集等于变相给这批数据加权重会放大它们带来的偏差。用感知哈希去重import hashlib import os from PIL import Image def dhash(img, hash_size16): 基于明暗变化的感知哈希裁剪/压缩后依然相近 img img.convert(L).resize((hash_size 1, hash_size)) px list(img.getdata()) return sum((px[i] px[i 1]) (i % 64) for i in range(hash_size * hash_size)) hash_map {} for root, _, files in os.walk(data_root): for f in files: if not f.lower().endswith((.jpg, .png, .jpeg)): continue path os.path.join(root, f) try: h dhash(Image.open(path)) except Exception: continue if h in hash_map: print(f疑似重复移除: {path}) os.remove(path) else: hash_map[h] pathdhash 比直接的 md5 文件哈希更实用同一张图改个尺寸、调个亮度md5 会变而 dhash 记录的是像素明暗变化关系剪裁之后依然一致。注意这里用 dict 做精确哈希去重只删完全重复的如果要处理“同一场景连拍但视角略不同”的近重复就得计算两两汉明距离几千张图片时开销不大可以按需扩展。2.3 划分训练/验证/测试按来源批次切分比随机切分更真实看到训练准确率 98%、验证准确率 96%先别高兴。如果数据是按“图片文件夹随机切分”做的这个数字很可能是虚高的。原因很现实同一朵蘑菇的多个角度照片某一张进了训练集另一张几乎同场景的图进了验证集模型等于提前见过答案。真实场景里用户拍到的照片和训练集完全不同验证时被这类“重复记忆”掩盖的问题全部会暴露。解决办法是按采集批次或来源文件夹整体切分。如果你的数据是用“source_a”、“source_b”这样的批次目录收集的切分时以批次为单位而不是以图片为单位import os import random from collections import defaultdict all_images [] for root, _, files in os.walk(data): for f in files: if f.lower().endswith((.jpg, .png, .jpeg)): all_images.append(os.path.join(root, f)) source_groups defaultdict(list) for img_path in all_images: parts img_path.split(os.sep) # 假设路径是 source/类别/图片.jpg source parts[1] source_groups[source].append(img_path) sources list(source_groups.keys()) random.shuffle(sources) n len(sources) train_sources sources[: int(n * 0.7)] val_sources sources[int(n * 0.7): int(n * 0.85)] test_sources sources[int(n * 0.85):]用这个列表把对应图片复制到训练/验证/测试目录即可。70/15/15 是我在数据批次较多时的默认划分如果只有三四个批次建议用 40/30/30否则验证集会太小指标波动大。关键点测试集必须是一次都没有参与过训练、也没有参与过验证集调参的“新批次”。这套切分逻辑不会让你的准确率显得那么漂亮但它测出来的是模型真实泛化能力。3. 用迁移学习训练识别模型EfficientNet 底座与数据增强参数怎么设数据准备好之后模型训练反而是最不玄学的部分。垂类图像识别选迁移学习几乎没有悬念真正需要花时间的是增强参数和学习率策略。3.1 为什么选迁移学习而不是自己从头搭 CNN蘑菇识别是一个典型的“小数据 细粒度”任务。一个类别可能只有两三百张图自己从零搭一个 ResNet 去训练网络第一层连“什么是边缘、什么是纹理”都要从数据里学样本量完全不够。迁移学习则把 ImageNet 上预训练好的通用视觉特征搬到你的模型里它已经知道菌盖边缘是边缘、菌褶是重复纹理你要做的只是在最后几层重新组织这些特征让模型把它们组合成“可食/有毒”的判断。这也是为什么迁移学习在几千张图的数据量下就能达到可用的准确率而从头训练可能要几万张。底座怎么选我常用的三档EfficientNet-B0精度与速度平衡最好CPU 推理也还能接受。EfficientNet-B1比 B0 精度高一点参数量大三分之一适合有 GPU 且推理机配置不差的情况。ResNet18推理速度最快适合部署到低配迷你主机或树莓派。我这里以 EfficientNet-B0 为例。换底座只改两行代码后面讲的增强与训练策略完全通用。3.2 训练脚本数据增强、类别权重、学习率三个必调参数这段是我训练垂类识别模型的基准配置复制改路径就能跑import torch import torch.nn as nn from torch.utils.data import DataLoader, WeightedRandomSampler from torchvision import datasets, models, transforms num_classes 4 batch_size 16 # 训练增强重点模拟自然环境中蘑菇的真实状态 train_tf transforms.Compose([ transforms.RandomResizedCrop(224, scale(0.7, 1.0)), transforms.RandomHorizontalFlip(), transforms.RandomVerticalFlip(p0.1), transforms.ColorJitter(brightness0.3, contrast0.3, saturation0.2), transforms.RandomErasing(p0.25, scale(0.02, 0.15)), transforms.ToTensor(), transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]), ]) val_tf transforms.Compose([ transforms.Resize(256), transforms.CenterCrop(224), transforms.ToTensor(), transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]), ]) # 计算每个类别样本数的倒数作为权重 labels [] # 从数据集读出的每张图真实标签 class_counts torch.bincount(torch.tensor(labels)).float() class_weights 1.0 / class_counts class_weights class_weights / class_weights.sum() sample_weights class_weights[torch.tensor(labels)] sampler WeightedRandomSampler(sample_weights, num_sampleslen(labels), replacementTrue) model models.efficientnet_b0(weightsmodels.EfficientNet_B0_Weights.IMAGENET1K_V1) model.classifier[1] nn.Linear(model.classifier[1].in_features, num_classes)三个参数值得单独说明。第一个是RandomResizedCrop的 scale。我刻意没有用默认的(0.08, 1.0)而是收紧到(0.7, 1.0)。原因是蘑菇识别属于细粒度任务菌盖上的鳞片、菌褶的疏密是关键判别点裁剪比例太小会把判别性细节裁掉。训练时随机裁掉边上一圈正好模拟拍照时蘑菇没拍全的状态。第二个是RandomErasing。野外蘑菇经常被松针、树叶遮挡擦除增强模拟的就是“菌盖被盖住了一半”的情形。p0.25相当于每四张图擦一张擦除面积控制在 2% 到 15%——超过 15% 会把整个菌盖擦掉模型只能靠背景猜效果反而变差。第三个是WeightedRandomSampler。蘑菇数据类别天然不均衡可食类因为样本好找可能占了一半有毒类数据难收可能只有少量。权重按样本数倒数计算让每次采样少量类别的概率更大。这里没有直接把权重乘到 loss 上是因为在几千张图的小数据里loss 加权容易放大噪声样本而采样加权更平滑、更稳。训练轮数少于 15 的话把RandomErasing的 p 降到 0.15否则增强太强模型学不完真实特征。训练循环的优化器配置optimizer torch.optim.AdamW(model.parameters(), lr1e-4, weight_decay1e-4) scheduler torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max20, eta_min1e-6) criterion nn.CrossEntropyLoss() for epoch in range(20): model.train() for images, labels in train_loader: optimizer.zero_grad() outputs model(images) loss criterion(outputs, labels) loss.backward() optimizer.step() scheduler.step()学习率用 1e-4不做“冻结骨干先训头再解冻”的策略。原因和增强一样数据太少冻结骨干那一两个 epoch 很容易让优化器一步迈到奇怪的位置。直接让预训练骨干跟着小学习率一起微调在 20 个 epoch 内更稳。AdamW 的 weight_decay 取 1e-4这个值对 ImageNet 预训练模型在小数据集上微调是常用区间太大可能压掉骨干已学好的特征太小则起不到抑制过拟合的作用。CosineAnnealing 搭配 T_max20 让学习率平滑降到 1e-6比按 epoch 手动降级省心。3.3 评估指标怎么看准确率会骗人要看混淆矩阵和各类别召回训练完成后的第一反应可能是看准确率95%看起来不错。但如果可食类样本占了一半有毒类只占 10%模型把有毒类全部判错准确率也只会掉 5 个百分点肉眼完全看不出问题。所以评估必须以混淆矩阵为中心重点看各类别召回率。from sklearn.metrics import classification_report, confusion_matrix, ConfusionMatrixDisplay import matplotlib.pyplot as plt preds, probs, y_true run_eval(model, val_loader) report classification_report( y_true, preds, target_names[edible, poisonous, uncertain, non_mushroom], digits3 ) print(report) cm confusion_matrix(y_true, preds) ConfusionMatrixDisplay(cm).plot() plt.savefig(confusion_matrix.png, dpi150)run_eval就是在验证集上跑一遍推理把预测标签和真实标签收集起来。看报告时我关注的顺序是每个类别的 recall尤其是 poisonous 的 recall。低于 0.8 就意味着相当一部分毒蘑菇被模型放走了这个方向不能接受。edible 被误判成 poisonous 的数量。误伤不致命但太多会影响体验。poisonous 被误判成 edible 的数量。出现一例都要把对应图片捞出来看是标签错标还是两个类群长得确实太像。后者要考虑把该样本移进 uncertain。uncertain 类上的表现。如果 uncertain 的召回高说明模型能表达“拿不准”这比硬给结论更有价值。同时把每个样本的 softmax 概率数组保存下来后面切阈值的时候要用eval 脚本里顺手存一份probs.npy就好。推理接口上线前这些数字必须了然于胸否则你根本不知道该把置信度阈值定在 0.5 还是 0.7。4. 蘑菇识别系统避坑记录数据、训练、推理三处高频翻车点这套系统我从数据到推理踩过不少坑挑四个最典型、最会影响交付质量的写在这里。每一条都是“现象 → 原因 → 解决”三段式可以直接对照排查。4.1 背景不是白底识别率直接掉十几个点现象训练时验证集准确率 95%拿到野外实拍图上测掉到 80% 以下。原因训练图片大多是白底或纯色背景的“图鉴风格”照片模型偷懒学到了“白底 菌盖”这个组合特征。实际场景里草地、松针、树皮背景一出现模型的判断逻辑就乱了。解决从两个方向下手。第一采集阶段刻意混入自然背景实拍图不要统一抠成白底。第二在训练增强里做随机背景破坏让模型无法依赖背景颜色。简单做法是随机把训练图的背景换成偏绿的噪声纹理import random import numpy as np from PIL import Image def random_background(img): if random.random() 0.5: noise np.random.randint(60, 150, size(img.size[1], img.size[0], 3)) bg Image.fromarray(noise.astype(uint8)) img Image.composite(img, bg, img.split()[-1]) return img不用 alpha 通道的话也可以直接在 numpy 层面把图片四周抹绿。核心目标只有一个让“背景 纯白”不再是可靠特征。4.2 标签被错标模型把毒蘑菇当可食现象训练完成后人工复核发现某类毒蘑菇的验证集图片被大量识别成可食单独看这些图有几张本身就是来源网站标错的可食图。原因网络搜集的图片标签可信度参差不齐文件名和同页文字描述都可能带误导。训练时一张错标图的影响会被模型放大尤其在某个类别样本本身就少的情况下。解决训练前对每个类别做全量人工复核不确定的图直接移到 uncertain 类。数据量实在太大就先随机抽 20%如果错标率超过 5%整个类的图片要重新过一遍不要急着训练。同时把疑似但无法确认的图片清单单独存放作为后续推理端的“拒绝输出名单”即使模型给出高置信度的可食结论只要图片和名单里的样本高度相似就返回“无法判断”。这套机制不是模型层面的但能在风险较高的冷门类上兜底。4.3 推理阶段内存暴涨加载方式问题现象用训练好的模型批量识别几百张图片内存从 500MB 一路涨到接近 3GB最后差点 OOM。原因一口气把所有图片读成 PIL 对象再批量送到 GPU 或 CPU 推理或者推理时忘记包torch.no_grad()PyTorch 默认记录梯度把中间激活值全部缓存下来。解决推理必须保持单张、无梯度torch.no_grad() def predict_one(model, image_path: str, tf: transforms.Compose): img Image.open(image_path).convert(RGB) img tf(img).unsqueeze(0) # 形状 [1, 3, 224, 224] logits model(img) prob torch.softmax(logits, dim1).squeeze(0) return probconvert(RGB)不能省有些蘑菇特写图会带 alpha 通道不转为三通道transforms 在Normalize阶段会报通道数不匹配。先convert再走 transform 是最稳的顺序。加no_grad之后推理阶段不再保存中间梯度内存占用基本等于一张图加一个模型本身。4.4 不把石头和树叶当蘑菇阈值与“非蘑菇”拒识现象拿一张树叶照片喂给系统模型硬是返回“可食68%”。原因标准分类模型只能在定义过的类别里选一个即使图片完全不像任何一类它也必须挑一个“最大概率”。模型缺少“全都不像”的出口。解决训练集里加 non_mushroom 类这一类找负样本很容易树叶、石头、松果、昆虫照片都能用。推理端再加一道置信度门槛prob_max, pred_idx torch.max(prob, dim0) if pred_idx ! NON_MUSHROOM_IDX and prob_max.item() 0.55: pred_idx UNCERTAIN_IDX0.55 是我在几千张实拍图上统计出来的经验值。阈值太高会把好多可识别的蘑菇拒之门外太低又会让非蘑菇溜进来。确认自己的阈值时用第 3 章保存的probs.npy对 test set 做一遍二分搜索找一个让“有毒误判为可食”概率小于 1% 的最低点即可。5. 把模型封装成可复用的识别接口命令行与 FastAPI 两种姿势模型训练好只是项目完成了一半另一半是让别人能方便地用起来。源码包通常要同时提供命令行工具和 HTTP 服务前者方便本地验证后者方便接入小程序、Web 前后端。5.1 命令行识别脚本单张图片最快验证命令行是最小可用单元任何环境第一件事都是跑通它。predict.py的骨架import argparse import torch from PIL import Image from torchvision import models, transforms ID2LABEL { 0: edible, 1: poisonous, 2: uncertain, 3: non_mushroom, } val_tf transforms.Compose([ transforms.Resize(256), transforms.CenterCrop(224), transforms.ToTensor(), transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]), ]) def load_model(model_path: str, device: str): model models.efficientnet_b0(weightsNone) model.classifier[1] torch.nn.Linear(1280, 4) state torch.load(model_path, map_locationdevice) model.load_state_dict(state) model.to(device).eval() return model def main(): parser argparse.ArgumentParser() parser.add_argument(--image, requiredTrue, help输入图片路径) parser.add_argument(--model, defaultmodels/mushroom_b0.pt) parser.add_argument(--device, defaultcpu) args parser.parse_args() model load_model(args.model, args.device) img Image.open(args.image).convert(RGB) with torch.no_grad(): prob torch.softmax(model(val_tf(img).unsqueeze(0)), dim1).squeeze(0) label ID2LABEL[torch.argmax(prob).item()] confidence float(torch.max(prob).item()) print(f识别结果: {label} 置信度: {confidence:.3f}) if __name__ __main__: main()load_model里要先声明模型结构再 load_state_dict顺序不能反否则会报维度不匹配。weightsNone是指不加载预训练权重因为训练时最后一行线性层已经改成 4 类了推理时只需要这份训练好的 state_dict。使用示例python predict.py --image data/test/edible/001.jpg --model models/mushroom_b0.pt输出就是一行识别结果方便在日志里持续记录。5.2 FastAPI 服务上传图片、返回标签与置信度要把识别能力开放给其他程序我给源码包配一个 FastAPI 服务。启动一次常驻模型对外提供 HTTP 接口。核心代码import io import threading from fastapi import FastAPI, UploadFile, File app FastAPI() device cpu model load_model(models/mushroom_b0.pt, device) # 推理锁并发同时进入时排队执行避免内存翻倍 infer_lock threading.Lock() app.post(/predict) async def predict(image: UploadFile File(...)): raw await image.read() img Image.open(io.BytesIO(raw)).convert(RGB) with infer_lock: with torch.no_grad(): prob torch.softmax(model(val_tf(img).unsqueeze(0)), dim1).squeeze(0) label ID2LABEL[torch.argmax(prob).item()] return { label: label, confidence: round(float(torch.max(prob).item()), 3), all_prob: {k: round(float(v), 3) for k, v in enumerate(prob.tolist())} }这里有两个细节。UploadFile读出来的是 bytes必须用io.BytesIO包一层再交给 PIL不能直接把UploadFile对象传给Image.open。infer_lock是必要的CPU 推理时多线程并发会让内存峰值成倍上涨蘑菇识别请求量不大串行推理完全够用内存更稳。启动命令uvicorn server:app --host 0.0.0.0 --port 8000然后用 curl 验证curl -X POST -F imagetest.jpg http://127.0.0.1:8000/predict返回的all_prob数组里四类概率和是 1调用方可以根据自己的业务场景重新切阈值而不是只能依赖服务端定死的策略。5.3 置信度阈值与 Top-N 输出把“不确定”说成“不确定”接口里只返回一个最大概率类别不够安全。两个相近菌种在湿润和干燥状态下外观变化很大模型给出“可食 0.55有毒 0.40”时单看第一名会误导使用者。我在接口里同时返回 Top-2topk torch.topk(prob, k2) response { label: ID2LABEL[torch.argmax(prob).item()], confidence: round(float(torch.max(prob).item()), 3), candidates: [ {label: ID2LABEL[idx.item()], probability: round(float(p), 3)} for idx, p in zip(topk.indices, topk.values) ], }加上之前说的阈值判断最终规则是最大概率低于 0.55 时即使顶部标签是可食也改成 uncertain 返回。这么做会牺牲一点表面上的“识别成功率”但换回来的是系统在“看不懂”时诚实表达的能力。识别类工具最怕的不是拒绝回答而是自信满满地给一个错误结论。6. 打包工程与源码交付从能跑到能交付的四个验证项源码 zip 交付出去对方第一件事并不是看模型效果而是解压后能不能跑起来。我在交付前会固定做四个验证每一项都卡得比较死。第一项干净环境安装。requirements.txt 里锁死关键版本torch1.13,2.1 torchvision0.14,0.16 fastapi0.100 uvicorn0.20 Pillow9.0 scikit-learn1.0torch 和 torchvision 版本跨度大时容易出现模型权重读取报错所以把这两个的范围收紧。REAME 里写清楚“用 Python 3.8 到 3.10pip install -r requirements.txt安装”并附带安装 numpy 和 matplotlib 的提示。第二项单图推理验证。README 给出从下载数据到训练、推理的最小命令序列python train.py完成后python predict.py --image data/test/edible/001.jpg --model models/mushroom_b0.pt必须能输出一条不报错的结果。第三项数据与模型目录约定。模型权重文件如果太大zip 里可以不塞但要留出一个models/目录文件名写成固定mushroom_b0.pt让训练脚本和推理脚本引用的是同一个相对路径而不是各自散乱的名字。第四项服务端自检。uvicorn server:app --port 8000起来之后至少用一张测试图跑一遍 POST确认返回 JSON 结构完整。我习惯把所有验证命令写进 README 的“快速开始”一节按命令执行完不出错才算交付完成。还有一个交付习惯想分享每次发新版本前我会新建一个 conda 环境严格按 README 重跑一遍。因为开发环境里装了一堆乱七八糟的包可能掩盖了缺依赖的问题。换新环境跑不通的一律按 bug 处理修到能通为止。源码包的价值不在代码本身有多漂亮而在别人拿到手后从解压 zip 到看到识别结果之间没有一座越不过去的山。这个习惯替我拦下了不少“你发的东西我跑不起来”的售后也希望帮到你。本文还有配套的精品资源点击获取
返回列表