ARTICLE DETAIL

资讯详情

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

基于CNN深度学习的大米识别实战:含图片数据集与PyTorch代码

基于CNN深度学习的大米识别实战:含图片数据集与PyTorch代码 简介本资源是一套基于PyTorch框架的CNN大米图像识别完整项目面向具备一定Python基础、希望入门深度学习图像分类的开发者与在校学生。项目围绕大米类别识别任务提供从数据预处理、模型训练到可视化交互的完整链路适合作为课程设计、毕业设计或练手项目的参考方案。压缩包共906个文件包含900张jpg图片构成的多类别数据集、3个txt说明与标签文件以及3个py脚本整体约11.98MB体积轻便易于本地运行。数据集已做短边补灰边成正方形与旋转角度等增强处理脚本依次完成标签文本生成、模型训练与PyQt界面调用训练过程会保存模型权重与逐epoch验证损失、准确率日志界面支持加载任意图片进行识别。目前已有152人学习适合想快速跑通CNN分类全流程并理解数据增强与训练监控细节的读者。1. 大米识别为什么值得用 CNN 做一遍从一张米粒图说起把一粒大米放在白纸上拍张照人眼能轻松分辨它是长粒香、珍珠米还是糯米。但换成工业分选线上的高速相机每秒几百帧、几万粒米同时过检靠人眼就彻底歇了。这正是「基于 CNN 深度学习的大米识别」要解决的问题用卷积神经网络对米粒图像做自动分类把品种、完整度、垩白度这些指标从像素里抠出来。它适合三类人——做农产品质检自动化的工程师、拿深度学习图像识别当毕设或实战项目练手的学生、以及手里已经有一批米粒图片、想跑通一个端到端分类流程的开发者。标题里那个「含图片数据集」是关键它意味着你不用从零去田里拍米拿到手就能直接进训练环节。但数据集能省掉采集的力气省不掉对 CNN 结构、预处理和调参的理解这篇就把这条链路从头到尾拆开讲清楚。2. 先搞懂 CNN 在大米识别里到底干了什么2.1 卷积、池化、全连接米粒特征是怎么被一层层抽出来的大米图像分类和通用图像分类在原理上没有本质区别都是让网络自己学会「什么样的纹理和形状组合对应哪个品种」。区别在于米粒的判别特征非常细长宽比、腹白面积、胚芽位置、表面裂纹这些在低分辨率下很容易糊成一团。CNN 的价值就在于它不需要你手工设计这些特征卷积核在图像上滑动时会先捕捉边缘和角点再往上组合成纹理块最后在全连接层形成类别判断。一个典型的卷积层做的是局部加权求和。假设输入是一张 224×224 的米粒图第一层用 3×3 的卷积核、步长 1、padding 1输出还是 224×224但通道数从 3 变成 32。这一步的意义是把原始 RGB 像素映射成 32 种「边缘响应图」。池化层紧接着把空间尺寸砍半保留最强响应降低计算量同时带来一点平移不变性——米粒在画面里稍微偏一点分类结果不该变。我一般会用一个轻量结构起步而不是直接上 ResNet。原因是米粒数据集通常只有几千到几万张太深的网络容易过拟合训练也慢。下面是一个可以直接跑的最小 CNN 定义用 PyTorch 写import torch import torch.nn as nn class RiceCNN(nn.Module): def __init__(self, num_classes5): super().__init__() # 第一组3通道输入 - 32通道适合捕捉米粒边缘 self.conv1 nn.Conv2d(3, 32, kernel_size3, padding1) self.bn1 nn.BatchNorm2d(32) # 批归一化稳定训练 self.conv2 nn.Conv2d(32, 32, kernel_size3, padding1) self.bn2 nn.BatchNorm2d(32) self.pool1 nn.MaxPool2d(2) # 224 - 112 # 第二组32 - 64通道捕捉纹理组合 self.conv3 nn.Conv2d(32, 64, kernel_size3, padding1) self.bn3 nn.BatchNorm2d(64) self.conv4 nn.Conv2d(64, 64, kernel_size3, padding1) self.bn4 nn.BatchNorm2d(64) self.pool2 nn.MaxPool2d(2) # 112 - 56 # 第三组64 - 128通道 self.conv5 nn.Conv2d(64, 128, kernel_size3, padding1) self.bn5 nn.BatchNorm2d(128) self.pool3 nn.AdaptiveAvgPool2d(1) # 全局平均池化输出 128x1x1 self.fc nn.Linear(128, num_classes) self.relu nn.ReLU() self.dropout nn.Dropout(0.5) # 防过拟合 def forward(self, x): x self.relu(self.bn1(self.conv1(x))) x self.relu(self.bn2(self.conv2(x))) x self.pool1(x) x self.relu(self.bn3(self.conv3(x))) x self.relu(self.bn4(self.conv4(x))) x self.pool2(x) x self.relu(self.bn5(self.conv5(x))) x self.pool3(x) x x.view(x.size(0), -1) # 展平 x self.dropout(x) return self.fc(x)这段代码里几个参数值得说清楚。num_classes要按你数据集里实际的大米类别数改常见的是 3 到 10 类。BatchNorm2d放在卷积和 ReLU 之间能让训练初期收敛快很多尤其是米粒图像亮度不均的时候。AdaptiveAvgPool2d(1)替代了传统的展平接大全连接层参数量小、抗过拟合对中小数据集更友好。Dropout(0.5)只在训练时生效推理时会自动关闭。2.2 为什么选 CNN 而不是传统机器学习或 Transformer传统机器学习做米粒分类典型流程是先做阈值分割把每粒米抠出来再手工量长宽、面积、周长、颜色矩最后丢给 SVM 或随机森林。这套方法在背景干净、光照稳定的产线上能跑但一旦米粒粘连、有阴影、品种间差异细微手工特征就崩了。CNN 的优势是端到端你给整张图它自己学该看哪里。那为什么不直接上 Transformer视觉 Transformer 在超大数据集上确实强但米粒数据集规模通常撑不起它的数据胃口训练成本也高。CNN 的归纳偏置——局部性、平移不变性——恰好匹配米粒这种「局部纹理决定类别」的任务。热搜里常出现「transformer和cnnrnn的区别」放到这个场景里结论很直接RNN 处理序列不适合图像Transformer 数据需求大CNN 是性价比最高的选择。2.3 数据集拿到手先做这三件事标题里带了图片数据集但别急着往网络里灌。我一般先做三件事。第一统计每个类别的图片数量如果某类只有几十张后面必然类别不平衡。第二抽查图片尺寸和通道有的数据集混了灰度图和 RGBA 图直接进网络会报错。第三可视化一批样本看有没有标错类、有没有几乎一样的重复图。这三步花不了半小时能省掉后面几小时的玄学报错。import os from PIL import Image from collections import Counter data_dir rice_dataset/train counter Counter() size_set set() mode_set set() for cls in os.listdir(data_dir): cls_dir os.path.join(data_dir, cls) if not os.path.isdir(cls_dir): continue for fname in os.listdir(cls_dir): path os.path.join(cls_dir, fname) try: img Image.open(path) counter[cls] 1 size_set.add(img.size) mode_set.add(img.mode) except Exception as e: print(坏图:, path, e) print(类别分布:, counter) print(出现的尺寸:, size_set) print(出现的色彩模式:, mode_set)Counter输出类别分布一眼看出平不平衡。size_set如果不止一个值说明图片尺寸不统一后面必须统一 resize。mode_set里如果出现L或RGBA要在 Dataset 里统一转成RGB否则Conv2d(3, ...)会直接报通道数不匹配。这个检查脚本我每次拿到新数据集都会跑一遍属于血泪经验。3. 把图片数据集喂进 CNN预处理、划分与训练循环3.1 图像预处理与数据增强的必调参数米粒图像预处理的核心目标是让网络看到的每一批图尺寸一致、亮度可比、且不因为拍摄角度和光照差异而学偏。必做的有 resize、归一化、转 RGB。增强则用来在小数据集上「造」出更多样本常用的是随机水平翻转、随机旋转小角度、颜色抖动。但增强不是越多越好。米粒的长宽比是重要判别特征如果你做随机拉伸或大角度旋转可能把长粒米变成短粒米的形状反而教坏网络。我的习惯是水平翻转可以开垂直翻转慎用米粒上下有别旋转限制在 ±15 度颜色抖动幅度调小。from torchvision import transforms train_tf transforms.Compose([ transforms.Resize((224, 224)), # 统一尺寸 transforms.RandomHorizontalFlip(p0.5), # 水平翻转 transforms.RandomRotation(15), # 小角度旋转 transforms.ColorJitter(brightness0.2, contrast0.2), # 轻微颜色扰动 transforms.ToTensor(), # 转张量像素归到[0,1] transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]) # ImageNet统计值 ]) 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]) ])Resize((224, 224))里的 224 是常见起点如果你的米粒在原图里占比很小可以先裁剪再 resize否则米粒被缩得太小细节全丢。Normalize用的均值方差是 ImageNet 的统计值如果你从零训练、不用预训练权重也可以换成自己数据集的均值方差但对结果影响不大。验证集绝对不能用增强否则评估结果会虚高这是新手最容易翻车的地方之一。3.2 训练集/验证集/测试集怎么分才不泄漏划分数据集有个隐蔽的坑如果同一个米粒被拍了多张相似照片随机划分会把「近重复图」同时分进训练和验证导致验证准确率虚高上线就露馅。稳妥做法是先按「拍摄批次」或「原始大图」分组再在组级别划分。如果数据集没有批次信息至少要用固定随机种子保证划分可复现。常见比例是 7:1.5:1.5 或 8:1:1。类别不平衡时用分层抽样让每个集合里各类比例接近。import random from torch.utils.data import DataLoader, Subset from torchvision.datasets import ImageFolder random.seed(42) # 固定种子保证可复现 full_ds ImageFolder(rice_dataset/train, transformtrain_tf) targets [s[1] for s in full_ds.samples] # 按类别分组索引 from collections import defaultdict cls_idx defaultdict(list) for i, t in enumerate(targets): cls_idx[t].append(i) train_idx, val_idx, test_idx [], [], [] for t, idxs in cls_idx.items(): random.shuffle(idxs) n len(idxs) n_train int(n * 0.7) n_val int(n * 0.15) train_idx idxs[:n_train] val_idx idxs[n_train:n_train n_val] test_idx idxs[n_train n_val:] train_ds Subset(full_ds, train_idx) val_ds Subset(ImageFolder(rice_dataset/train, transformval_tf), val_idx) test_ds Subset(ImageFolder(rice_dataset/train, transformval_tf), test_idx) train_loader DataLoader(train_ds, batch_size32, shuffleTrue, num_workers4) val_loader DataLoader(val_ds, batch_size32, shuffleFalse, num_workers4) test_loader DataLoader(test_ds, batch_size32, shuffleFalse, num_workers4)这里按类别分别划分保证每类在三个集合里都有。batch_size32是显存和稳定性的折中显存小就降到 16 或 8。num_workers4在 Linux 上加速数据加载Windows 上如果报错就改成 0。注意验证和测试用的是val_tf没有增强。3.3 训练循环、学习率与早停的实操写法训练循环本身不复杂难的是监控和止损。我一般会记录每个 epoch 的训练损失、验证损失、验证准确率一旦验证损失连续几个 epoch 不降反升就触发早停同时保存验证准确率最高的那版权重。import torch.optim as optim device torch.device(cuda if torch.cuda.is_available() else cpu) model RiceCNN(num_classeslen(full_ds.classes)).to(device) criterion nn.CrossEntropyLoss() optimizer optim.Adam(model.parameters(), lr1e-3, weight_decay1e-4) scheduler optim.lr_scheduler.ReduceLROnPlateau( optimizer, modemin, factor0.5, patience3 ) best_acc 0.0 patience_counter 0 early_stop_patience 7 for epoch in range(50): model.train() running_loss 0.0 for imgs, labels in train_loader: imgs, labels imgs.to(device), labels.to(device) optimizer.zero_grad() outputs model(imgs) loss criterion(outputs, labels) loss.backward() optimizer.step() running_loss loss.item() * imgs.size(0) train_loss running_loss / len(train_ds) # 验证 model.eval() correct, total, val_loss 0, 0, 0.0 with torch.no_grad(): for imgs, labels in val_loader: imgs, labels imgs.to(device), labels.to(device) outputs model(imgs) loss criterion(outputs, labels) val_loss loss.item() * imgs.size(0) preds outputs.argmax(dim1) correct (preds labels).sum().item() total labels.size(0) val_loss / len(val_ds) val_acc correct / total scheduler.step(val_loss) print(fEpoch {epoch1}: train_loss{train_loss:.4f} fval_loss{val_loss:.4f} val_acc{val_acc:.4f}) if val_acc best_acc: best_acc val_acc torch.save(model.state_dict(), best_rice_cnn.pth) patience_counter 0 else: patience_counter 1 if patience_counter early_stop_patience: print(早停触发停止训练) breaklr1e-3是 Adam 的常用起点如果损失一开始就震荡降到 1e-4。weight_decay1e-4是 L2 正则抑制过拟合。ReduceLROnPlateau在验证损失停滞时把学习率减半比手动调省心。early_stop_patience7意味着连续 7 个 epoch 没刷新最佳准确率就停避免白跑。保存的是验证集上最好的权重不是最后一个 epoch 的这点很关键。4. 大米识别模型调优与评估别只看准确率4.1 混淆矩阵告诉你哪两类米在互相误判准确率是个笼统指标。如果 5 类米里有 4 类都识别得很好只有两类长得很像总混整体准确率可能还有 85%但实际业务里这两类的误判就是致命的。混淆矩阵能直接暴露这个问题。from sklearn.metrics import confusion_matrix, classification_report import numpy as np model.load_state_dict(torch.load(best_rice_cnn.pth)) model.eval() all_preds, all_labels [], [] with torch.no_grad(): for imgs, labels in test_loader: imgs imgs.to(device) outputs model(imgs) preds outputs.argmax(dim1).cpu().numpy() all_preds.extend(preds) all_labels.extend(labels.numpy()) cm confusion_matrix(all_labels, all_preds) print(混淆矩阵:\n, cm) print(classification_report(all_labels, all_preds, target_namesfull_ds.classes))classification_report会给出每一类的 precision、recall、f1-score。如果某一类 recall 特别低说明这类米大量被误判成别的类回去看这类样本是不是太少、或者和某个类视觉上太接近。解决办法可以是补样本、加类别权重、或者针对这两类做更细的预处理。4.2 类别不平衡时用加权损失和重采样米粒数据集里某些品种可能天然就少。这时候不加处理网络会倾向于预测多数类少数类 recall 惨不忍睹。两种常用手段给损失函数加类别权重或者对少数类过采样。# 按类别频率计算权重频率越低权重越高 class_counts np.bincount(all_labels) class_weights 1.0 / (class_counts 1e-6) class_weights class_weights / class_weights.sum() * len(class_counts) class_weights torch.tensor(class_weights, dtypetorch.float).to(device) criterion nn.CrossEntropyLoss(weightclass_weights)class_weights让少数类的损失被放大网络会更在意它们。注意权重别设得太极端否则多数类又被牺牲。另一种做法是用WeightedRandomSampler在采样阶段就让每个 batch 里各类比例均衡效果类似看个人习惯。4.3 用预训练权重把准确率再抬一截从零训练一个小 CNN在几千张米粒图上通常能到 80% 到 90%。如果想再往上走最省力的办法是拿 ImageNet 预训练的主干网络做迁移学习只换掉最后的分类头。热搜里「深度学习模型」「深度学习实战项目案例」经常提到迁移学习放到大米识别上确实好用。import torchvision.models as models model models.resnet18(weightsmodels.ResNet18_Weights.DEFAULT) # 冻结前面的层只训练最后的全连接 for param in model.parameters(): param.requires_grad False model.fc nn.Linear(model.fc.in_features, len(full_ds.classes)) model model.to(device) # 只优化 fc 层 optimizer optim.Adam(model.fc.parameters(), lr1e-3)先冻结主干只训分类头几个 epoch 后如果验证准确率上来了再解冻最后几个 block 做微调学习率调小到 1e-4。这样既利用了预训练特征又不会因为数据少而把主干带偏。注意weightsmodels.ResNet18_Weights.DEFAULT会自动下载权重离线环境要提前准备好。5. 避坑与排查大米识别训练里最容易翻车的五件事5.1 现象训练准确率 99%测试准确率 60%原因典型过拟合或者训练集和测试集划分时发生了数据泄漏比如同一粒米的近重复图同时进了两边。解决先检查划分逻辑按原始图或批次分组再加强正则提高 dropout、加 weight_decay、减少模型容量最后看训练集是不是太小考虑迁移学习。5.2 现象损失一直是 nan原因学习率太大、输入没归一化、或者某张图数据损坏。解决先把学习率降到 1e-4 试确认ToTensor()和Normalize都在用前面那个检查脚本扫一遍坏图。我遇到过一张全黑的图导致整个 batch 梯度爆炸删掉就好了。5.3 现象验证准确率剧烈震荡一会儿 90% 一会儿 70%原因batch_size 太小、验证集样本太少、或者学习率偏高。解决把 batch_size 提到 32 或 64验证集每类至少留几十张用ReduceLROnPlateau让学习率自动降。震荡严重时也可以对验证结果做滑动平均再判断。5.4 现象报错 Expected 3 channels but got 1原因数据集里混了灰度图Image.open出来是单通道而网络第一层要 3 通道。解决在 Dataset 的__getitem__里统一img img.convert(RGB)或者在预处理 transform 里加一步转换。这个坑几乎每个做图像分类的人都踩过。5.5 现象GPU 显存够但训练特别慢原因num_workers设成 0数据加载成了瓶颈或者每张图都实时做复杂增强。解决Linux 上把num_workers设成 CPU 核数的 2 到 4 倍Windows 上如果多进程报错就退回 0把 resize 后的图缓存成小图减少重复解码开销。6. 把模型用起来单张米粒图推理与批量质检的落地技巧训练完拿到best_rice_cnn.pth真正的价值在于能不能稳定地对新图做判断。单张推理的写法很直接但有几个细节决定它能不能上产线。from PIL import Image def predict(image_path, model, class_names, device): model.eval() img Image.open(image_path).convert(RGB) # 强制三通道 tensor val_tf(img).unsqueeze(0).to(device) # 加 batch 维度 with torch.no_grad(): logits model(tensor) probs torch.softmax(logits, dim1)[0] conf, pred probs.max(dim0) return class_names[pred.item()], conf.item() label, confidence predict(test_rice.jpg, model, full_ds.classes, device) print(f预测: {label}, 置信度: {confidence:.4f})convert(RGB)是必须的产线相机如果出灰度图不转就会报通道错。unsqueeze(0)补上 batch 维度因为网络期望输入是[N, C, H, W]。softmax把 logits 变成概率置信度低于某个阈值时应该拒绝判断、转人工复检而不是硬给一个结果。这个阈值我一般设在 0.7 到 0.8具体看业务对误判的容忍度。批量质检时别一张张读用 DataLoader 批量推理能快好几倍。另外产线光照会漂移模型上线前最好用现场新拍的图做一次验证如果准确率掉得厉害说明训练集和现场分布不一致得补现场样本重新微调。这一步没有捷径属于上线前的后悔药。一个我踩过的坑推理时忘了model.eval()dropout 和 batchnorm 还在训练模式同一张图两次预测结果都不一样排查了半天才想起来。现在我的习惯是只要不是训练第一行就写model.eval()第二行写torch.no_grad()。这套流程从数据检查、训练、评估到推理跑通一遍之后换成其他农产品图像分类也能直接套用。希望帮到你。本文还有配套的精品资源点击获取
返回列表