尧图网站设计 尧图网站设计YAOTU DESIGN
ARTICLE DETAIL

资讯详情

深耕网站设计与一线实操的经验洞察。

基于CNN的大米识别实战:数据集处理、模型训练与产线部署

基于CNN的大米识别实战:数据集处理、模型训练与产线部署 简介本资源是一套基于PyTorch框架的CNN深度学习大米识别实战项目面向具备Python基础、希望入门图像分类的开发者与在校学生可用于课程设计、毕业项目或算法练手。压缩包共906个文件包含900张jpg图片构成的多类别大米数据集以及3个py脚本和3个txt说明文件整体约11.98MB体积轻便便于本地运行。项目对数据做了较完整的预处理通过短边补灰边将图片统一为正方形并叠加旋转、翻转等操作扩增样本脚本依次完成数据集文本生成、模型训练与PyQt可视化界面调用训练过程会输出日志记录每个epoch的验证集损失与准确率并保存本地模型权重。已有152人学习读者可借此掌握从数据增强、模型训练到界面交互的完整图像分类流程并直接加载自选图片进行识别验证。1. 大米识别为什么值得用 CNN 做一遍从一张米粒图说起把一粒大米放在白纸上用手机拍一张照片人眼能轻松分辨它是长粒香、珍珠米还是糯米。但换成流水线上的工业相机每秒过几百粒米还要把碎米、黄粒米、异品种米自动剔出来靠人眼盯屏幕就不现实了。基于 CNN 深度学习的大米识别解决的正是这个场景输入一张米粒图像输出它的品种或等级分类。它适合三类人——做农产品分选设备的一线工程师、带学生做深度学习图像识别毕设的高校老师以及想找一个完整数据集把 CNN 分类流程跑通的自学者。这个方向门槛不高但要把准确率从 90% 推到 98% 以上坑都在数据集的细节里。下面按「数据集怎么用 → 模型怎么搭 → 训练怎么调 → 坑在哪 → 怎么验证」的顺序把这条链路讲透。2. 大米图片数据集怎么读目录结构、类别分布与预处理2.1 拿到 zip 后先做三件事解压、数类别、看尺寸很多人拿到「含图片数据集.zip」直接解压就开跑结果训练到一半发现某个类别只有十几张图或者图片尺寸从 200×200 到 2000×2000 都有DataLoader 直接报错。我一般会先写一个统计脚本把数据集的底细摸清楚。import os from PIL import Image from collections import Counter data_root ./rice_dataset # 解压后的根目录 class_counts Counter() size_samples [] for cls_name in sorted(os.listdir(data_root)): cls_dir os.path.join(data_root, cls_name) if not os.path.isdir(cls_dir): continue imgs [f for f in os.listdir(cls_dir) if f.lower().endswith((.jpg, .jpeg, .png, .bmp))] class_counts[cls_name] len(imgs) # 每个类别抽前 5 张看尺寸和通道 for f in imgs[:5]: with Image.open(os.path.join(cls_dir, f)) as im: size_samples.append((cls_name, f, im.size, im.mode)) print(类别分布:, dict(class_counts)) print(总图片数:, sum(class_counts.values())) for s in size_samples[:20]: print(s)这段脚本做三件事统计每个类别的图片数量、抽样查看图片尺寸和色彩模式、给出总样本量。参数上data_root指向解压后的根目录脚本默认每个子文件夹是一个类别这是 ImageFolder 的标准约定。如果输出里出现某个类别少于 50 张或者尺寸跨度超过 3 倍后面的预处理就要针对性处理不能一把梭。2.2 类别不均衡与尺寸不统一两个必须先处理的现实问题大米数据集常见的问题是类别不均衡。比如「正常米」有 800 张「碎米」只有 60 张「黄粒米」只有 40 张。直接训练模型会偏向多数类碎米和黄粒米的召回率极低。常见做法有两种一是对少数类做数据增强后再采样二是在损失函数里加类别权重。我一般先用加权交叉熵改动最小。import torch from torch.utils.data import DataLoader from torchvision import datasets, transforms train_tf transforms.Compose([ transforms.Resize((224, 224)), # 统一到 224适配主流 backbone transforms.RandomHorizontalFlip(p0.5), # 米粒方向随机水平翻转合理 transforms.RandomRotation(15), # 轻微旋转模拟摆放角度 transforms.ColorJitter(0.2, 0.2, 0.2), # 模拟不同光照 transforms.ToTensor(), transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]), ]) dataset datasets.ImageFolder(data_root, transformtrain_tf) counts [0] * len(dataset.classes) for _, label in dataset.samples: counts[label] 1 total sum(counts) weights torch.tensor([total / (len(counts) * c) for c in counts], dtypetorch.float32) print(类别权重:, weights)Resize((224,224))是为了对齐预训练 backbone 的输入RandomRotation(15)不要开太大米粒旋转 90 度后形态语义会变15 度以内比较安全ColorJitter模拟产线光照波动。权重计算用的是「总数 / (类别数 × 该类样本数)」少数类权重自然更大。这个权重后面传给CrossEntropyLoss(weightweights)即可。2.3 训练集 / 验证集 / 测试集怎么切才不泄漏同一个米粒可能被拍了好几张如果随机切分同一粒米的不同照片可能同时出现在训练集和验证集里验证准确率会虚高。稳妥做法是按「拍摄批次」或「原始文件名前缀」分组切分保证同一粒米只出现在一个集合里。如果数据集没有批次信息退而求其次用random_split时固定随机种子并检查类别比例是否一致。from torch.utils.data import random_split torch.manual_seed(42) n len(dataset) n_train int(0.7 * n) n_val int(0.15 * n) n_test n - n_train - n_val train_set, val_set, test_set random_split( dataset, [n_train, n_val, n_test], generatortorch.Generator().manual_seed(42)) train_loader DataLoader(train_set, batch_size32, shuffleTrue, num_workers4) val_loader DataLoader(val_set, batch_size32, shuffleFalse, num_workers4) test_loader DataLoader(test_set, batch_size32, shuffleFalse, num_workers4)batch_size32是 224 输入下的常用起点显存不够就降到 16num_workers4在 Linux 上比较稳Windows 下如果报错就改成 0。切分比例 7:1.5:1.5 适合千级样本量样本更少时验证集比例可以再降。3. 从零搭一个能打的大米分类 CNNbackbone 选型与训练循环3.1 别从 LeNet 开始ResNet18 微调是性价比最高的起点大米识别本质是细粒度图像分类类间差异小靠浅层网络提取的边缘纹理特征不够用。常见做法是用 ResNet18 或 EfficientNet-B0 做迁移学习。ResNet18 参数量约 1100 万在单张消费级显卡上就能跑ImageNet 预训练权重已经把通用纹理特征学好了微调几十个 epoch 就能收敛。下面给出完整的模型改造和训练代码。import torch.nn as nn from torchvision import models def build_model(num_classes, freeze_backboneTrue): model models.resnet18(weightsmodels.ResNet18_Weights.IMAGENET1K_V1) if freeze_backbone: for p in model.parameters(): p.requires_grad False # 替换最后的全连接层适配大米类别数 in_features model.fc.in_features model.fc nn.Sequential( nn.Dropout(0.3), nn.Linear(in_features, num_classes) ) return model num_classes len(dataset.classes) model build_model(num_classes, freeze_backboneTrue) device torch.device(cuda if torch.cuda.is_available() else cpu) model model.to(device) print(类别数:, num_classes, 设备:, device)freeze_backboneTrue表示先冻结卷积层只训练新的全连接头适合样本量几千张以内的场景如果数据量上万可以解冻后面几个 stage 一起微调。Dropout(0.3)是防止全连接层过拟合的常规手段。ResNet18_Weights.IMAGENET1K_V1是 torchvision 提供的预训练权重标识不同版本 API 略有差异以你本地 torchvision 文档为准。3.2 训练循环里必须记录的四个量loss、acc、lr、混淆矩阵训练循环本身不复杂但要把关键指标记下来否则调参就是玄学。下面这个循环在每个 epoch 结束后打印训练损失、验证损失、验证准确率并保存验证准确率最高的权重。import torch.optim as optim from sklearn.metrics import confusion_matrix criterion nn.CrossEntropyLoss(weightweights.to(device)) optimizer optim.AdamW(filter(lambda p: p.requires_grad, model.parameters()), lr1e-3, weight_decay1e-4) scheduler optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max20) best_acc 0.0 for epoch in range(20): 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) scheduler.step() model.eval() correct, total 0, 0 all_preds, all_labels [], [] with torch.no_grad(): for imgs, labels in val_loader: imgs, labels imgs.to(device), labels.to(device) outputs model(imgs) preds outputs.argmax(dim1) correct (preds labels).sum().item() total labels.size(0) all_preds.extend(preds.cpu().numpy()) all_labels.extend(labels.cpu().numpy()) acc correct / total print(fEpoch {epoch1}: train_loss{running_loss/len(train_set):.4f}, fval_acc{acc:.4f}, lr{scheduler.get_last_lr()[0]:.6f}) if acc best_acc: best_acc acc torch.save(model.state_dict(), best_rice_cnn.pth) cm confusion_matrix(all_labels, all_preds) print(混淆矩阵:\n, cm)AdamW的lr1e-3是微调全连接层的常用值weight_decay1e-4抑制过拟合CosineAnnealingLR让学习率按余弦曲线下降T_max20对应总 epoch 数。混淆矩阵只在最佳模型时打印方便看哪两个类别最容易混。如果发现「碎米」和「正常米」互相误判严重说明特征区分度不够要么加数据要么换更强的 backbone。3.3 数据增强的边界哪些增强对大米有效哪些会帮倒忙数据增强不是越多越好。对大米识别来说水平翻转、±15 度旋转、轻微亮度对比度扰动是有效的因为米粒在传送带上本来就是随机朝向和光照。但垂直翻转要谨慎米粒的胚芽位置有方向性垂直翻转可能产生现实中不存在的形态。另外CutMix 和 MixUp 在细粒度分类上有时反而掉点因为混合后的图像语义模糊模型学不到清晰的类间边界。我一般先只用几何增强跑一版 baseline再逐步加颜色扰动每加一项看验证集准确率变化涨了才保留。4. 训练大米识别模型时最容易翻车的五个地方4.1 现象验证准确率 99%测试集只有 70%原因同一粒米的多张照片被随机切分到了训练集和验证集验证集泄漏。解决按拍摄批次或文件名前缀分组切分确保同一粒米只出现在一个集合。如果数据集没有批次信息至少固定随机种子并检查类别分布。4.2 现象训练 loss 一直不降准确率卡在随机水平原因学习率太大导致梯度爆炸或者标签编码有问题比如标签从 1 开始而不是 0。解决先把学习率降到 1e-4 试一个 epoch确认 loss 有下降趋势再检查dataset.classes和标签映射是否连续。另外如果用了weight参数但权重计算错误比如除零也会导致 loss 异常。4.3 现象少数类召回率极低混淆矩阵里全被预测成多数类原因类别不均衡没有处理或者加权损失权重不够大。解决先确认weights是否正确传入CrossEntropyLoss如果仍然不行对少数类做过采样或者用 Focal Loss 替代交叉熵。Focal Loss 的gamma2是常用起点能进一步压低易分类样本的权重。4.4 现象GPU 显存够但训练速度极慢GPU 利用率只有 20%原因num_workers设置不当或数据预处理在 CPU 上成为瓶颈。解决把num_workers调到 CPU 核心数的 1/2 到 2/3开启pin_memoryTrue并把Resize等操作放在 DataLoader 的 transform 里而不是训练循环里。如果图片原始尺寸很大先用脚本离线缩放到 256×256 再训练能省掉大量重复解码时间。4.5 现象模型在本地跑得好部署到产线相机上准确率暴跌原因训练时的预处理和推理时的预处理不一致比如训练用了 ImageNet 归一化推理时忘了减均值除方差或者产线光照色温与训练集差异大。解决把预处理封装成一个函数训练和推理共用同一份代码产线部署前用现场相机拍 50 张图做一次快速验证看准确率是否在可接受范围。如果色温差异大在训练集里加入现场光照条件下的样本重新微调。5. 怎么验证你的大米识别模型真的能用混淆矩阵、置信度与现场抽检5.1 混淆矩阵要按类别看不能只看总体准确率总体准确率 95% 听起来不错但如果「黄粒米」的召回率只有 60%产线上每 10 粒黄粒米就漏掉 4 粒这个模型就不能上线。验证时先把测试集的混淆矩阵打出来逐类看召回率和精确率。下面这段代码输出每个类别的分类报告。from sklearn.metrics import classification_report 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) all_preds.extend(preds.cpu().numpy()) all_labels.extend(labels.cpu().numpy()) print(classification_report(all_labels, all_preds, target_namesdataset.classes, digits4))classification_report会给出每个类别的 precision、recall、f1-score 和支持样本数。重点看少数类的 recall如果低于 0.85就要回到第 4 章排查类别不均衡或增强策略。digits4保留四位小数方便对比不同模型版本的细微差异。5.2 置信度阈值给模型一个「拒识」的后悔药产线上不是所有米粒都清晰可辨有些模糊、遮挡、反光的样本模型强行分类反而添乱。稳妥做法是设一个置信度阈值低于阈值的样本送入人工复检通道。下面代码演示如何根据验证集确定阈值。import numpy as np model.eval() confidences, correctness [], [] with torch.no_grad(): for imgs, labels in val_loader: imgs, labels imgs.to(device), labels.to(device) outputs model(imgs) probs torch.softmax(outputs, dim1) max_probs, preds probs.max(dim1) confidences.extend(max_probs.cpu().numpy()) correctness.extend((preds labels).cpu().numpy()) confidences np.array(confidences) correctness np.array(correctness) for thresh in [0.5, 0.6, 0.7, 0.8, 0.9]: mask confidences thresh acc correctness[mask].mean() if mask.sum() 0 else 0 coverage mask.mean() print(f阈值 {thresh}: 覆盖率{coverage:.3f}, 准确率{acc:.4f})这段代码遍历不同阈值输出「覆盖率」和「准确率」的权衡。覆盖率指有多少样本被模型自信地分类准确率指这些样本里分对的比例。产线上一般要求覆盖率不低于 0.9同时准确率不低于 0.98。如果阈值 0.8 时覆盖率只有 0.7说明模型整体置信度偏低需要检查训练是否充分或数据分布是否偏移。5.3 现场抽检用产线相机拍 50 张图做一次快速验证实验室测试集和产线现场永远有差距。我的习惯是模型上线前拿产线相机在真实光照、真实传送速度下拍 50 张图覆盖每个类别至少 5 张跑一遍推理看准确率和置信度分布。如果现场准确率比测试集低超过 5 个百分点先检查预处理是否一致再检查光照和白平衡。这一步花不了半小时但能避免上线后批量误判的血泪教训。模型不是训完就完事现场抽检才是最后一道关。希望帮到你。本文还有配套的精品资源点击获取
返回列表