ARTICLE DETAIL

资讯详情

深耕网站建设、视觉设计与SEO优化的一线实战洞察。

蘑菇识别系统源码实战:从图像分类到PyTorch模型训练全流程解析

蘑菇识别系统源码实战:从图像分类到PyTorch模型训练全流程解析 简介这份Python蘑菇识别系统源码是一套基于深度学习的完整图像识别项目面向熟悉Python基础、希望系统学习图像分类与计算机视觉的开发者也适合生物、农业领域需要自动化识别蘑菇种类的技术人员作为参考。压缩包共54个文件、大小约31MB其中9个py源码负责模型搭建、训练与预测20个pyc为编译缓存23张png为示例图片或界面素材另有md说明文档与txt配置文件目录结构清晰便于定位与二次开发。系统涉及图像预处理、卷积神经网络特征提取、数据集划分、模型调优与结果评估等关键环节基本覆盖了图像分类任务的完整流程可帮助读者把一个想法落地为可运行的程序。已有382人学习下载适合需要完整可运行样例、快速上手深度学习项目的读者无论是课程设计还是实际应用都具备不错的参考价值。1. 拿到这份蘑菇识别源码先想清楚它解决的是什么问题“Python蘑菇识别系统源码.zip”——我猜很多人解压后的第一件事是双击打开 train.py然后盯着几百行代码发呆。我建议你先别急着跑花五分钟想清楚这件事蘑菇识别在深度学习里属于典型的图像分类问题输入一张蘑菇照片输出一个品种标签可能附带置信度。它的使用场景很具体野外采菌的人想确认手里这朵能不能吃食品行业做菌类分级农业科普机构批量鉴定标本。市面上的免费 python 源码包很多这套系统的价值不在模型多新而在于它把数据整理、训练、推理、评估串成了一条完整链路。适合刚入门 CV 的 Python 开发者照着跑通也适合做过分类任务但没处理过细粒度识别的熟手研究它的数据策略——毒蘑菇和食用菌长得极像这种类间差异小的问题恰恰比模型结构更考验工程细节。这篇笔记按“解压→配环境→备数据→训练→排错”的路径把源码里最值得抄的几段拆开讲。2. 拆解源码结构先别双击 train.py按这三步把包吃透很多人在拿到 zip 包这一步就翻车了原因不是代码而是压缩包本身。这一步的目标是把源码完整落地到本地、配好隔离环境、再根据目录结构判断这套系统的“骨架”长什么样。急不得。2.1 源码包落地解压编码、目录确认与虚拟环境隔离如果你是 Windows 用户直接右键“全部解压缩”通常没问题。但我更推荐用 7-Zip尤其是源码包是从网盘或群聊里转手来的压缩软件对文件名编码的处理差异很大Windows 自带的解压工具遇到 GBK 编码的文件名偶尔会解出一堆乱码目录你后面加载数据时路径全对不上还以为是代码写错了。Linux 用户就用标准命令unzip Python蘑菇识别系统源码.zip -d mushroom_project cd mushroom_project find . -maxdepth 2 -type d | sort-d指定解压到目标目录避免把一堆散文件直接撒在当前文件夹里污染工作区。解压后先看目录结构一个规范的蘑菇识别项目通常包含这几个目录data/或dataset/存放图片和划分脚本、models/网络结构定义、utils/公共工具函数、train.py、predict.py、requirements.txt。如果解压后只有一个孤零零的.py文件说明这不是完整源码只是核心逻辑片段你得自己补数据管线和训练循环。目录确认后第一件事是建虚拟环境。用系统全局 Python 直接pip install -r requirements.txt大概率把 PyTorch、opencv 装到你不想碰的系统环境里等下次跑其他项目时发现问题你根本不知道是谁改坏了什么。建议在 VSCode 里做 Python 环境配置时顺手把虚拟环境选好后续调试少很多麻烦。命令如下python -m venv venv # Windows PowerShell: venv\Scripts\Activate.ps1 source venv/bin/activate python -m pip install --upgrade pip我的习惯是先把venv建好再读requirements.txt解耦这两个步骤。因为这套源码如果是两年前写的里面的依赖版本可能和当前 PyTorch 版本冲突先配环境再改依赖能少走一段弯路。2.2 依赖清单逐行过torch、torchvision、opencv 三方怎么配不打架打开requirements.txt不要一键安装完就看下一行。重点核对三组依赖。第一组是torch和torchvision这两个版本必须匹配。源码里如果写torch1.10.0那torchvision很可能是0.11.0你要是直接把 torch 升到 2.x 而 torchvision 没跟着升models.resnet18(pretrainedTrue)这一行就能报出“找不到属性”或者底层的 C 算子错误。第二组是opencv-python蘑菇识别里图片读取经常用到它但 opencv 的版本如果和 numpy 的编译版本不对齐cv2.imread返回的数组形状和你预期不一样查起来极其隐蔽。第三组是scikit-learn它一般用在评估阶段画混淆矩阵、算 precision/recall版本差点问题不大但低于 0.24 时有些 API 名不一样。如果这份源码没有附requirements.txt常见做法是手写一个最小依赖torch、torchvision、opencv-python、pillow、numpy、matplotlib、tqdm、scikit-learn。逐个安装pip install torch torchvision --index-url https://download.pytorch.org/whl/cpu pip install opencv-python pillow numpy matplotlib tqdm scikit-learn如果你有 N 卡且装了 CUDA把--index-url .../cpu去掉换成默认的 PyPI 源就行。装完后跑一句最简验证确认 torch 和 torchvision 能正常导入python -c import torch, torchvision; print(torch.__version__, torchvision.__version__)torch 和 torchvision 的版本兼容性是这个项目里最玄学的一环。版本对不上时训练经常在中途报错错误信息指向torchvision某个.so文件加载失败。我的建议是如果源码没锁版本优先装“最新稳定 torch 配套最新 torchvision”用那句验证命令确认能打印出版本号再往下走。2.3 入口脚本定位从 train.py / predict.py 反推整个工作流依赖配完回到源码本身。通常项目里有train.py和predict.py两个入口但没有的话也别慌用一条命令快速定位所有 Python 文件找出带__main__或者if __name__ __main__的grep -rl __main__ --include*.py .找到入口后按住Ctrl点击关键函数跳转把main()里调用的函数列表列出来就能画出这套系统的工作流。我一般会看三处第一数据是从文件夹读还是从 CSV 读——这决定了你是否需要按它的目录规范重排图片第二模型是用torchvision.models里现成的还是自定义结构——现成的可以省事自定义的结构则要检查 forward 里有没有用到特定输入尺寸第三predict.py里加载权重的路径是写死的还是命令行参数——写死的话你要么把权重放到指定位置要么改源码。做完这三步你对这套系统的认知就已经从“一团代码”变成了“有输入有输出的工具”。这也是把源码读进脑子的最短路径不是逐行读而是找入口、找数据流、找边界。3. 数据集这一关不过模型训练就是玄学蘑菇识别和猫狗分类最大的区别在于数据集普遍很小而且类之间的视觉差异极端。有的品种只差在菌褶颜色或菌柄根部的一个环人眼都要凑近看半天。这种场景下数据整理工作的优先级高于模型选型。很多人拿源码直接跑发现 loss 降不下去以为是模型问题其实是训练集里混着大量背景杂乱、光照失衡的图片。3.1 蘑菇数据集长什么样从爬虫抓图到 ImageFolder 目录规范先明确一个事实公开的蘑菇图片分类数据集很少且大多是表格型数据——比如 UCI 的 Mushroom 数据集里面是蘑菇的形态学属性菌盖形状、气味、颜色不是图片。所以做蘑菇图片识别数据通常要自己收集。常见做法是用 Python 爬虫从公开图片站按品种关键词抓图这一步在学术用途和自用场景下是合法的。抓图脚本的核心逻辑很简单请求图片 URL按品种名建目录把图片写进去。# crawl_mushroom.py —— 按关键词抓公开图片按品种归档 import os import time import requests species_list [Amanita muscaria, Cantharellus cibarius] save_root raw_images for species in species_list: os.makedirs(os.path.join(save_root, species), exist_okTrue) # 这里以接口返回的图片URL列表为例实际按公开图片站的限制写 image_urls get_image_urls(species, limit300) for i, url in enumerate(image_urls): try: r requests.get(url, timeout10) ext url.rsplit(., 1)[-1] if . in url else jpg out os.path.join(save_root, species, f{i:04d}.{ext}) with open(out, wb) as f: f.write(r.content) time.sleep(1) # 控制抓取频率 except Exception as e: print(species, i, e)代码里的get_image_urls是个占位函数实际按你选的图片源实现关键是三个细节time.sleep(1)控制频率避免被限流按品种建目录省去后面打标签的功夫异常捕获放在单张图片级别别让一张坏图中断整个爬取任务。抓完图片后必须做一轮人肉清洗把所有.png、.webp混入的、分辨率过低小于 300×300、明显是插画或示意图的图片删掉。这一步没有捷径我一般会写个小工具把所有图片拼成网格图快速扫一眼。数据清洗完接下来就是把raw_images整理成 PyTorch 的 ImageFolder 规范目录。ImageFolder 会自动把一级子目录名当作类别名读取时按字母序生成索引 0、1、2……这个特性很方便但也是个隐性坑——训练和推理时的类别编号取决于目录名的字母序和你的直觉顺序不一定一致。3.2 数据增强和样本不均衡蘑菇识别最常见的两个翻车点蘑菇数据集的典型病症是有的品种你能抓到一千张图有的品种全网只有四五十张。直接用原始分布训练模型会对样本多的品种过拟合对稀有品种几乎不学习。两个常用解法是类别权重和过采样。类别权重在损失函数上做文章给稀有类别更大的梯度权重过采样则是在派生数据集层面把稀有类别的图片重复采样让它每次 epoch 出现的次数和其他类别接近。数据增强同样关键。蘑菇图片的一大特点是拍摄环境杂乱树叶、土壤、手部阴影占了半张图而且菌盖角度稍微一偏轮廓就完全变了。我一般会用一组温和的增强不做极端裁剪因为蘑菇的判别特征可能就在图片边缘from torchvision import transforms train_transforms transforms.Compose([ transforms.Resize((224, 224)), transforms.RandomHorizontalFlip(p0.5), transforms.RandomRotation(degrees15, fill0), transforms.ColorJitter(brightness0.2, contrast0.2, saturation0.1), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]) ])RandomRotation(15)这个幅度对蘑菇是合适的旋转超过 30 度菌盖和菌柄的比例关系就失真了模型学到的是“倾斜的蘑菇”而不是“蘑菇”。ColorJitter的饱和度扰动调得保守有些毒的蘑菇和可食蘑菇的区别就在颜色深浅扰动太大等于人为抹掉了关键特征。归一化参数用的 ImageNet 统计量是因为接下来要用预训练模型做迁移学习输入分布保持一致。3.3 一份可直接改用的数据整理脚本把爬取的零散图片整理成训练/验证两套目录是每次都要重复的工序。我一般写一个脚本输入raw_images输出按 8:2 划分的data/train和data/val。脚本逻辑不复杂但能避免手动拖文件时把类目弄混# organize_data.py —— 把 raw_images 整理成 ImageFolder 格式 import os import shutil import random from pathlib import Path random.seed(42) raw_root Path(raw_images) target_root Path(data) val_ratio 0.2 for species_dir in sorted(raw_root.iterdir()): if not species_dir.is_dir(): continue species_name species_dir.name images [p for p in species_dir.iterdir() if p.suffix.lower() in {.jpg, .jpeg, .png}] random.shuffle(images) n_val int(len(images) * val_ratio) val_images, train_images images[:n_val], images[n_val:] for split, imgs in [(train, train_images), (val, val_images)]: out_dir target_root / split / species_name out_dir.mkdir(parentsTrue, exist_okTrue) for img in imgs: shutil.copy2(img, out_dir / img.name) print(f{species_name}: total{len(images)}, train{len(train_images)}, val{n_val}) print(done:, target_root)这个脚本的要点有两个。第一random.seed(42)必须显式设置否则每次划分的验证集不一样你两次实验的结果就没有可比性第二copy2保留文件元信息实际也可以换成os.symlink做软链省磁盘空间但软链在 Windows 上权限要求多跨平台项目我用复制更省心。划分完成后你可以直接数一下每个类的文件数如果某个品种的train目录少于 20 张图我建议你直接删掉这个类或者回源头再补抓一轮否则这个类在训练时就是“可有可无”的噪声。4. 训练与推理把模型跑通的命令、参数和代码数据目录准备好接下来就是核心环节训练和推理。这一章的目的不是教你写一个不能再精简的 demo而是让你理解训练脚本里那些参数的可调范围、每个参数失控时的表现以及推理脚本为什么必须和训练时的预处理保持一致。4.1 训练脚本骨架与核心参数表用 PyTorch 的torchvision作为 backbone 是这类小项目最靠谱的起点。下面是精简后的训练骨架去掉日志和可视化保留完整训练链路# train.py —— 蘑菇识别训练骨架 import argparse import torch import torch.nn as nn from torch.utils.data import DataLoader from torchvision import datasets, transforms, models def get_transforms(image_size): # 数据增强和归一化训练集和验证集分开处理 train_tf transforms.Compose([ transforms.Resize((image_size, image_size)), transforms.RandomHorizontalFlip(p0.5), transforms.RandomRotation(degrees15, fill0), transforms.ColorJitter(brightness0.2, contrast0.2, saturation0.1), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]), ]) val_tf transforms.Compose([ transforms.Resize((image_size, image_size)), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]), ]) return train_tf, val_tf def main(args): train_tf, val_tf get_transforms(args.image_size) train_ds datasets.ImageFolder(args.data_dir /train, transformtrain_tf) val_ds datasets.ImageFolder(args.data_dir /val, transformval_tf) train_loader DataLoader(train_ds, batch_sizeargs.batch_size, shuffleTrue, num_workersargs.num_workers) val_loader DataLoader(val_ds, batch_sizeargs.batch_size, shuffleFalse, num_workersargs.num_workers) # 迁移学习加载预训练模型替换最后一层全连接 model models.resnet18(weightsmodels.ResNet18_Weights.IMAGENET1K_V1) in_features model.fc.in_features model.fc nn.Linear(in_features, len(train_ds.classes)) device torch.device(cuda if args.gpu 0 and torch.cuda.is_available() else cpu) model model.to(device) criterion nn.CrossEntropyLoss() optimizer torch.optim.Adam(model.parameters(), lrargs.lr, weight_decay1e-4) best_acc 0.0 for epoch in range(args.epochs): model.train() total_loss, correct 0.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() total_loss loss.item() * images.size(0) epoch_loss total_loss / len(train_ds) # 每个 epoch 结束跑一遍验证集保留准确率最高的权重 model.eval() val_correct, val_total 0, 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) val_correct (preds labels).sum().item() val_total labels.size(0) val_acc val_correct / val_total print(fepoch {epoch1}/{args.epochs} loss{epoch_loss:.4f} val_acc{val_acc:.4f}) if val_acc best_acc: best_acc val_acc torch.save(model.state_dict(), best_mushroom.pth) if __name__ __main__: parser argparse.ArgumentParser() parser.add_argument(--data_dir, typestr, defaultdata) parser.add_argument(--epochs, typeint, default30) parser.add_argument(--batch_size, typeint, default32) parser.add_argument(--lr, typefloat, default1e-3) parser.add_argument(--image_size, typeint, default224) parser.add_argument(--num_workers, typeint, default2) parser.add_argument(--gpu, typeint, default0) args parser.parse_args() main(args)训练脚本里最需要解释的参数有五个它们也是调参时最先要动的参数取值范围作用与调整逻辑--image_size224 / 299 / 384输入分辨率。224 是 ImageNet 预训练的标准尺寸数据量大或想提升小纹理识别就上 299显存压力也线性增长--batch_size16~64受显存限制。数据集小的时候 batch 太大容易收敛到尖锐局部最优我用 16~32--lr1e-4~1e-3迁移学习场景下 1e-3 稍偏激进预训练权重只需要微调建议从 1e-3 跑几个 epoch 后观察 loss 是否震荡震荡就降一半--epochs20~50蘑菇数据集小30 epoch 左右 base 模型就能收敛过多会过拟合--num_workers2~8数据加载线程数。Windows 上设太大会报 “DataLoader worker” 错误一般 2~4 即可执行训练的命令python train.py --data_dir data --epochs 30 --batch_size 32 --lr 1e-3训练过程中如果val_acc在某个 epoch 后开始小幅度波动但不上升说明学习率该调低了如果验证集准确率远低于训练集说明过拟合增大weight_decay或加 Dropout。这些都是通用策略但适用面很广。4.2 迁移学习选型为什么蘑菇这种小数据集不推荐从头训练蘑菇识别的数据集往往只有几百到两三千张图从头训练一个深度模型几乎必然过拟合。最省事、最稳定的是用 ImageNet 预训练权重微调。ResNet18、ResNet34、EfficientNet-B0、MobileNetV3都是不错的选择它们的取舍很清晰ResNet18结构简单、对设备要求低、调试 BillGates 式 bug 最容易新手拿来跑通流程最合适CPU 上也能训。ResNet34比 18 多了一层保留了结构简洁的优点适合图片分辨率不高、但类别数超过十几类的情况。EfficientNet-B0性价比高参数量少但准确率不差不过它的输入尺寸不是默认的 224代码里要显式改成它的image_size。MobileNetV3模型小适合未来要部署到手机端或树莓派的场景训练时收敛略慢。在小数据集上ResNet18和EfficientNet-B0的最终正确率常常只差一两个百分点不要为了“看起来很新”的模型牺牲调试便利性。真正的训练技巧是把预训练权重冻结住只训练最后一层全连接跑几个 epoch 后再解冻全部层做微调。这个技巧在很多源码里没有体现但对小数据集非常有效。实现方式是在初始化后用requires_grad_(False)冻结主干层训练后半段再设回True。4.3 推理脚本单张图片识别从加载权重到输出 Top-5训练完落地到实际使用就要靠推理脚本。推理脚本最常见的错误是预处理和训练不一致导致模型看到的世界完全不是它训练时熟悉的样子。下面的推理脚本保持了与训练一致的 Resize、Normalize# predict.py —— 单张蘑菇图片推理 import argparse import torch from PIL import Image from torchvision import transforms, models CLASS_NAMES [鸡油菌, 牛肝菌, 毒蝇伞, 香菇] # 按 data/train 下的目录字母序填 def main(): parser argparse.ArgumentParser() parser.add_argument(--image, typestr, requiredTrue, help输入图片路径) parser.add_argument(--weights, typestr, defaultbest_mushroom.pth) parser.add_argument(--num_classes, typeint, default4) args parser.parse_args() # 与训练完全一致的预处理 val_tf transforms.Compose([ transforms.Resize((224, 224)), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]), ]) model models.resnet18(weightsNone) model.fc torch.nn.Linear(model.fc.in_features, args.num_classes) model.load_state_dict(torch.load(args.weights, map_locationcpu)) model.eval() img Image.open(args.image).convert(RGB) tensor val_tf(img).unsqueeze(0) # 增加 batch 维度 with torch.no_grad(): logits model(tensor) probs torch.softmax(logits, dim1)[0] top5 torch.topk(probs, kmin(5, len(CLASS_NAMES))) print(预测结果) for idx, prob in zip(top5.indices, top5.values): print(f {CLASS_NAMES[idx]:10s} {prob.item():.3f}) if __name__ __main__: main()推理脚本里两个细节值得解释。第一Image.open(...).convert(RGB)这张图如果是 RGBA 四通道 PNG不做 convert 就会直接变成 4 通道输入模型报错如果图是灰度图convert(RGB)会把单通道复制成三份保持和训练数据分布一致。第二torch.topk取 Top-5 而不是 Top-1这是蘑菇识别的实际需要毒蝇伞和可食用的橙盖鹅膏在视觉上高度相似模型输出第二或第三高概率的类别往往能提醒用户“这个预测不确定性很高别吃”。执行推理python predict.py --image test_photo.jpg --weights best_mushroom.pth看到输出的概率分布后我的经验是最高概率低于 0.7 的预测结果直接当成“未知”处理不要信任。这个置信度阈值在后续批量识别里会是一个重要的过滤参数。5. 避坑跑这套系统最容易翻车的 5 个位置以下五条踩坑记录有的是我帮人排错时遇到的有的是自己第一次跑类似项目时浪费过一晚上的。每条都按“现象→原因→解决”写你遇到时可以直接对照。5.1 解压时提示“文件损坏”或要求输入密码现象从群聊或网盘下下来的 zip 包双击解压到一半报“不可预料的压缩文件末端”或者直接弹出密码框。原因一是文件在传输过程中被截断二是包被二次压缩或伪加密——伪加密是 zip 格式的常见问题压缩工具把“加密标志位”置位但你输入什么密码都不对。解决先用 7-Zip 打开这个 zip在“信息”面板看压缩文件是否完整如果提示头部损坏重新下载源文件。如果只是伪加密7-Zip 可以直接把文件拖出来Windows 自带解压器反而会拦截。5.2 torch 和 torchvision 版本不匹配训练中途报底层错误现象train.py能启动前几分钟也正常跑到第二个 epoch 时突然弹出AttributeError: module torch has no attribute ...或者一条指向 torchvision 动态库加载失败的 C 错误。原因torch 和 torchvision 是分开编译的版本错位时 C 算子表对不上。解决在虚拟环境里统一重装匹配版本不要单独升级某个库。先卸载再装pip uninstall torch torchvision -y pip install torch torchvision --index-url https://download.pytorch.org/whl/cpu装完一定要跑一遍上一章那条版本验证命令。这个组合是这套系统里最容易造成的“黑匣子”问题错误信息完全不指向版本排查优先级排最前面。5.3 图片读取后多了一个通道模型前向直接报错现象用PIL.Image.open()读图然后datasets.ImageFolder训练时报Expected 3-channel input, got 4-channel input。原因数据里有 RGBA 的 PNG 或者带透明通道的图片ImageFolder 不会自动丢弃 alpha 通道。解决在organize_data.py的清洗阶段就把所有图片统一转成 RGB 并重存一遍from PIL import Image for img_path in all_images: img Image.open(img_path).convert(RGB) img.save(img_path)这比在 Dataset 里做pydicom之类的兼容处理要彻底得多顺手把 EXIF 旋转信息也抹掉了避免图片进模型前被ImageFolder错误旋转。5.4 训练和推理时的类别顺序对不上识别结果全是错的现象训练出的模型在验证集上正确率 90%但用predict.py识别单张图片时明明是一张鸡油菌结果却打印出“香菇”。原因ImageFolder 按目录名字母序生成类别索引训练脚本里len(train_ds.classes)拿到的顺序和你在predict.py里手写的CLASS_NAMES顺序不一致。解决不要手写类别列表从训练好的数据里导出import torch from torchvision import datasets train_ds datasets.ImageFolder(data/train) print(train_ds.classes)把输出的类别列表原样粘贴到predict.py的CLASS_NAMES。如果你在训练后改了data/train的目录结构必须重新生成一次这个列表否则错位会悄无声息地存在。5.5 显存不够batch_size 调小后反而更慢现象CUDA out of memory报错后把--batch_size从 32 改成 8训练倒是能跑但每个 epoch 的耗时反而比之前更长。原因batch 太小导致 GPU 吞吐量未打满数据加载和 GPU 计算之间产生了大量空闲等待。解决先确认瓶颈在显存还是数据加载。如果是显存batch_size16是甜点值同时把--num_workers调大到 4如果还是慢检查--image_size是不是被设到了 384回退到 224 会省一半显存。实在不行就在训练代码里启用梯度累积用 4 个小 batch 模拟一个大 batch 的梯度更新而不是一味调小 batch。6. 进阶把单图识别扩展成批量识别并用混淆矩阵验证可靠性源码里通常只提供单张图片推理但实际使用时你手头可能有一整个文件夹的标本照片。批量识别的意义不只是省去逐张执行命令的时间而是能顺带统计所有预测的置信度分布用数据判断这套模型在你的采集场景下到底可不可靠。# batch_predict.py —— 批量识别并输出置信度过滤结果 import argparse from pathlib import Path import torch from PIL import Image from torchvision import transforms, models def main(): parser argparse.ArgumentParser() parser.add_argument(--input_dir, typestr, requiredTrue) parser.add_argument(--weights, typestr, defaultbest_mushroom.pth) parser.add_argument(--threshold, typefloat, default0.7) args parser.parse_args() model models.resnet18(weightsNone) model.fc torch.nn.Linear(model.fc.in_features, 4) # 按实际类别数改 model.load_state_dict(torch.load(args.weights, map_locationcpu)) model.eval() tf transforms.Compose([ transforms.Resize((224, 224)), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]), ]) for img_path in sorted(Path(args.input_dir).glob(*.jpg)): img Image.open(img_path).convert(RGB) tensor tf(img).unsqueeze(0) with torch.no_grad(): probs torch.softmax(model(tensor), dim1)[0] conf, idx torch.max(probs, 0) print(f{img_path.name}: class{idx.item()} conf{conf.item():.3f} f{(跳过) if conf.item() args.threshold else }) if __name__ __main__: main()批量识别跑完后你手上实际上有了一份“模型自评报告”如果超过一半图片的置信度低于阈值说明这批图片的拍摄环境、角度或品种分布和训练集差异太大此时回来用验证集算混淆矩阵效果比盲目调模型更直接。混淆矩阵用scikit-learn一行就能算配合matplotlib画出来能一眼看出哪些蘑菇品种互相混淆——通常这也是真正需要补充训练数据的方向而不是换更深的网络。我自己第一次跑通这套系统时卡在 5.2 节的版本问题上整整一个晚上最后发现是 pip 默认装了不兼容的 torchvision气得不行。后来养成一个习惯任何源码项目先花十分钟把环境、数据和入口理顺再跑训练绝不直接双击。这套流程救过我后续很多次“那种明明代码没问题但就是跑不动”的场面。希望帮到你。本文还有配套的精品资源点击获取
返回列表