基于CNN深度学习的大米识别实战:从数据集到PyQt可视化界面

发布时间:2026/10/5 15:32:50
基于CNN深度学习的大米识别实战:从数据集到PyQt可视化界面
简介这份资源面向深度学习入门者与计算机视觉方向的开发者提供一套基于PyTorch框架的CNN大米图像识别完整实现方案可用于学习图像分类项目的全流程搭建。压缩包共906个文件包含900张jpg格式的大米类别图片、3个txt说明与日志文本以及3个py脚本整体约11.98MB体积轻便便于本地运行。代码对数据集做了预处理通过短边补灰边将图片统一为正方形并施加旋转角度来扩增增强样本覆盖数据读取、模型训练与可视化推理三个环节。运行01脚本可生成记录图片路径与标签的文本02脚本完成训练并保存模型与逐epoch验证损失、准确率日志03脚本则提供PyQt界面点击按钮即可加载图片进行识别。目前已有152人学习适合希望掌握CNN图像分类实战、理解数据增强与模型评估流程的读者参考。1. 从一堆文件名到可训练模型这套大米识别资源到底能跑出什么如果你手头正躺着一批按类别分文件夹的稻米图片文件名里还带着rotated45、flip这类后缀想快速验证一个 CNN 分类流程能不能跑通这套「基于 CNN 深度学习的大米识别-含图片数据集」就是冲这个场景来的。它把 Ipsala 等类别的大米图片按目录组织好配套三个脚本先扫数据集生成训练/验证用的 txt 清单再用 PyTorch 训练 CNN 并保存模型和日志最后用 PyQt 拉起一个可视化界面点按钮就能加载图片看识别结果。适合刚接触深度学习 CNN、想拿一个完整闭环练手的人也适合需要快速搭图像分类 demo 的从业者。它不追求 SOTA 精度胜在流程完整、依赖清晰、每一步都能自己复现。2. 环境与数据准备requirement.txt 之外你还要盯住什么2.1 为什么这套代码对 PyTorch 版本比较敏感这套代码的核心是torch、torchvision加PyQt5训练脚本里大概率用了transforms、DataLoader、nn.Conv2d这些标准件。问题在于PyTorch 从 1.x 到 2.x 在部分 API 上有过调整比如torch.load的weights_only默认值变化、transforms.Normalize的调用方式没变但周边工具链容易出岔子。如果你直接pip install torch拉最新版训练脚本里如果有旧写法可能报TypeError或AttributeError。常见做法是打开requirement.txt看它有没有钉版本号如果没钉我一般会选一个 LTS 性质的组合比如 Python 3.83.10 配 PyTorch 1.122.0这个区间对大多数教学级 CNN 代码兼容性最好。另一个容易被忽略的是 CUDA 版本。训练脚本如果写了.cuda()或device cuda而你机器上没有对应 CUDA 驱动就会直接抛错。没有 GPU 也能跑但要在代码里把设备改成 CPU或者用torch.cuda.is_available()做判断。数据集是图片Ipsala 类别单类图片数量不算大CPU 训练虽然慢但跑通流程没问题。2.2 数据集目录结构与 01 脚本的读取逻辑数据集文件夹按类别分子目录每个子目录里是图片。01 脚本的任务就是遍历这些子目录把每张图片的路径和对应标签写进 txt。它的读取逻辑通常是先拿到类别列表给每个类别分配一个数字标签然后对每个类别目录下的图片逐个拼路径。这里有个细节文件名里带rotated45、flip的图片是原始图片经过旋转和翻转增强后的产物它们和原图在同一个类别目录下所以标签是一致的。01 脚本不需要区分这些后缀它只认目录名。你可以先手动确认一下目录层级常见结构是dataset/ ├── Ipsala/ │ ├── 10074_rotated45.jpg │ ├── 10074_flip.jpg │ ├── 10031_rotated45.jpg │ └── ... ├── 类别B/ │ └── ... └── 类别C/ └── ...如果目录层级和脚本预期不一致比如多了一层或类别名有中文01 脚本可能读不到图片或标签错乱。我一般会先跑一个最小检查import os root dataset for cls in os.listdir(root): cls_dir os.path.join(root, cls) if os.path.isdir(cls_dir): imgs [f for f in os.listdir(cls_dir) if f.lower().endswith((.jpg, .png, .jpeg))] print(cls, len(imgs))这段代码只做一件事确认每个类别目录下到底有多少张可读图片。如果某个类别输出 0说明扩展名不匹配或路径写错了。参数root要改成你本地数据集的实际路径endswith里把常见图片格式都列上避免漏掉。2.3 安装依赖时容易翻车的两个点第一个点是 PyQt5 和 PyTorch 的安装顺序。如果你先装 PyQt5 再装 PyTorch一般没事但如果你用 conda 装 PyTorch、用 pip 装 PyQt5偶尔会出现 Qt 插件路径冲突表现为运行 03 脚本时弹窗报Could not find the Qt platform plugin windows。解决办法是统一用 pip 或统一用 conda或者手动设置QT_QPA_PLATFORM_PLUGIN_PATH环境变量指向 PyQt5 的 plugins 目录。第二个点是requirement.txt里可能列了opencv-python而 01 脚本如果用cv2.imread读图遇到中文路径会返回 None。数据集路径里如果有中文要么改成英文要么用cv2.imdecode(np.fromfile(path, dtypenp.uint8), -1)替代。这个坑在 Windows 上尤其常见血泪经验是只要路径带中文先怀疑编码问题。提示环境装完后先跑python -c import torch; print(torch.__version__); import cv2; print(cv2.__version__)确认两个核心库都能正常导入再往下走。3. 训练脚本拆解02 脚本里 CNN 结构、数据增强与日志到底怎么读3.1 数据预处理灰边填充与旋转增强的实际作用项目正文提到代码对数据集做了预处理包括在较短边增加灰边使图片变为正方形以及旋转角度来扩增增强数据集。这个操作的目的很直接CNN 的输入通常要求固定尺寸比如 224×224 或 128×128。如果原始图片是长方形直接 resize 会拉伸变形影响纹理特征。加灰边变成正方形后再 resize能保留原始比例灰边本身不携带类别信息不会干扰分类。旋转增强则是针对大米颗粒的方向不变性。同一类大米横着拍和斜着拍应该被识别为同一类旋转 45 度生成新样本相当于告诉模型「方向不重要」。文件名里的rotated45和flip就是这些增强结果的标记。你在 01 脚本生成的 txt 里会看到它们和原图混在一起训练时会被随机分到训练集或验证集。这里有个参数要注意旋转角度如果太大比如 90 度对于某些本身有方向性特征的类别可能反而引入噪声。45 度是一个折中值。翻转同理水平翻转对大米通常安全垂直翻转要看具体场景。3.2 02 脚本的训练流程与关键参数02 脚本的典型流程是读 txt 清单 → 定义 Dataset 类 → 定义 CNN 模型 → 定义损失函数和优化器 → 循环 epoch 训练 → 每个 epoch 后在验证集上评估 → 保存模型和日志。下面是一个和该项目结构对齐的代码骨架你可以对照自己的 02 脚本看差异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 RiceDataset(Dataset): def __init__(self, txt_path, transformNone): self.samples [] self.transform transform with open(txt_path, r, encodingutf-8) as f: for line in f: line line.strip() if not line: continue path, label line.rsplit( , 1) # 假设 txt 每行是 路径 标签 self.samples.append((path, int(label))) def __len__(self): return len(self.samples) def __getitem__(self, idx): path, label self.samples[idx] img Image.open(path).convert(RGB) if self.transform: img self.transform(img) return img, label train_tf transforms.Compose([ transforms.Resize((128, 128)), transforms.ToTensor(), transforms.Normalize(mean[0.5, 0.5, 0.5], std[0.5, 0.5, 0.5]) ]) train_ds RiceDataset(train.txt, transformtrain_tf) train_loader DataLoader(train_ds, batch_size32, shuffleTrue, num_workers0) class SimpleCNN(nn.Module): def __init__(self, num_classes): super().__init__() self.features nn.Sequential( nn.Conv2d(3, 16, 3, padding1), nn.ReLU(), nn.MaxPool2d(2), nn.Conv2d(16, 32, 3, padding1), nn.ReLU(), nn.MaxPool2d(2), nn.Conv2d(32, 64, 3, padding1), nn.ReLU(), nn.MaxPool2d(2) ) self.classifier nn.Sequential( nn.Flatten(), nn.Linear(64 * 16 * 16, 128), nn.ReLU(), nn.Linear(128, num_classes) ) def forward(self, x): return self.classifier(self.features(x)) model SimpleCNN(num_classes5) # 类别数按实际改 criterion nn.CrossEntropyLoss() optimizer optim.Adam(model.parameters(), lr1e-3) for epoch in range(20): model.train() total_loss 0 for imgs, labels in train_loader: optimizer.zero_grad() outputs model(imgs) loss criterion(outputs, labels) loss.backward() optimizer.step() total_loss loss.item() print(fepoch {epoch1}, loss {total_loss/len(train_loader):.4f})这段代码里RiceDataset负责按 txt 读路径和标签rsplit( , 1)是为了兼容路径里可能带空格的情况。Resize((128, 128))是输入尺寸你可以改成 224但要注意全连接层的输入维度也要跟着改。batch_size32在 CPU 上可能偏大如果内存吃紧就降到 8 或 16。num_workers0在 Windows 上比较稳设大了有时会卡死。lr1e-3是 Adam 的常见起点训练不收敛就降到 1e-4。3.3 日志里的验证集损失和准确率怎么看训练完成后本地会保存 log 日志里面记录了每个 epoch 的验证集损失值和准确率。这两个指标要结合起来看训练损失下降但验证损失上升是过拟合的典型信号说明模型记住了训练集的噪声两者都下降但准确率卡在某个值不动可能是学习率太大或模型容量不够验证准确率波动大可能是验证集太小或 batch 里类别分布不均。我一般会先把 log 里的数字拉出来画个曲线不用 matplotlib 也行直接看最后几个 epoch 的验证准确率是否稳定。如果 20 个 epoch 后验证准确率还在 60% 以下先检查标签有没有错位再检查图片有没有读成黑图。标签错位的排查方法是从 txt 里抽几条手动打开图片看类别是否和标签一致。注意log 文件如果是在训练脚本里用open(..., a)追加写的重复运行会叠加历史记录。每次重新训练前先删掉旧 log或者改成带时间戳的文件名。4. PyQt 界面与推理03 脚本怎么把模型变成可点的按钮4.1 03 脚本的加载逻辑与模型路径03 脚本做三件事加载训练好的模型权重、构建一个带按钮的窗口、点击按钮后选图片并显示识别结果。模型加载部分通常长这样import torch from torchvision import transforms from PIL import Image from PyQt5.QtWidgets import QApplication, QWidget, QPushButton, QLabel, QFileDialog, QVBoxLayout from PyQt5.QtGui import QPixmap class RiceUI(QWidget): def __init__(self, model_path, class_names): super().__init__() self.class_names class_names self.model torch.load(model_path, map_locationcpu) self.model.eval() self.tf transforms.Compose([ transforms.Resize((128, 128)), transforms.ToTensor(), transforms.Normalize([0.5]*3, [0.5]*3) ]) self.label QLabel(等待加载图片) self.btn QPushButton(选择图片) self.btn.clicked.connect(self.load_image) layout QVBoxLayout() layout.addWidget(self.label) layout.addWidget(self.btn) self.setLayout(layout) def load_image(self): path, _ QFileDialog.getOpenFileName(self, 选图片, , Images (*.jpg *.png)) if not path: return img Image.open(path).convert(RGB) x self.tf(img).unsqueeze(0) with torch.no_grad(): out self.model(x) pred out.argmax(1).item() self.label.setText(f识别结果{self.class_names[pred]}) self.label.setPixmap(QPixmap(path).scaled(300, 300))torch.load(model_path, map_locationcpu)里的map_location很关键如果你在 GPU 上训练、在 CPU 上推理不加这个参数会报错。class_names要和训练时的标签顺序一致否则识别结果会张冠李戴。unsqueeze(0)是给图片加一个 batch 维度因为模型期望输入是(N, C, H, W)。4.2 界面卡顿与图片格式的常见问题PyQt 界面在加载大图时可能卡顿因为QPixmap直接加载原图会占内存。解决办法是在显示前先缩放到固定尺寸比如QPixmap(path).scaled(300, 300)。另外如果图片是灰度图或带透明通道Image.open(path).convert(RGB)能统一成三通道避免ToTensor后维度不对。还有一个坑是模型保存方式。如果 02 脚本保存的是整个模型对象torch.save(model, path)03 脚本加载时需要能 import 到模型类定义否则会报AttributeError。更稳的做法是只保存state_dict加载时先实例化模型再load_state_dict。如果你拿到的代码是保存整个对象就把模型类定义放在 03 脚本能访问到的地方或者改成 state_dict 方式。4.3 从点击按钮到输出标签的完整链路完整链路是用户点按钮 →QFileDialog返回路径 →PIL读图并转 RGB →transforms做 resize、ToTensor、Normalize →unsqueeze加 batch → 模型前向 →argmax取最大概率索引 → 用class_names映射成类别名 → 更新QLabel文本和图片。每一步的参数都要和训练时对齐尤其是Resize尺寸和Normalize的 mean/std。如果训练时用了mean[0.5,0.5,0.5]推理时也必须一样否则输入分布偏移准确率会掉。5. 避坑与排查这套代码跑不起来时先查这五条5.1 现象01 脚本生成的 txt 是空的原因数据集路径写错或者脚本里的root指向了上一级目录导致os.listdir找不到类别文件夹。解决在 01 脚本开头打印os.path.abspath(root)确认路径存在再打印os.listdir(root)看是否包含类别目录名。5.2 现象训练时 loss 一直是 nan原因学习率太大或者输入图片没有归一化像素值 0-255 直接进网络导致梯度爆炸。解决把lr降到 1e-4 或 1e-5确认transforms.ToTensor()在Normalize之前Normalize的 mean/std 按实际数据调整。5.3 现象03 脚本报ModuleNotFoundError: No module named PyQt5原因PyQt5 没装或者装到了另一个 Python 环境里。解决pip install PyQt5然后用python -c import PyQt5; print(PyQt5.__file__)确认路径和你运行 03 脚本的 Python 是同一个。5.4 现象验证集准确率远高于训练集准确率原因验证集太小或者验证集和训练集有重叠图片。解决检查 01 脚本划分训练/验证的比例确保同一张图片不会同时出现在两个集合里如果数据集本身小可以增大验证集比例或做交叉验证。5.5 现象界面能弹出但点按钮没反应原因btn.clicked.connect(self.load_image)里的self.load_image没绑定成功或者load_image内部异常被吞了。解决在load_image开头加print(clicked)看是否触发再用try/except把异常打印出来定位是读图失败还是模型推理失败。6. 进阶技巧把单张推理改成批量验证并核对混淆矩阵跑通单张图片识别之后我习惯做一件事用验证集做一次批量推理生成混淆矩阵。这一步能暴露很多单张测试看不出来的问题比如某个类别总是被误判成另一个类别或者某个类别的召回率特别低。具体做法是读验证集 txt逐张推理把真实标签和预测标签存下来再用 sklearn 的confusion_matrix输出。import torch from torchvision import transforms from PIL import Image from sklearn.metrics import confusion_matrix, classification_report model torch.load(model.pth, map_locationcpu) model.eval() tf transforms.Compose([ transforms.Resize((128, 128)), transforms.ToTensor(), transforms.Normalize([0.5]*3, [0.5]*3) ]) y_true, y_pred [], [] with open(val.txt, r, encodingutf-8) as f: for line in f: line line.strip() if not line: continue path, label line.rsplit( , 1) img Image.open(path).convert(RGB) x tf(img).unsqueeze(0) with torch.no_grad(): pred model(x).argmax(1).item() y_true.append(int(label)) y_pred.append(pred) print(confusion_matrix(y_true, y_pred)) print(classification_report(y_true, y_pred))这段代码的关键参数是val.txt的路径和模型路径要和你本地实际文件对应。classification_report会输出每个类别的 precision、recall、f1-score比单纯看准确率更有信息量。如果某个类别的 recall 明显低就去看看这个类别的图片是不是数量太少或者增强后的图片质量有问题。我一般还会把误判的图片路径单独存下来人工翻一遍看看是不是标签标错了。数据集里带rotated45和flip的图片多了之后偶尔会出现某张增强图旋转后主体跑出画面这种图训练时就是噪声。从那以后我每次拿到带增强后缀的数据集都强制先抽样看一批增强图确认没有无效样本再开训。希望帮到你。本文还有配套的精品资源点击获取