植物识别全流程:ResNet训练与PyQt5界面部署实战
简介这套基于残差网络ResNet的植物类别识别代码采用 Python 与 PyTorch 编写面向需要从零搭建图像分类训练与推理流程的开发者或课程设计用户。资源包包含 8 个文件涵盖 3 个 Python 脚本、3 张文件夹结构提示图、1 份环境依赖文本和 1 篇说明文档整体仅 214KB轻量便于阅读。代码先扫描数据集文件夹生成带标签的 txt 文件自动划分训练集与验证集训练脚本会自动读取内容并适配新增分类目录而无需修改代码训练时实时显示进度条、准确率与损失值每个训练轮次结束后记录日志最终保存 model.ckpt 模型文件随后可运行 PyQt 界面脚本加载模型完成图片识别。由于压缩包不含数据集图片需自行收集图片按提示图放入对应文件夹下载后即可按照说明文档操作。目前已有 70 人学习下载适合快速上手深度学习图像分类实践。1. 三个 py 文件闭环的植物识别项目从数据到界面拿到手的这份资源是三个 Python 文件、一份说明文档和一个 requirement.txt模型采用 ResNet 结构基于 PyTorch 训练。项目不含数据集图片目录里只放了裸子植物、蕨类植物、被子植物三个分类的占位提示图真正的图片需要自己收集后填入对应文件夹。比较有参考价值的一点是整套流程从生成标签 txt、训练 CNN 到 PyQt5 图形界面推理全部打通数据量小的时候几分钟就能跑完一个 epoch适合作为深度学习中图像分类任务的入门框架来拆解。需要关注的核心问题有三个数据路径如何自动生成、ResNet 预训练模型如何接入自建数据集、以及训练好的 model.ckpt 怎样在界面里加载并完成预测。2. 数据文件夹与标签 txt01生成txt.py 的路径划分逻辑CNN 训练前要解决的第一个问题不是网络结构而是数据怎么喂进 DataLoader。工程上最省事的做法就是用文件夹名当标签遍历目录生成 txt每一行写“图片路径 类别索引”训练时再按行读回。这个项目里 01生成txt.py 做的正是这件事。2.1 目录结构与标签映射规则数据集按生物学分类组织成三个文件夹被子植物、裸子植物、蕨类植物。每个文件夹里放对应类别的图片文件名任意但建议统一为英文或数字以避免编码问题。文件夹名会按 ASCII 顺序排序后映射为数字标签这个映射关系会写进 txt训练和推理阶段都依赖同一份映射所以后续如果增加分类文件夹排序规则不变旧模型权重就不能直接复用需要重新训练。文件夹名排序索引标签值被子植物00裸子植物11蕨类植物22txt 文件的格式是每行图片绝对路径 标签用空格分隔。训练集和验证集按 8:2 比例随机划分这个比例写在脚本里需要精细调参时可以改成 9:1 或 7:3。2.2 生成脚本的核心逻辑import os import random from glob import glob # 数据集根目录 data_root 数据集 # 支持常见图片格式避免把系统隐藏文件带进来 exts (*.jpg, *.jpeg, *.png, *.bmp) train_lines [] val_lines [] # 遍历一级子目录每个子目录代表一个类别 for label, class_dir in enumerate(sorted(os.listdir(data_root))): class_path os.path.join(data_root, class_dir) if not os.path.isdir(class_path): continue # 收集该类别下所有图片 img_paths [] for ext in exts: img_paths.extend(glob(os.path.join(class_path, ext))) # 按 8:2 划分训练集和验证集 random.seed(42) random.shuffle(img_paths) split_idx int(len(img_paths) * 0.8) for p in img_paths[:split_idx]: train_lines.append(f{os.path.abspath(p)} {label}\n) for p in img_paths[split_idx:]: val_lines.append(f{os.path.abspath(p)} {label}\n) # 写入 txt 文件 with open(train.txt, w, encodingutf-8) as f: f.writelines(train_lines) with open(val.txt, w, encodingutf-8) as f: f.writelines(val_lines) print(f训练集 {len(train_lines)} 张验证集 {len(val_lines)} 张)这里有几个关键点。enumerate(os.listdir())的 enumerate 会自动按文件夹排序结果生成 0、1、2 等标签索引排序依据是文件系统返回的字节顺序。如果你的文件夹数量不固定这种动态映射方式能保证新增类别时训练脚本不用改动。random.seed(42)是固定随机种子确保每次划分结果一致方便复现实验。由于程序以绝对路径写入 txt后续不管终端工作目录在哪儿都能正确找到图片文件。注意一个问题如果不设置glob模式直接os.listdir会把 .DS_Store 之类的系统文件也当作图片导致 ImageFolder 加载时报UnidentifiedImageError。本脚本用glob(*.jpg)模式匹配来规避这点。另外中文路径在 Windows 下配合 OpenCV 的imread容易乱码这里写入 txt 的是os.path.abspath返回的原始路径训练端用 PIL/PyTorch 的Image.open读取则没有这个问题因为它是通过 Python 文件协议打开的不过还是建议数据路径中不要出现中文和空格。3. ResNet50 接入训练流程02CNN训练数据集.py 的模型替换与 epoch 监控训练脚本 02CNN训练数据集.py 是整套代码的重头戏。它的设计思路是类目数量完全从 txt 的标签最大值动态获取因此新增植物类别时训练脚本无需改代码。模型主干使用 ResNet默认加载 ImageNet 预训练权重最后一层全连接替换成自己的分类头。3.1 ResNet 结构选型与残差机制ResNet 解决的是深层网络退化问题。在植物图像这种细粒度分类任务上叶子纹理、花瓣形状的差异非常细微浅层网络容易欠拟合深层纯卷积网络在反向传播时梯度容易消失。ResNet 的核心是残差块让网络去学习输入与输出的差值公式化表示是H(x) F(x) x。如果身份映射恒等映射已经是最优解网络只要把 F(x) 的权重逼近零即可训练难度大幅下降。本项目的模型文件是 model.ckpt脚本里默认用 18 层 ResNet。如果你的图片是 224x224 输入处理 1000 类以内的分类任务ResNet18 和 ResNet34 在推理速度与精度之间更平衡。增加深度到 ResNet101 对植物细粒度分类的提升有限但训练时间几乎翻倍。import torch.nn as nn from torchvision import models num_classes 3 # 实际从 txt 读取最大标签 1 model models.resnet18(weightsmodels.ResNet18_Weights.IMAGENET1K_V1) # 替换最后一层全连接输出维度改为当前数据集类别数 model.fc nn.Linear(model.fc.in_features, num_classes)如果本地没有预训练权重PyTorch 会自动下载到~/.cache/torch/hub/checkpoints/。下载失败时可以把 URL 复制到浏览器手动下载再放到指定缓存目录。model.fc.in_features是 ResNet18 最后一个全连接层的输入维度 512这一行写法比硬编码 512 更通用换成 ResNet502048 维时不会出错。3.2 训练循环与监控项设计import torch import torch.nn as nn import torch.optim as optim from torch.utils.data import Dataset, DataLoader from torchvision import transforms from PIL import Image class PlantDataset(Dataset): 从 txt 读取图片路径与标签 def __init__(self, txt_path, transform): self.samples [] self.transform transform with open(txt_path, r, encodingutf-8) as f: for line in f: img_path, label line.strip().split( ) self.samples.append((img_path, int(label))) def __len__(self): return len(self.samples) def __getitem__(self, idx): img_path, label self.samples[idx] img Image.open(img_path).convert(RGB) return self.transform(img), label # 数据增强训练集随机裁剪加翻转验证集等比缩放 train_tf transforms.Compose([ transforms.RandomResizedCrop(224), transforms.RandomHorizontalFlip(), 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]) ]) train_loader DataLoader( PlantDataset(train.txt, train_tf), batch_size32, shuffleTrue, num_workers2 )RandomResizedCrop会随机裁剪出 224x224 的区域并缩放到 224相当于一种在线数据增强可以缓解植物图片背景复杂导致的过拟合。Normalize中的四个数字是 ImageNet 数据集的均值和标准差加载预训练权重时不能修改否则会破坏特征分布。num_workers2在 Windows 上如果报多进程错误改成 0 并用主进程加载数据。训练循环里需要同时关注三件事每个 epoch 的训练损失、验证集准确率和单 epoch 耗时。损失下降但准确率不升说明存在过拟合或学习率过大准确率上升但损失波动可能是 batch size 太小造成梯度震荡。训练结束后模型保存为 model.ckpt这个后缀名可以改成 .pth格式本就是 PyTorch 的 state_dict。4. PyQt5 推理界面与模型加载03pyqt界面.py 的前向预测链路模型训练完成后考验工程能力的是如何把权重文件变成别人能用的工具。03pyqt界面.py 用 PyQt5 搭建了一个最小可用的推理界面交互逻辑是点击按钮选择本地图片程序加载 model.ckpt对图片做与训练时完全一致的预处理最后在窗口上显示预测类别和置信度。4.1 模型加载与类别映射还原训练脚本保存的model.ckpt是model.state_dict()的结果不包含网络结构。加载时必须先实例化一个同名 ResNet 模型再load_state_dict把权重灌进去。这里最容易犯的错误是忘记修改model.fc输出维度直接加载权重会报size mismatch。import torch from torchvision import models from PyQt5.QtWidgets import QApplication, QLabel, QPushButton, QFileDialog, QVBoxLayout, QWidget class Predictor: def __init__(self, ckpt_path, class_names): self.class_names class_names self.model models.resnet18(weightsNone) self.model.fc torch.nn.Linear(self.model.fc.in_features, len(class_names)) self.model.load_state_dict(torch.load(ckpt_path, map_locationcpu)) self.model.eval() # 预处理参数必须与训练时一致 self.transform 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 predict(self, img_path): img Image.open(img_path).convert(RGB) tensor self.transform(img).unsqueeze(0) with torch.no_grad(): logits self.model(tensor) prob torch.softmax(logits, dim1) conf, idx torch.max(prob, dim1) return self.class_names[idx.item()], conf.item()map_locationcpu确保在无 GPU 机器上也能加载模型否则用 GPU 训练的权重在 CPU 环境会报 CUDA 相关错误。.eval()是必须的它会关闭 Dropout 和 BatchNorm 的滚动更新否则同样的输入两次预测结果可能不一致。torch.no_grad()上下文管理器避免构建计算图批量推理时能省下不少显存。4.2 界面事件与置信度显示PyQt 的交互逻辑本质上就是信号槽。点击按钮触发open_file_dialog拿到文件路径后调用predictor.predict再把结果写到 QLabel 上。界面控件类型作用选择图片按钮QPushButton触发文件选择对话框原始图片标签QLabel显示待识别图片缩略图结果标签QLabel显示类别名称与置信度百分比推理时需要注意 OpenCV 的cv2.imread读取路径带中文时返回空数组导致cv2.cvtColor报错。推荐用 PIL 的Image.open读取因为 PyQt 的QFileDialog.getOpenFileName返回的是本地文件系统路径中文路径概率很高。还有一点unsqueeze(0)是把单张 3x224x224 的图片变成 1x3x224x224 的 batch这是 PyTorch 卷积层要求的四维输入格式漏掉会直接报维度错误。5. 模型验证与三个常见坑从 model.ckpt 到界面识别5.1 快速验证模型是否真的学到特征训练结束后不要急着打开界面。先用命令行做一次批量验证统计每个类别的精确率和召回率这能帮你判断是数据问题还是模型问题。python -c from PIL import Image from torchvision import transforms import torch model torch.load(model.ckpt, map_locationcpu) 实际上 model.ckpt 只是 state_dict需要先构建模型再加载。常见的做法是在 02 脚本里加一个验证函数遍历 val.txt统计预测正确的图片数量。如果验证集准确率低于 90%先检查 train.txt 和 val.txt 是否混入了同一个文件夹下的重复图片随机种子固定后每次划分结果一致不会出现混叠。5.2 坑一加载模型报 size mismatchfc层输出维度和保存时不一致就会报这个错。很多人在训练时把num_classes写死成 3后期增加到 5 类重新训练加载旧权重时忘了同步修改就出现 mismatch。解决办法是保存模型时把num_classes一起存进文件。torch.save({ state_dict: model.state_dict(), num_classes: num_classes, class_names: [被子植物, 裸子植物, 蕨类植物] }, model.ckpt)加载时使用torch.load读取字典内容再动态创建模型结构这样无论分类数怎么变都不会出问题。5.3 坑二推理界面精度低而训练精度高这种症状十有八九是预处理不一致。训练时的RandomResizedCrop(224)会随机裁剪验证和推理时用的是Resize(256)加CenterCrop(224)。如果推理代码里图片没有缩放就直接送入模型输入尺寸就不匹配PyTorch 虽然会报错但如果你用torchvision.transforms.ToTensor()后忘了 Normalize结果会差很多。排查方式是打印输入 tensor 的均值标准差对比训练集预处理后的分布。5.4 坑三CPU 推理太慢没有 GPU 的机器跑 ResNet18 推理一张图大约需要 200 到 500 毫秒如果每次都要等待模型加载就不适合做实时识别。可以改成启动时加载一次模型界面保持常驻或者把model.half()转换成半精度推理速度提升约一倍但 CPU 上部分算子不支持 half 需要做兼容处理。还可以用torch.jit.script对模型做 TorchScript 导出去掉 Python 动态图开销在不改精度的前提下通常能快 15% 到 30%导出后直接torch.jit.load加载即可。本文还有配套的精品资源点击获取
上一篇/下一篇内容由系统自动关联
返回资讯列表 →