基于PyTorch的CNN柑橘成熟度识别实战:从数据预处理到模型部署

发布时间:2026/10/4 7:48:03
基于PyTorch的CNN柑橘成熟度识别实战:从数据预处理到模型部署 简介基于卷积神经网络的柑橘成熟度识别项目面向图像分类初学者、农业院校学生以及希望快速搭建果实检测原型的开发者。资源以PyTorch为运行框架配套自制的柑橘果实数据集图片按Healthy、Greening等成熟/病害类别分目录存放且已通过短边灰边填充及随机旋转完成数据增强可直接投入模型训练。包体共120个文件其中113张jpg图片构成本体3个Python脚本分别负责数据集标签生成、卷积模型训练与PyQt交互界面3个txt文件记录环境依赖与标签映射另有1张示例图片整体压缩后仅10.8MB下载与复现成本低。目前已有237人学习使用可作为课程设计或小规模科研实验的参考基线。代码流程连贯按序执行即可完成从图像预处理、模型训练到权重保存与可视化识别的全部环节环境依赖清单一并提供便于快速搭建运行环境也可替换数据集迁移至其他果蔬成熟度或品质识别任务。1. 拿到这份 PyTorch 柑橘分类资源先搞清它到底是干什么的拿到这份《基于卷积神经网络的柑橘成熟度识别-含数据集.zip》压得挺实。解压后你看到的不是花哨大工程而是一套能直接跑通的 PyTorch 分类任务数据集文件夹按类别放着柑橘图片01 脚本负责把图片路径和标签整理成文本02 脚本用卷积神经网络CNN训练并保存模型03 脚本用 PyQt 打开一个可视化识别界面。它解决的是“手里有几百张柑橘照片想快速分清健康果和病果又不想从零搭数据加载和训练流程”的问题。适合已经会跑 Python、但没完整走通训练管线的初学者也适合想拿真实数据集练手 CNN 全流程的人。整个包不需要做目标检测标注走的是最直接的分类路线环境按 requirement.txt 装好后三个脚本按顺序跑完就能看到模型文件落地。2. 先把数据吃透类别分布、灰边补方和旋转增强决定上半场成败这套代码最容易劝退的不是训练而是环境安装。压缩包里给了 requirement.txt也给了环境配置的参考说明链接说明里还提到一个付费的免安装环境包。按我的习惯不会直接买现成环境包而是用 conda 建一个干净的 Python 3.8 环境然后pip install -r requirement.txt几分钟的事之后出问题也更好排查。环境就位后先别急着双击 01 脚本数据这一关值得你先花十分钟看明白。2.1 解压后先别跑脚本先看目录和文件名数据集文件夹里存放的是各类别图片。从文件名看这个项目至少分两类Healthy是健康果Greening是染病果也就是果面绿化病这一类病果。文件名像Healthy (9).jpg、Greening (1).jpeg这样的是原始图而Healthy (9_rotated45.jpg、Healthy (9_flip.jpg这种带_rotated45、_flip后缀的说明包里已经有做过旋转和翻转的增强图片了。换台电脑跑之前我一般会先写个几行的小脚本看一眼类别分布避免后面标签和文件夹名对不上。import os from collections import Counter data_root path/to/dataset exts (.jpg, .jpeg, .png, .bmp) counts, samples Counter(), {} for cls in os.listdir(data_root): cls_dir os.path.join(data_root, cls) if not os.path.isdir(cls_dir): continue files [f for f in os.listdir(cls_dir) if f.lower().endswith(exts)] counts[cls] len(files) samples[cls] files[:3] print(counts)这段逻辑很直白遍历数据集根目录下的每个子文件夹把图片文件按扩展名筛出来计数。重点在于这里的cls必须是文件夹本身的名称因为 01 脚本生成标签时依据的就是文件夹名如果某个类目录名带空格或者中文后面训练会出各种莫名其妙的路径问题。扩展名默认加上.jpeg是因为包里的Greening (1).jpeg确实是这种后缀漏掉的话那一张图会被静默跳过。2.2 01 脚本到底做了什么从文件夹到 train.txt / val.txt01 脚本的名字在不同版本里可能写作“01数据集文本生成制作.py”或“01数据集文本制作.py”内容一致。它的工作是把每个类别目录里的图片路径和对应的类别编号拼成一行文本再按比例拆成 train.txt 和 val.txt。注意它只生成路径与标签不做像素级操作像素级的预处理放在后面的 DataLoader 或者 transforms 里这是这类小工程的常见做法。import os, random from sklearn.model_selection import train_test_split data_root path/to/dataset out_dir path/to/output train_ratio 0.8 random.seed(42) lines [] for cls_idx, cls_name in enumerate(sorted(os.listdir(data_root))): cls_dir os.path.join(data_root, cls_name) if not os.path.isdir(cls_dir): continue for f in os.listdir(cls_dir): if f.lower().endswith((.jpg, .jpeg, .png)): lines.append(f{os.path.join(cls_dir, f)} {cls_idx}) train_lines, val_lines train_test_split( lines, test_size1 - train_ratio, random_state42 ) with open(os.path.join(out_dir, train.txt), w) as fp: fp.write(\n.join(train_lines) \n) with open(os.path.join(out_dir, val.txt), w) as fp: fp.write(\n.join(val_lines) \n) print(ftrain: {len(train_lines)}, val: {len(val_lines)})这里的核心是“图片路径 空格 类别编号”的文本格式PyTorch 的 Dataset 读 txt 时按空格拆开即可。train_ratio设成 0.8 是常规值数据量小的时候可以调到 0.85 甚至 0.9但超过 0.9 验证集会太薄评估结果波动会很大。random_state42一定要固定这样每次跑 01 脚本生成的划分是一致的方便复现和对比训练效果。2.3 灰边补方为什么是短边加灰边而不是直接拉伸这个项目的预处理有个关键策略如果图片不是正方形就在短边两侧补灰边让图片变回正方形再进行缩放。如果图片本来就是正方形则不会加灰边。为什么补灰而不是直接 resize 成正方形因为柑橘本身是类圆形物体直接拉伸会让果形变扁变长CNN 学到的是被扭曲的形状比例补灰边则保留了原始比例只是把空白区域填充成中性灰度值。import cv2 import numpy as np def square_pad(image, value114): h, w image.shape[:2] side max(h, w) padded np.full((side, side, 3), value, dtypenp.uint8) y0 (side - h) // 2 x0 (side - w) // 2 padded[y0:y0 h, x0:x0 w] image return padded img cv2.imread(path/to/Healthy (9).jpg) img cv2.cvtColor(img, cv2.COLOR_BGR2RGB) img square_pad(img, value114) img cv2.resize(img, (224, 224))value114是目标检测和分类任务里常见的填充值具体数值不是玄学关键是训练和推理要保持一致如果你改成 128那推理脚本也得跟着改成 128。这段代码里我先把 BGR 转成 RGB因为 OpenCV 读图默认是 BGR 通道顺序如果不转后面给 PyTorch 的图颜色会偏蓝偏暗。resize 到 224 是 ResNet 系列网络的默认输入尺寸这种小规模分类项目用 224 是最稳妥的选择。2.4 旋转增强角度别贪验证集别污染项目里提到的扩增方式是旋转角度。旋转增强对分类任务来说非常安全因为柑橘从哪个角度拍都还是柑橘标签不会变。但角度本身是个超参数我一般不会用 45° 这种大幅旋转而是 8° 到 15° 的小角度加水平翻转就够。旋转角度太大果柄、果脐这些有方向性的细节会被扭曲模型会去学一些不该学的姿态不变性。def rotate_aug(image, angle15): h, w image.shape[:2] M cv2.getRotationMatrix2D((w / 2, h / 2), angle, 1.0) return cv2.warpAffine(image, M, (w, h), borderValue(114, 114, 114))旋转后空出来的四角用同样的灰值填充这张增强图就相当于一张“新样本”。不过这里有一个很容易翻车的点如果数据集包里已经带了一批_rotated45、_flip的预增强图片01 脚本会把它们也当成独立样本写入 txt如果你在训练脚本的 transforms 里又做了一遍随机旋转相当于二次增强同一个原始图会被制造出大量高度相似的副本训练集和验证集之间会出现“近亲”样本训练完 val_acc 虚高换成现场实拍图立刻露馅。提示训练集的 transforms 可以加随机旋转和翻转验证集的 transforms 不要加。验证集只需要 resize、归一化这些确定性操作否则你评估的不是模型能力而是随机种子。2.5 数据加载阶段再做增强transforms 分离写法项目里如果已经把增强图预生成好了我建议你在训练脚本里把 transform 写得更克制一点。下面这种写法是分类任务的标配训练集用带随机性的 transform验证集用纯确定性的 transform。import torchvision.transforms as T train_transform T.Compose([ T.Resize((224, 224)), T.RandomRotation(10), T.RandomHorizontalFlip(p0.5), T.ToTensor(), T.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]), ]) val_transform T.Compose([ T.Resize((224, 224)), T.ToTensor(), T.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]), ])这里的 Normalize 用的是 ImageNet 统计量如果你 02 脚本里加载了 torchvision 的预训练权重就必须用这组数值如果模型是从零训练可以用自己的均值方差或者干脆把 Normalize 去掉也能跑只是收敛可能慢一点。Resize 到 224 之前其实还应该先补方如果前面已经用 OpenCV 把图补成正方形并存好了这里 Resize 就不会变形如果没存T.Resize会直接拉伸效果差一截。3. 02 训练脚本网络选型、数据集读取和模型保存的关键点01 脚本把文本准备好之后02 脚本才进入真正的深度学习训练。这一章的三个关键词是网络选型、不会被污染的数据加载、以及能顺利被 03 界面加载的模型保存方式。这三个点有一个没做好后面全白跑。3.1 小数据集该选什么网络ResNet18 迁移学习是稳妥起点这份资源的数据集规模从文件名就能看出来属于几十到几百张的量级。这种量级下从零训练一个深 CNN 很难收敛常见做法是用 torchvision 里带预训练权重的 ResNet18 做微调把最后一层全连接改成二分类输出。预训练权重的意义在于模型已经在大规模数据集上学过纹理、边缘、颜色分布这些通用特征柑橘的果皮纹理和病斑特征不需要从头学起训练起来几轮就能看到 loss 快速下降。import torchvision.models as models import torch.nn as nn model models.resnet18(weightsmodels.ResNet18_Weights.IMAGENET1K_V1) model.fc nn.Linear(model.fc.in_features, 2)in_features是 ResNet18 倒数第二层的维度用model.fc.in_features取而不是写死数字换网络结构时不用改。二分类输出 2 个神经元配合nn.CrossEntropyLoss()使用时不需要手动加 softmax损失函数内部会算。内存紧张的话可以换mobilenet_v3_small精度略低但训练和推理都快很多数据量超过每类 200 张时可以考虑 ResNet34。3.2 Dataset 与训练循环txt 路径读取和指标打印02 脚本的典型结构是读 train.txt / val.txt → 自定义 Dataset → DataLoader → epoch 循环 → 保存模型。Dataset 的写法直接决定训练脚本会不会在数据读取阶段崩掉。from torch.utils.data import Dataset, DataLoader from PIL import Image class FruitDataset(Dataset): def __init__(self, txt_path, transformNone): self.pairs [] with open(txt_path) as fp: for line in fp: path, label line.strip().rsplit( , 1) self.pairs.append((path, int(label))) self.transform transform def __len__(self): return len(self.pairs) def __getitem__(self, idx): path, label self.pairs[idx] img Image.open(path).convert(RGB) if self.transform: img self.transform(img) return img, label这里rsplit( , 1)是从右侧只切一刀防止图片路径里带空格时把标签位切坏。用 Pillow 读图然后convert(RGB)是省事且不容易出错的写法绕开了 OpenCV 的 BGR 通道问题但要注意如果前面预处理依赖 OpenCV 的补方函数就需要在__getitem__里先读成 numpy 数组补方再转回 PIL二者必须二选一别混用。训练循环是整套脚本里最不需要花哨技巧的部分把标准流程写稳就行。device torch.device(cuda if torch.cuda.is_available() else cpu) model model.to(device) criterion nn.CrossEntropyLoss() optimizer torch.optim.Adam(model.parameters(), lr1e-4) scheduler torch.optim.lr_scheduler.StepLR(optimizer, step_size10, gamma0.5) epochs 30 for epoch in range(epochs): model.train() total_loss, correct, total 0.0, 0, 0 for images, labels in train_loader: images, labels images.to(device), labels.to(device) optimizer.zero_grad() outputs model(images) loss criterion(outputs, labels) loss.backward() optimizer.step() preds outputs.argmax(dim1) total_loss loss.item() * images.size(0) correct (preds labels).sum().item() total labels.size(0) val_acc evaluate(model, val_loader, device) print(fepoch {epoch1}/{epochs} loss {total_loss/total:.4f} facc {correct/total:.4f} val_acc {val_acc:.4f}) scheduler.step()lr1e-4是微调预训练模型的常见起点一般不需要更大如果 loss 迟迟不降可以试 3e-4但稳一手的话先跑 5 轮看曲线。Adam 在这里是图方便换成 SGD momentum 0.9 也是常见做法。evaluate函数就是遍历验证集算平均准确率每次 epoch 结束后打印一次不用等 30 轮全跑完才开始判断。3.3 保存模型state_dict 加 meta 信息别让 03 界面猜参数训练完保存模型文件时我强烈建议不要直接torch.save(model, model.pth)保存整个模型对象那样换个环境或者改一行网络定义后加载时经常报类名不匹配。保存state_dict加一些附加信息是更通用的做法。torch.save({ model_state_dict: model.state_dict(), num_classes: 2, input_size: 224, }, fruit_cnn.pth)03 界面脚本加载时先用num_classes和input_size重建模型结构再load_state_dict这样训练和推理的解耦最干净。如果别人给你的工程里已经用了整模型保存方式加载时记得写torch.load(path, map_locationcpu)并且不要再去改网络结构定义否则键名对不上。3.4 训练不收敛时先看症状别急着调参运行 02 脚本时loss 不降是很多人最慌的时候。不要一上来就调学习率先看症状loss 一开始就是 NaN多半是学习率太大或者标签越界loss 一直稳定在 0.69 不动二分类数据集随机猜就是 0.69说明模型没学到东西检查输入图和标签是否对齐训练几轮后验证集接近 100%先怀疑验证集被污染而不是模型有多强。症状可能原因先查什么loss 为 NaN学习率过大 / 标签越界打印 labels 最大值loss 停在 0.69类别不平衡 / 输入全是噪声检查数据集路径与增强验证集 100%验证集与训练集有同源副本检查重复文件与增强叠加验证集震荡大验证集样本太少调大 train_ratio 或做 K 折这几种情况里验证集污染最容易踩且最隐蔽具体排查方法在第 5 章展开。4. 03 界面与真实推理从 pth 到一次前向的完整链路03 脚本是 PyQt5 写的可视化识别界面它的功能套路是选择一张图片在界面上显示原图同时输出预测类别和置信度。但界面只是一个壳真正决定识别效果的是它背后的推理函数。先把推理链路单独拉出来验证再进界面这是我最推荐的调试顺序。4.1 先写一个独立的单图预测函数无论 UI 怎么包03 脚本最终调用的都是同一套流程读图 → 补方 → resize → 归一化 → 前向 → softmax。下面这个函数可以单独放到一个predict.py里界面和命令行共用。def predict_single_image(model, image_path, transform, device): img cv2.imread(image_path) if img is None: raise ValueError(fcannot read {image_path}) img cv2.cvtColor(img, cv2.COLOR_BGR2RGB) img square_pad(img, value114) img cv2.resize(img, (224, 224)) from PIL import Image tensor transform(Image.fromarray(img)).unsqueeze(0).to(device) model.eval() with torch.no_grad(): logits model(tensor) prob torch.softmax(logits, dim1)[0] cls_id int(prob.argmax().item()) return cls_id, float(prob[cls_id].item())这一段里面最关键的是“预处理和训练保持一致”补灰值都是 114尺寸都是 224Normalize 参数都是同一组。如果训练时用的是 OpenCV 补边、推理时换成了 PIL 直接 resize尺寸和通道顺序不对齐界面上可能显示能出结果但换一张光线暗一点的实拍图就翻车了。unsqueeze(0)把单张图变成 batch 维度为 1 的输入这是 PyTorch 前向推理最容易忘的一步。4.2 PyQt 显示图片的三个细节界面脚本里显示图片时容易出问题的地方集中在三个小点。第一OpenCV 读出来是 BGR显示到 QLabel 上之前要转成 QImage 的 RGB888 格式否则柑橘颜色会偏蓝第二大图显示前要做缩放setScaledContents(True)可以省事但缩放会改变视觉比例如果同时要在图上画框需要自己按比例换算坐标第三预测过程不要阻塞 UI 主线程单张图推理通常几十毫秒问题不大但图片很大时建议放到按钮的槽函数里加个忙碌光标不要用多线程硬撑。界面上的类别显示宁可显示类别名字符串也不要显示 0 和 1。类别编号的顺序取决于 01 脚本里sorted(os.listdir(...))的结果直接显示数字用户根本不知道 0 是 Healthy 还是 Greening这是界面设计里很low但很常见的体验坑。4.3 命令行先行验证模型产物在打开 PyQt 界面之前先用命令行跑一次单图预测可以快速排除一大半问题。下面这个命令直接在终端执行能把模型加载和推理链路一起验证掉。python -c import torch from your_model_code import build_model model build_model(num_classes2) ckpt torch.load(fruit_cnn.pth, map_locationcpu) model.load_state_dict(ckpt[model_state_dict]) rid, conf predict_single_image(model, test_img.jpg, val_transform, cpu) print(rid, conf) 这里map_locationcpu是故意的因为本机不一定有 CUDA而训练时可能是在 GPU 上保存的加了它就能避免设备不匹配的报错。如果这一行命令能跑通再进 03 界面问题就只可能出在 PyQt 的图片读取或显示层如果这里就报错优先检查state_dict的键名和num_classes是否一致。4.4 把预处理统一成同一个公共函数训练脚本、命令行预测、PyQt 界面三处对图片的处理必须一模一样。一个很实用的做法是把“读图 补方 resize 转 PIL”抽成一个公共函数训练和推理都调用它。这样只需要维护一份预处理代码减少训练和推理不一致的概率。项目里如果 01、02、03 三个脚本各自写了各自的预处理强烈建议在复现时重构成一个data_utils.py这是这套工程里最值得动手改的地方。5. 复现路上的五个常见问题现象、原因、解决办法这个工程本身不大能踩的坑基本都藏在“图片尺寸、类别不平衡、路径和模型加载”这几个黑匣子里。下面五条都是我在类似分类任务上实际遇到过的翻车记录按“现象 → 原因 → 解决”的顺序给你排过雷。5.1 训练时报形状错误从补方到 batch 维度现象02 脚本刚开始训练就报类似Expected input batch_size (1) to match target size或者 tensor 维度对不上的错误有时候会直接崩在 Dataset 的__getitem__里。原因最常见的两种情况。一是某张图不是正方形补方逻辑没生效resize 之后长宽不一致二是__getitem__返回的图片张量形状不对比如少了 batch 维度或者通道维度被 squeeze 掉了。解决在训练循环第一轮打印images.shape确认输出是[B, 3, 224, 224]。然后单独测试补方函数对一张 1000×800 的图补方后打印padded.shape应该是[1000, 1000, 3]。补方代码最好是先补方再 resize顺序反了会得到一张带灰边的变形图。5.2 验证集虚高换实拍图就废现象训练日志里 val_acc 高得离谱接近 100%但拿一张手机拍的现场柑橘图去 03 界面测试识别结果完全不能用。原因数据集包里已经预生成了大量旋转和翻转的副本图片01 脚本把这些副本和原图同时写进了 train.txt 和 val.txt同一个原始图的不同副本分别出现在训练集和验证集里。这就是验证集污染模型相当于已经见过验证集图片的“亲戚”指标当然好看。解决跑 01 脚本之前先对数据集里的图片做去重。按文件名解析出原始编号把_rotated45、_flip这类后缀的副本单独拿出来检查。最干净的做法是原始图和增强图全放进训练集验证集只用原始图并且保证验证集的原始图不出现在训练集里。增强应该在训练脚本的 transforms 里实时做而不是预生成一堆文件混在数据目录里。5.3 换台电脑训练txt 里的路径全部失效现象项目从一台电脑拷贝到另一台电脑跑 02 脚本直接报 FileNotFoundError指向 train.txt 里的某一行路径。原因01 脚本生成 txt 时写的是绝对路径比如C:/user/data/Healthy (9).jpg换电脑后目录前缀变了路径自然失效。解决改 01 脚本txt 里不要写绝对路径而是写“相对路径 类别名 标签”运行时再拼接data_root。比如每行存Healthy/Healthy (9).jpg 0训练脚本里用os.path.join(data_root, line_path)拼完整路径。这样换电脑只需要改一个data_root变量不用重新跑 01。注意路径里不要带中文跨平台时用/分隔符最省心。5.4 加载 pth 报 missing keys / unexpected keys现象03 脚本加载模型时报Missing key(s) in state_dict或Unexpected key(s)有时候明明刚训练完同一份代码加载却失败。原因训练时保存的是整个模型对象加载时用了load_state_dict或者训练时模型在 GPU 上加载到 CPU 没有加map_location还有一种情况是训练脚本改了model.fc的输出维度保存 meta 信息里却没更新num_classes加载时的网络结构和保存的权重维度不一致。解决统一用state_dict保存加载时显式传map_locationcpu。并且像第 3.3 节那样把num_classes、input_size一起存进去加载时先读 meta 再重建网络。如果手里的 pth 已经是整模型格式加载后先用model.state_dict().keys()打印键名和当前模型的键名逐个对比差在哪里一眼就能看到。5.5 PyQt 界面读不了中文路径的图片现象03 界面里选了路径带中文的图片预测结果空白或者直接报错但同一张图用命令行跑预测函数却正常。原因OpenCV 的cv2.imread对中文路径支持不好在 Windows 上尤其明显会静默返回 None后续代码拿到 None 再算形状就崩了。解决这个问题的后悔药是绕开cv2.imread用numpy.fromfile加cv2.imdecode读图。读图函数单独封装训练脚本和 UI 共用。import numpy as np import cv2 def imread_unicode(image_path): data np.fromfile(image_path, dtypenp.uint8) img cv2.imdecode(data, cv2.IMREAD_COLOR) if img is None: raise ValueError(ffailed to read image: {image_path}) return imgnp.fromfile按字节读文件不经过路径编码解析cv2.imdecode直接从内存解码中文路径就不再是问题。这个函数应该放在公共工具文件里01 脚本做数据检查时也可以先用它验证所有图片都能正常读出。6. 再深入一步用混淆矩阵和阈值校准验证“能不能用”训练完看 val_acc 是远远不够的。二分类任务里两个类别样本数量不均时准确率会被多数类掩盖。比如模型把所有图都判成 Healthy总体准确率可能也有 80%但健康果分拣这种场景根本没法用。所以训练完之后我会再多走两步打印混淆矩阵再看置信度分布。混淆矩阵可以用 sklearn 一行算出来把验证集的真实标签和预测标签收集齐直接传进去。from sklearn.metrics import confusion_matrix, classification_report # y_true 和 y_pred 分别是验证集遍历时收集的标签列表 print(confusion_matrix(y_true, y_pred, labels[0, 1])) print(classification_report(y_true, y_pred, target_names[Healthy, Greening]))输出结果里看 Healthy 被误判成 Greening 的比例和 Greening 被误判成 Healthy 的比例这两类错误在实际场景里后果完全不同。病果被当成健康果流入市场是质量问题健康果被当成病果扔掉是损耗问题。如果后者的比例偏高应该考虑调整阈值而不是盲目加数据。阈值调整的思路很简单默认 argmax 相当于在 0.5 处截断但实际使用中可以把 softmax 输出的判断阈值抬高到 0.7置信度低于 0.7 的样本归入“人工复核”类别。对于柑橘分拣这种场景拿不准的总比判错强。UI 界面里如果只显示高分预测就额外加一个“置信度低”的提示框这是实用性和准确性之间的平衡点。最后再说一个我自己的习惯。拿到这种分类数据集我以前是真的会直接跑 02结果踩过一次验证集泄漏的坑之后现在哪怕是一个 10 分钟的小实验也强制走一遍流程先跑目录统计脚本看类别分布和重复文件再跑一次单图预测验证模型链路最后用混淆矩阵收尾。这套顺序看着多花十分钟实际能省下大量对着 log 猜谜的时间。希望帮到你。本文还有配套的精品资源点击获取