ARTICLE DETAIL

资讯详情

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

PyTorch实战:苹果腐烂识别分类模型训练与部署全流程

PyTorch实战:苹果腐烂识别分类模型训练与部署全流程 简介面向Python深度学习初学者的苹果腐烂识别项目基于PyTorch框架构建附带完整图片数据集。项目针对苹果品质检测场景实现新鲜与腐烂的二分类识别适合学习图像分类、数据增强、模型训练与界面部署等技术环节。资源包内共601个文件包括531张jpg图片、63张png图片等苹果样本数据以及标签txt、环境配置说明与三个Python源码文件压缩包总体积约64.41MB。目前已有129人学习下载。代码需要依次运行数据分割脚本、模型训练脚本和PyQt界面脚本其中数据预处理在短边增加灰边使图片变为正方形并辅以旋转、翻转等增强手段以提升数据多样性与模型鲁棒性。训练完成后权重保存在本地可直接用于推理测试或二次开发。整体流程清晰、依赖明确是实战驱动的图像分类入门参考。1. 苹果腐烂识别一个能直接跑通的 PyTorch 图像分类闭环把苹果腐烂识别这个任务交给深度学习最直接的路线就是拿 PyTorch 训练一个图像二分类模型。《基于python深度学习对苹果是否腐烂识别-含图片数据集》这套资源恰好把整个闭环做全了自带 freshApple / rottenApple 两类图片数据集内置了灰边补正方形和旋转、翻转增强跑通后能训练出区分腐烂苹果和新鲜苹果的模型最后还有一个 PyQt5 界面可以做图片实测。整套代码基于 Python PyTorch核心是三个脚本按顺序运行01 生成数据集文本、02 训练模型、03 启动识别界面。适合两类人——刚学完深度学习理论、想找一个完整图像分类项目练手的人以及手头有类似水果外观检测需求、想快速验证 PyTorch 全流程的从业者。这个项目对硬件要求很低普通 CPU 也能跑完小轮次训练。下面我从数据预处理开始把每个脚本的作用、关键参数和路上的坑挨个拆开讲。2. 项目拆解数据集组织、灰边补正方形与旋转增强的细节拿到压缩包第一件事不是解压后立刻跑代码而是先花几分钟把目录结构和脚本职责理清楚。这个项目的组织方式非常直白属于典型的 PyTorch 分类项目布局。2.1 目录结构与脚本调用顺序project_root/ ├── dataset/ │ ├── freshApple/ │ │ ├── freshApple (184).jpg │ │ ├── freshApple (184_flip).jpg │ │ └── freshApple (184_rotated45).jpg │ └── rottenApple/ │ ├── rottenApple (1).jpeg │ ├── rottenApple (463).JPG │ ├── rottenApple (463_flip).jpg │ └── rottenApple (463_rotated45).jpg ├── 01数据集文本生成制作.py ├── 02深度学习模型训练.py ├── 03pyqt_ui界面.py └── requirement.txt数据集按类别分子目录这是 PyTorch 数据加载最通用的组织方式torchvision.datasets.ImageFolder也是按目录名自动生成标签。细看文件名会发现每张原始图都配了_flip和_rotated45两个后缀版本说明作者先拍了原始照片再做了固定的水平翻转和 45 度旋转增强样本量翻了三倍。提示三个脚本有严格的依赖顺序——01 生成的 txt 是 02 的输入02 生成的 best_model.pth 是 03 的输入。跳步跑必然报文件不存在。跑之前先打开 requirement.txt 看一眼。里面应该列了 torch、torchvision、PyQt5、scikit-learn、Pillow 这些依赖。环境方面Python 3.8 或 3.9 搭配 PyTorch 1.10 以上都能跑torchvision 版本跟着 torch 走即可不需要追最新版。2.2 灰边补正方形letterbox 预处理与像素填充细节数据集预处理里最值得单独讲的一步是通过在较短边增加灰边把图片变为正方形。这个操作在目标检测里叫 letterbox套用到分类任务上同样有意义。为什么不直接 resize 成正方形直接拉伸会把苹果形状压变形圆的不圆、长的更长模型学到的是这种畸变特征测试时遇到真实比例的图片准确率会明显打折扣。补灰边相当于在短边方向加 padding图片内容比例完全不变苹果还是那个苹果只是四周多了灰色背景模型不会因为形状畸变产生误判。代码层面常见做法是读图后比较宽高把短边扩到与长边一致from PIL import Image def make_square(img, fill_value114): 把图片补成正方形短边填充灰色长边保持不变原图居中 w, h img.size if w h: return img side max(w, h) new_img Image.new(RGB, (side, side), (fill_value, fill_value, fill_value)) # 原图居中粘贴到正方形画布上 new_img.paste(img, ((side - w) // 2, (side - h) // 2)) return new_imgfill_value 取 114 是 YOLO 系列流传下来的默认灰值。这个值不是玄学它处在 0-255 的中间偏暗区间既不会像纯白 255 那样在归一化后产生过强的边缘响应也不会像纯黑 0 那样和真实阴影混淆。分类模型对填充色的敏感度不如检测模型高但统一用 114 是个好习惯以后迁移到检测任务也不用改。补成正方形还有一个附带好处后续做随机裁剪、随机旋转时正方形图片不容易出现边角越界数据增强的容错空间更大。如果图片本身是正方形这段逻辑会直接跳过不产生额外计算。2.3 翻转与旋转增强小数据集扩大样本量的固定搭配数据增强是这个项目的另一个关键设计。作者没有用在线随机增强而是把增强结果直接生成到数据集里让训练脚本读到的就是已经扩增过的图片。水平翻转实现简单img.transpose(Image.FLIP_LEFT_RIGHT)。腐烂斑块长在苹果左侧还是右侧是随机的翻转能让模型学到左右对称的腐烂特征对类别的判断不会因为位置变化而失效。旋转 45 度值得单独分析。这个角度属于中等强度变换能让模型适应苹果倾斜摆放的情况又不会像旋转 90 度那样出现大面积背景填充导致苹果特征占比骤降。实际拍摄时苹果在桌面上滚动、放置角度本身就随机这种增强贴合真实分布。PIL 的 rotate 实现如下from PIL import Image import os src_dir dataset/rottenApple for fname in os.listdir(src_dir): if _flip in fname or _rotated45 in fname: continue # 跳过已生成的增强图避免重复处理 img Image.open(os.path.join(src_dir, fname)).convert(RGB) base os.path.splitext(fname)[0] flipped img.transpose(Image.FLIP_LEFT_RIGHT) flipped.save(os.path.join(src_dir, f{base}_flip.jpg)) rotated img.rotate(45, expandFalse, resampleImage.BICUBIC, fillcolor(114, 114, 114)) rotated.save(os.path.join(src_dir, f{base}_rotated45.jpg))注意 rotate 的expandFalse表示旋转后保持画布尺寸不变旋转产生的四个三角区域用 fillcolor 填灰resampleImage.BICUBIC是三次插值质量比最近邻好。我一般会在代码里加一个跳过逻辑避免脚本重复运行时把_flip的图再翻一次生成出_flip_flip这种连环增强。这种固定增强的优点是可复现、可审查增强图直接落盘训练时不需要在 DataLoader 里再算一遍省 CPU 开销。缺点是多样性有限模型见过的变形只有翻转和 45 度旋转两种。如果后续发现模型泛化不足可以在 transform 里补随机亮度抖动和轻微模糊这个后面再展开。3. 跑通 01 脚本路径读取、标签映射与 train/val 文本生成01 脚本是整个流程的地基职责是把 dataset 目录下的图片路径和类别标签翻译成纯文本文件。虽然代码量不大但路径解析的边界情况很多值得逐行讲清楚。3.1 读取逻辑类别目录名到标签的映射最核心的映射规则是文件夹名即标签。freshApple 文件夹下的所有图片标签记为 0rottenApple 文件夹下的标签记为 1。这种命名方式在 PyTorch 生态里是默认约定ImageFolder也是按目录名字典序自动生成标签。脚本需要处理两个容易翻车的细节。第一是图片后缀不统一数据集里有.jpg、.JPG和.jpeg三种后缀如果只匹配小写.jpg会漏掉一批文件训练集数量悄悄变少准确率却查不出原因。第二是文件名带空格和括号像rottenApple (463).JPG这种名字写入 txt 后如果按空格切分路径会被拦腰截断训练时读图必然失败。所以读取时要统一处理后缀写入时用制表符做分隔符。我习惯把所有样本先收集成一个列表再做训练集和验证集划分而不是先写全部路径再手动拆分。3.2 代码实现按类别分层划分训练集与验证集完整实现如下可以直接对照项目里的 01 脚本理解import os from sklearn.model_selection import train_test_split dataset_root dataset classes [freshApple, rottenApple] class_to_label {name: idx for idx, name in enumerate(classes)} all_samples [] for cls in classes: cls_dir os.path.join(dataset_root, cls) for fname in os.listdir(cls_dir): # 统一转小写判断后缀覆盖 .jpg / .JPG / .jpeg if fname.lower().endswith((.jpg, .jpeg, .png)): path os.path.join(cls_dir, fname) all_samples.append((path, class_to_label[cls])) print(f共读取 {len(all_samples)} 张图片) # 按标签分层划分保证训练集和验证集里两类比例一致 train_samples, val_samples train_test_split( all_samples, test_size0.2, stratify[s[1] for s in all_samples], random_state42, ) with open(train.txt, w, encodingutf-8) as f: for path, label in train_samples: f.write(f{path}\t{label}\n) with open(val.txt, w, encodingutf-8) as f: for path, label in val_samples: f.write(f{path}\t{label}\n) print(f训练集 {len(train_samples)} 张验证集 {len(val_samples)} 张)逻辑说明脚本按类别目录逐个遍历fname.lower().endswith(...)统一处理后缀大小写问题。train_test_split的stratify参数按标签分层切分确保验证集里腐烂和新鲜的比例与全集一致。如果省掉这行随机划分可能把某一类的图片大量分到验证集导致训练集某类样本严重不足训练出来的模型对那类苹果几乎不识别。参数说明test_size0.2表示 20% 的图片划给验证集random_state42固定随机种子让每次运行划分结果一致排查问题时可以复现。txt 每行格式是「图片路径 制表符 标签」训练脚本按\t切分就能拿到完整路径文件名里的空格不会造成干扰——这是用制表符而不是空格做分隔符的关键原因。3.3 产物校验跑完 01 后先检查这三项01 脚本跑完后项目根目录下应出现 train.txt 和 val.txt。先别急着训练花一分钟做校验wc -l train.txt val.txt cut -f2 train.txt | sort | uniq -c head -3 train.txt我一般会检查三件事第一行数之和是否等于图片总数对不上说明有文件后缀漏匹配第二两类标签的计数是否都有如果只有 0 没有 1说明某个类别目录没被扫描到多半是目录名拼写不一致第三随机抽几行确认路径存在Windows 下直接os.path.exists(path)判断。有一回我在类似项目里发现 val.txt 全是同一类图片就是因为漏了stratify验证集准确率虚高模型上线后被真实数据打回原形。txt 内容是后面所有流程的地基这一步检查到位后面能少踩一半坑。4. 02 训练脚本实战ResNet18 选型、超参数与模型保存策略02 脚本承担真正的训练任务。这个项目是典型的二分类图像任务输入是预处理后的正方形图片输出是 freshApple / rottenApple 两个类别的概率。模型选型、超参设置和保存策略三个点逐一拆开。4.1 模型选型小数据集为什么优先考虑 ResNet18模型选型上ResNet18 是这类小数据集最稳妥的选择没有之一。理由有三层第一ResNet18 参数量约 1100 万相比 ResNet50 的 2500 万省一半以上显存训练速度快不少第二残差连接结构解决了深层网络的梯度消散问题即使只训练 30-50 个 epoch 也能稳定收敛第三torchvision 直接提供 ImageNet 预训练权重把预训练模型迁移到苹果腐烂识别这种小数据集上效果远好于从零训练。如果机器显存特别紧张可以考虑 MobileNetV3 或 EfficientNet-B0但第一次跑通流程建议先用 ResNet18。这个模型的社区生态最好网上踩坑记录最多出问题一搜就有答案不会卡在奇怪的报错上。4.2 训练代码结构自定义 Dataset、DataLoader 与超参设置训练脚本的核心结构如下import torch import torch.nn as nn from torch.utils.data import Dataset, DataLoader from torchvision import models, transforms from PIL import Image class AppleDataset(Dataset): 读取 01 脚本生成的 txt按行解析图片路径和标签 def __init__(self, txt_path, transformNone): self.samples [] with open(txt_path, r, encodingutf-8) as f: for line in f: parts line.strip().split(\t) self.samples.append((parts[0], int(parts[1]))) self.transform transform 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, labelAppleDataset必须实现__len__和__getitem__两个接口这是 PyTorch 自定义数据集的固定要求。__getitem__里convert(RGB)很关键数据集里混有.JPG后缀文件部分图片可能是灰度模式不做 RGB 转换会导致张量通道数不一致训练到中途才报 shape 错误。transform 部分有三步不能省transform 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))是 ResNet 的标准输入尺寸对应 ImageNet 预训练时的输入分布。Normalize的 mean 和 std 用的是 ImageNet 统计值因为加载了预训练权重这个必须对齐否则预训练的效果会被严重削弱——这是新手最容易忽略的一步。数据加载和模型构建train_loader DataLoader(AppleDataset(train.txt, transformtransform), batch_size32, shuffleTrue, num_workers2) val_loader DataLoader(AppleDataset(val.txt, transformtransform), batch_size32, shuffleFalse, num_workers2) model models.resnet18(weightsmodels.ResNet18_Weights.IMAGENET1K_V1) model.fc nn.Linear(model.fc.in_features, 2) # 替换最后一层为 2 类输出batch_size32在 4GB 显存的普通显卡上跑 ResNet18 没有问题显存不够就降到 16。num_workers2开两个子进程做数据加载Windows 下如果报多进程相关错误直接改成 0 最省事。替换model.fc那行是整个迁移学习的核心操作保留预训练的特征提取层只重新训练最后的全连接分类层。训练循环与超参criterion nn.CrossEntropyLoss() optimizer torch.optim.Adam(model.parameters(), lr0.001) scheduler torch.optim.lr_scheduler.StepLR(optimizer, step_size10, gamma0.5) epochs 50 best_acc 0.0 for epoch in range(epochs): model.train() running_loss 0.0 for inputs, labels in train_loader: outputs model(inputs) loss criterion(outputs, labels) optimizer.zero_grad() loss.backward() optimizer.step() running_loss loss.item() scheduler.step() model.eval() correct total 0 with torch.no_grad(): for inputs, labels in val_loader: outputs model(inputs) _, preds torch.max(outputs, 1) correct (preds labels).sum().item() total labels.size(0) acc correct / total print(fEpoch {epoch1}: loss{running_loss/len(train_loader):.4f}, acc{acc:.4f}) if acc best_acc: best_acc acc torch.save(model.state_dict(), best_model.pth)关键参数汇总如下参数取值说明batch_size32显存不足时降为 16不建议超过 64optimizerAdam迁移学习首选对学习率不敏感lr0.001Adam 的默认推荐值收敛快schedulerStepLR(10, 0.5)每 10 个 epoch 学习率减半epochs50小数据集 50 轮足够收敛保存策略验证集 acc 最高时保存避免过拟合轮次的权重覆盖最优模型scheduler.step()让训练后期学习率降下来loss 更容易收敛到更低平台。if acc best_acc是典型的模型保存策略只保留验证集准确率最高的权重避免最后几轮过拟合导致保存的模型反而更差。torch.max(outputs, 1)返回每行最大值的索引也就是预测类别。4.3 训练产物与收敛判断训练完成后目录下会出现 best_model.pth这是 03 界面要加载的权重文件。判断训练是否正常看两个指标loss 下降曲线和验证集准确率。二分类任务、CrossEntropyLoss 的初始值约 0.69。因为两个类别均等初始概率各 50%而 -ln(0.5)≈0.693。如果第一轮 loss 远高于这个值比如直接是 2 以上说明模型没有正确初始化或者数据标签错位。验证集准确率稳定在 95% 以上、loss 进入平台期基本可以认为训练收敛。如果出现训练集准确率 99% 而验证集只有 80% 的情况就是过拟合。常见对策是加强数据增强、加 Dropout 或提前停止这个项目固定增强已经做了一部分还可以在 transform 里补随机翻转和亮度抖动来缓解。5. 避坑清单环境安装、路径解析与训练推理中的四个常见问题这一章把最容易翻车的四个点整理成现象、原因、解决三段式每一条都是实际跑这类项目时的高频故障。5.1 conda 装 PyTorch 反复失败import 直接报错现象pip install torch卡在下载阶段下载完成但安装报错或者装完执行import torch直接报DLL load failed。原因PyTorch 的 wheel 包体积超过 800MB默认 PyPI 源在国内网络环境下极易超时中断DLL 报错则多半是把 GPU 版 torch 装到了没有对应 CUDA 运行库的机器上。解决先确认机器有没有独立显卡、支持哪个 CUDA 版本。没有 GPU 就直接装 CPU 版清华源指定安装最快pip install torch torchvision -i https://pypi.tuna.tsinghua.edu.cn/simple注意这个命令装的是 CPU 版本因为 PyPI 默认源里 torch 就是 CPU 版。有 GPU 想用 GPU 训练需要到 PyTorch 官网选对应的 CUDA 版本用官方给出的命令安装。requirement.txt 里的版本号是参考不是铁律Python 3.8 torch 1.13 torchvision 0.14 的组合在大多数教学项目里都能正常跑不用盲目追求最新版。5.2 训练时报 FileNotFoundError路径解析失败现象01 脚本生成的 txt 一切正常但训练时 Dataset 读取第一张图就报FileNotFoundError或者PIL.UnidentifiedImageError。原因两个常见因素。一是文件名里带空格和括号写入 txt 后解析方式不对导致路径分裂二是项目路径中包含中文个别 PIL 版本在 Windows 下对非 ASCII 路径处理不兼容。解决01 脚本写入时用制表符分隔训练脚本按\t切分这是最稳妥的方案。另外在 Dataset 里加一个路径存在性检查跑训练前把问题暴露出来assert os.path.exists(parts[0]), f图片不存在: {parts[0]}如果项目在中文目录下最省时间的做法是直接把项目移到纯英文路径再跑。这个坑在 Windows 下尤其多换路径比 debug PIL 源码高效得多。5.3 训练 loss 不下降甚至直接变成 nan现象第一轮 loss 就输出nan或者训了 20 轮 loss 纹丝不动像一条水平线。原因loss 变 nan 最常见是学习率过大导致梯度爆炸或者 transform 里漏了归一化、像素值以 0-255 的原始范围直接进模型数值范围异常。loss 不下降则可能和预训练权重不匹配用了 ImageNet 预训练权重但 Normalize 的 mean/std 没对齐。解决按顺序排查。先确认 transform 里有ToTensor()和Normalize没归一化先补上有 nan 就把 lr 从 0.001 降到 0.0001 再试最后检查 Normalize 参数是否为 ImageNet 标准值mean[0.485, 0.456, 0.406]、std[0.229, 0.224, 0.225]。每改一个变量只跑 3-5 个 epoch 验证不要攒一堆改动一起跑否则定位不到根因。5.4 加载权重报尺寸不匹配界面起不来现象03 脚本运行后torch.load加载 best_model.pth 时报RuntimeError: size mismatch for fc.weight: copying a param with shape torch.Size([2, 512]) ...。原因保存的模型是把最后全连接层替换成 2 类输出后训练的fc 层权重形状是 [2, 512]加载时如果重新构建的模型没有改这一层默认还是 ImageNet 的 1000 类输出形状是 [1000, 512]必然对不上。解决加载权重前先重建模型结构保持和训练脚本完全一致model models.resnet18(weightsNone) model.fc nn.Linear(model.fc.in_features, 2) model.load_state_dict(torch.load(best_model.pth, map_locationcpu)) model.eval()这是典型的训练能过、加载翻车场景。跑 03 界面之前先确认模型结构定义和 02 训练脚本完全一致特别是最后一行的model.fc替换。map_locationcpu是另一个习惯保证 GPU 上训练出的权重在无 GPU 的机器上也能加载。6. 进阶验证PyQt5 界面推理与模型鲁棒性自检03 脚本启动 PyQt5 窗口用户选一张图片界面显示预测类别和置信度。界面本身不复杂核心还是推理代码。推理的预处理必须和训练完全对齐这是整个流程最容易在当前环节出错的地方。def predict_image(model, img_path, transform): img Image.open(img_path).convert(RGB) img make_square(img) # 与训练一致先补正方形 tensor transform(img).unsqueeze(0) # 增加 batch 维度 model.eval() with torch.no_grad(): outputs model(tensor) probs torch.softmax(outputs, dim1) confidence, idx torch.max(probs, 1) return class_names[idx.item()], confidence.item()unsqueeze(0)是给单张图加一个 batch 维模型要求输入形状是 [N, 3, 224, 224]单张图只有 [3, 224, 224]必须补维。softmax把输出转成概率分布torch.max取最大概率和对应类别。推理时如果跳过make_square模型看到的图片分布就和训练时不一致——训练时每张图都是正方形推理时来了个长方形直接 resize苹果被压扁准确率莫名下降还查不出原因。模型训练完我习惯做一轮鲁棒性自检。光有验证集准确率不够验证集图片和训练集同源分布太接近。真正要测的是没参与训练的真实场景图侧面摆放的苹果、带阴影的苹果、光照偏暗的苹果各挑几张逐个推理。自检的做法很具体对同一张测试图生成三个变体——亮度降低 20%、顺时针旋转 10 度、轻微高斯模糊分别推理并记录置信度变化。如果置信度从 0.98 掉到 0.7甚至类别翻转说明模型对光照和模糊敏感需要补增强后重新训练。这套资源自带的数据增强覆盖了翻转和 45 度旋转但光照和模糊需要自己改 transform 补上transform transforms.Compose([ transforms.Resize((224, 224)), transforms.ColorJitter(brightness0.2, contrast0.2), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]), ])ColorJitter的 brightness 和 contrast 参数分别控制在 0.2 范围内随机扰动模拟真实光照变化。从那以后我每次训完分类模型都会强制把变体测试走一遍同一个错误不犯第二次。界面演示翻车比训练报错更打击信任——训练报错能解释演示翻车解释起来像黑匣子。希望这套拆解和避坑记录能帮你在自己的数据集上少走几步弯路。本文还有配套的精品资源点击获取
返回列表