ARTICLE DETAIL

资讯详情

深耕商务建站与企业官网运营的一线实战洞察。

基于CNN深度学习的大米识别:PyTorch训练与PyQt界面全流程实战

基于CNN深度学习的大米识别:PyTorch训练与PyQt界面全流程实战 简介本资源是一套基于PyTorch框架的CNN深度学习大米识别实战项目面向具备Python基础、希望入门图像分类的开发者与在校学生。项目以大米品种识别为任务场景完整覆盖从数据预处理、模型训练到可视化交互的全流程适合作为课程设计或深度学习练手案例。压缩包共906个文件包含900张jpg图片构成的多类别数据集、3个txt说明与日志文件、3个py脚本整体约11.98MB体积轻量便于本地运行。数据集已做短边补灰边成正方形及旋转角度等增强处理脚本依次完成标签文本生成、模型训练与PyQt界面调用训练过程会保存模型权重并记录每个epoch的验证损失与准确率。目前已有152人学习读者可借此掌握图像分类数据组织、CNN训练调参与桌面端推理界面的完整实现思路。1. 从一堆文件名到可训练数据集这套大米识别资源到底能跑出什么如果你手头正好有一批按类别分文件夹存放的大米图片文件名里还带着rotated45、flip这类后缀想快速验证一个 CNN 分类模型能不能把它们区分开这套「基于 CNN 深度学习的大米识别-含图片数据集」就是冲这个场景来的。它把数据增强、标签生成、模型训练和 PyQt 可视化界面串成了一条完整链路技术栈是 Python PyTorch适合刚接触深度学习图像分类、需要跑通一个完整项目的人也适合想拿它当模板改自己数据集的从业者。数据集里能看到 Ipsala 等类别图片已经做过旋转和翻转增强省掉了从零写增强脚本的功夫。下面按「资源是什么、怎么装、怎么跑、坑在哪」的顺序拆开讲。2. 环境配置与依赖安装把 PyTorch 和 PyQt 装进同一个解释器2.1 为什么这套代码对版本敏感这套资源的核心依赖只有两块PyTorch 负责 CNN 训练PyQt 负责最后的可视化界面。问题在于PyTorch 的版本和 CUDA 驱动、Python 版本之间是强绑定的而 PyQt 又对 Python 版本有自己的一套要求。我一般会先确认本机 Python 版本再决定装哪个 PyTorch 轮子。如果你用的是较新的 PythonPyQt5 的某些旧版本可能装不上这时候要么降 Python要么换 PyQt6 并改几行导入代码。资源里给了requirement.txt但环境仍然需要自行配置这一步没有后悔药装错了后面训练和界面都会报错。常见做法是建一个独立虚拟环境避免和系统里已有的包打架。下面这套命令是我在 Windows 和 Linux 上都跑过的顺序先建环境再装依赖最后单独确认 PyTorch 能不能调用 GPU。# 创建虚拟环境Python 版本建议 3.8 到 3.10 python -m venv rice_env # 激活环境Windows rice_env\Scripts\activate # 激活环境Linux / macOS source rice_env/bin/activate # 先装 PyTorch具体命令按你的 CUDA 版本去官网生成 pip install torch torchvision --index-url https://download.pytorch.org/whl/cu118 # 再装其余依赖requirement.txt 里通常包含 numpy、Pillow、PyQt5 等 pip install -r requirement.txt逻辑说明先建虚拟环境是为了隔离依赖避免 PyQt 和系统里其他 Qt 程序冲突。PyTorch 单独装是因为它的安装源和普通 PyPI 不一样混在requirement.txt里容易装成 CPU 版还不自知。参数上--index-url后面的cu118对应 CUDA 11.8你要按自己显卡驱动支持的版本换装完用下面这段验证。import torch print(torch version:, torch.__version__) print(cuda available:, torch.cuda.is_available()) print(device name:, torch.cuda.get_device_name(0) if torch.cuda.is_available() else CPU only)如果cuda available是 False训练会退回 CPU速度差很多但代码本身能跑。这时候别急着怀疑代码先查驱动和 PyTorch 版本是否匹配。2.2 requirement.txt 里没写全的隐式依赖requirement.txt通常只列了直接依赖但实际跑起来还会用到tqdm做进度条、matplotlib画损失曲线、scikit-learn算混淆矩阵这类库。我一般会先按 txt 装一遍然后直接跑01数据集文本生成制作.py报什么缺什么再补。这样比一次性装一大堆用不上的包更省事。另外PyQt 界面在部分 Linux 发行版上需要额外的系统库比如libxcb-xinerama0缺了会直接闪退报错信息里会提到xcb遇到时用系统包管理器补上即可。提示装完 PyTorch 后先跑一次torch.cuda.is_available()确认 GPU 可用再往下走否则后面训练慢到你会怀疑人生。3. 数据增强与标签生成01 脚本到底对图片做了什么3.1 短边补灰边与旋转增强的实现逻辑这套资源对数据集的预处理有两个关键动作一是把非正方形图片通过短边补灰边变成正方形二是做旋转角度增强。补灰边的好处是避免直接拉伸导致大米形状变形旋转增强则是为了让模型见过更多角度下的米粒。资源里的文件名已经带了rotated45和flip后缀说明增强后的图片已经落盘01数据集文本生成制作.py要做的是读取这些图片路径并生成标签文本。我一般会先看一眼数据集目录结构确认每个类别一个文件夹文件夹名就是类别名。下面这段代码是我按常见做法补的目录检查逻辑你可以直接拿去对照自己的数据集。import os data_root dataset # 数据集根目录按实际路径改 classes sorted(os.listdir(data_root)) print(类别列表:, classes) for cls in classes: cls_dir os.path.join(data_root, cls) if not os.path.isdir(cls_dir): continue imgs [f for f in os.listdir(cls_dir) if f.lower().endswith((.jpg, .png, .jpeg))] print(f{cls}: {len(imgs)} 张图片)逻辑说明sorted保证类别顺序稳定避免每次运行标签编号乱跳。os.path.isdir过滤掉非文件夹项防止误读。图片后缀判断覆盖 jpg、png、jpeg 三种常见格式。参数上data_root要改成你解压后的实际路径如果数据集里还有子文件夹嵌套需要递归遍历这里只做了一层。3.2 生成训练集与验证集 txt 的步骤01数据集文本生成制作.py的核心产出是一个 txt 文件每行记录图片路径和对应标签。常见做法是按比例划分训练集和验证集比如 8:2。下面这段是我常用的划分逻辑和资源里脚本的思路一致。import os import random data_root dataset output_txt train_val.txt val_ratio 0.2 random.seed(42) # 固定随机种子保证每次划分一致 lines [] classes sorted(os.listdir(data_root)) for idx, cls in enumerate(classes): cls_dir os.path.join(data_root, cls) if not os.path.isdir(cls_dir): continue imgs [f for f in os.listdir(cls_dir) if f.lower().endswith((.jpg, .png, .jpeg))] random.shuffle(imgs) val_count int(len(imgs) * val_ratio) for i, img in enumerate(imgs): split val if i val_count else train lines.append(f{os.path.join(cls_dir, img)}\t{idx}\t{split}) with open(output_txt, w, encodingutf-8) as f: f.write(\n.join(lines)) print(已写入, len(lines), 条记录)逻辑说明random.seed(42)是血泪经验不固定种子的话每次划分不同验证集准确率会飘没法复现。val_count按比例算验证集数量split字段标记 train 或 val训练脚本读取时按这个字段分流。参数上val_ratio可以按数据集大小调数据少的时候 0.1 到 0.2 都行数据多可以到 0.3。标签用类别索引idx训练时配合CrossEntropyLoss使用。注意如果某个类别图片特别少按比例划分后验证集可能一张都没有这时候要么调低val_ratio要么保证每类至少留一张验证图。4. 模型训练与日志分析02 脚本里的 CNN 结构和训练参数4.1 一个够用的 CNN 分类网络长什么样这套资源的训练脚本02深度学习模型训练.py里应该定义了一个 CNN 网络。按大米识别这个任务的特点输入图片经过补灰边后是正方形尺寸通常会被缩放到 224 或 128。我一般会用一个三层卷积加全连接的结构够用且不容易过拟合。下面这段是我按常见做法写的网络定义你可以对照资源里的实现看差异。import torch import torch.nn as nn class RiceCNN(nn.Module): def __init__(self, num_classes): super().__init__() self.features nn.Sequential( nn.Conv2d(3, 32, 3, padding1), nn.ReLU(), nn.MaxPool2d(2), nn.Conv2d(32, 64, 3, padding1), nn.ReLU(), nn.MaxPool2d(2), nn.Conv2d(64, 128, 3, padding1), nn.ReLU(), nn.MaxPool2d(2), ) self.classifier nn.Sequential( nn.Flatten(), nn.Linear(128 * 28 * 28, 256), nn.ReLU(), nn.Dropout(0.5), nn.Linear(256, num_classes), ) def forward(self, x): return self.classifier(self.features(x))逻辑说明三层卷积逐步把通道数从 3 提到 128每次池化把空间尺寸减半。Dropout(0.5)是防过拟合的关键大米图片类间差异小不加 dropout 训练集准确率会飙到很高但验证集上不去。参数上num_classes按你数据集实际类别数传128 * 28 * 28这个维度要按输入图片尺寸算如果输入是 128 而不是 224这里要改成对应的数。4.2 训练循环与 log 日志里该看什么训练脚本跑起来后会在本地存 log记录每个 epoch 的验证集损失和准确率。我一般会重点看两个信号验证集损失是否在连续几个 epoch 后开始上升以及验证集准确率和训练集准确率的差距。前者说明过拟合后者说明泛化能力。下面这段是训练循环的核心骨架。import torch from torch.utils.data import DataLoader device torch.device(cuda if torch.cuda.is_available() else cpu) model RiceCNN(num_classeslen(classes)).to(device) criterion nn.CrossEntropyLoss() optimizer torch.optim.Adam(model.parameters(), lr1e-3) for epoch in range(30): model.train() for imgs, labels in train_loader: imgs, labels imgs.to(device), labels.to(device) optimizer.zero_grad() loss criterion(model(imgs), labels) loss.backward() optimizer.step() model.eval() correct, total 0, 0 with torch.no_grad(): for imgs, labels in val_loader: imgs, labels imgs.to(device), labels.to(device) preds model(imgs).argmax(dim1) correct (preds labels).sum().item() total labels.size(0) print(fepoch {epoch}, val acc {correct / total:.4f})逻辑说明model.train()和model.eval()切换影响 dropout 和 batchnorm 行为漏写会导致验证结果不准。optimizer.zero_grad()必须在反向传播前调用否则梯度会累加。参数上lr1e-3是 Adam 的常用起点如果损失不下降可以降到 1e-4如果下降太慢可以升到 3e-3。epoch数量按数据集大小调大米识别这种任务通常 20 到 50 轮够用。提示log 里如果验证集准确率一直在 0.2 到 0.3 徘徊先检查标签生成时类别索引和训练时是否一致这是最常见的翻车点。5. PyQt 界面与推理03 脚本怎么把模型变成可点的按钮5.1 界面加载模型与图片的流程03pyqt_ui界面.py的作用是提供一个可视化窗口点按钮加载图片然后显示识别结果。常见做法是在界面初始化时加载训练好的模型权重点击按钮后用QFileDialog选图片再用和训练时一致的预处理流程把图片转成张量送进模型。下面这段是界面核心逻辑的骨架。import sys import torch from PyQt5.QtWidgets import QApplication, QWidget, QPushButton, QLabel, QFileDialog, QVBoxLayout from PIL import Image from torchvision import transforms class RiceUI(QWidget): def __init__(self, model, classes): super().__init__() self.model model self.classes classes self.transform transforms.Compose([ transforms.Resize((224, 224)), transforms.ToTensor(), ]) self.label QLabel(等待加载图片) btn QPushButton(选择图片) btn.clicked.connect(self.load_image) layout QVBoxLayout() layout.addWidget(btn) layout.addWidget(self.label) self.setLayout(layout) def load_image(self): path, _ QFileDialog.getOpenFileName(self, 选择图片, , Images (*.jpg *.png)) if not path: return img Image.open(path).convert(RGB) tensor self.transform(img).unsqueeze(0) with torch.no_grad(): pred self.model(tensor).argmax(dim1).item() self.label.setText(f识别结果: {self.classes[pred]})逻辑说明transforms.Resize((224, 224))必须和训练时的输入尺寸一致不一致会导致识别结果乱跳。unsqueeze(0)是加 batch 维度模型 forward 要求四维输入。torch.no_grad()关闭梯度计算推理时省显存。参数上classes要和训练时sorted(os.listdir(data_root))的顺序完全一致否则标签会对错。5.2 界面与训练脚本的标签对齐问题界面里显示的类别名来自classes列表这个列表的顺序必须和训练时生成标签的顺序一致。我一般会把类别列表存成一个 json 文件训练和界面都读同一个文件避免两边各写一份导致对不上。如果你发现界面识别结果总是固定错到某一类先查这个顺序而不是怀疑模型没训练好。注意PyQt 界面在部分高分辨率屏幕上会出现控件错位这是 Qt 的缩放策略问题可以在QApplication创建前设置QApplication.setAttribute(Qt.AA_EnableHighDpiScaling)缓解。6. 避坑与排查这套资源跑不起来时先查这五条6.1 现象01 脚本报 FileNotFoundError原因data_root路径写的是相对路径但脚本运行目录和你解压目录不一致。解决把data_root改成绝对路径或者在脚本开头用os.chdir切到数据集所在目录。我一般直接在代码里写绝对路径省得来回调。6.2 现象训练时 loss 变成 nan原因学习率太大或者图片预处理时没有归一化导致输入值范围过大。解决把lr从 1e-3 降到 1e-4并在transforms.Compose里加transforms.Normalize(mean[0.5,0.5,0.5], std[0.5,0.5,0.5])。如果已经加了归一化还 nan检查标签里有没有超出类别数的索引。6.3 现象验证集准确率始终等于类别数的倒数原因标签生成时类别索引和训练时读取的索引不一致或者验证集图片路径写错导致读进来全是同一类。解决打印前 10 条 txt 记录确认路径和标签对应关系再打印train_loader里一个 batch 的标签分布。6.4 现象PyQt 界面点按钮没反应原因QFileDialog返回的路径为空时没有做判断或者模型加载失败但异常被吞了。解决在load_image里加if not path: return并在模型加载处用 try-except 打印异常。我一般会在界面启动时先打印一行「模型加载成功」确认。6.5 现象GPU 显存不足报 CUDA out of memory原因batch size 太大或者图片输入尺寸设得过高。解决把 batch size 从 32 降到 8 或 16或者把输入尺寸从 224 降到 128。如果还不行在训练循环里加torch.cuda.empty_cache()但根本办法还是减小 batch。7. 进阶技巧把验证集准确率从 0.85 推到 0.95 的几个实操手段这套资源跑通之后你大概率会看到验证集准确率在 0.85 到 0.9 之间。想再往上推我一般会从三个地方下手。第一是数据增强再加强资源里已经做了旋转和翻转你可以再加随机裁剪和颜色抖动让模型见过更多变化。第二是换优化器策略把固定学习率换成CosineAnnealingLR让学习率随 epoch 下降后期收敛更稳。第三是保存最佳模型而不是最后一个 epoch 的模型验证集准确率最高的时候存一次避免过拟合后的权重被用到界面上。下面这段是加余弦退火和最佳模型保存的改动直接嵌到训练循环里就行。from torch.optim.lr_scheduler import CosineAnnealingLR scheduler CosineAnnealingLR(optimizer, T_max30) best_acc 0.0 for epoch in range(30): # ... 训练和验证代码同上 ... scheduler.step() if val_acc best_acc: best_acc val_acc torch.save(model.state_dict(), best_model.pth) print(f已保存最佳模型准确率 {best_acc:.4f})逻辑说明CosineAnnealingLR的T_max设成总 epoch 数学习率会按余弦曲线从初始值降到接近 0。best_acc记录历史最高验证准确率只有超过时才覆盖保存。参数上T_max要和实际训练的 epoch 数一致设小了学习率提前降到 0设大了退火效果不明显。另外如果你发现某些类别之间总是互相混淆可以把混淆矩阵打出来看是哪两类在打架。常见做法是用sklearn.metrics.confusion_matrix把验证集的预测结果和真实标签传进去一眼就能看出问题类别。针对混淆的类别可以单独多补一些该角度的图片或者把输入尺寸从 128 提到 224让模型看到更多纹理细节。从那以后我每次跑这类图像分类项目都会先把类别列表存成 json训练和界面共用一份再固定随机种子最后一定保存最佳模型而不是最后一个 epoch。这三步做完复现性会好很多也不会出现界面识别结果和训练日志对不上的情况。希望帮到你。本文还有配套的精品资源点击获取
返回列表
PREV
查看更多资讯
NEXT
返回资讯列表