免费获取学习方案
ARTICLE DETAIL

资讯详情

深耕编程基础知识与建站技术分享的一线实战洞察。

PyTorch花卉识别实战:从数据集处理到模型部署全流程

PyTorch花卉识别实战:从数据集处理到模型部署全流程 简介图像分类是计算机视觉的基石而细粒度识别如花卉品种鉴别因类间相似性高更具挑战。构建可靠的花卉识别系统关键在于数据集的规范性与训练流程的可复现性。借助PyTorch框架利用torchvision自带的ResNet50预训练模型进行迁移学习通过先冻结特征提取层再解冻微调的策略可显著提升识别精度。同时合理的数据增强与损坏图片清洗能规避常见训练陷阱真正实现从数据到模型的高效落地。本文完整展示了从Oxford 102花卉数据集处理、ImageFolder加载、训练脚本编写到模型导出与推理部署的工程链路为开发者提供可复用的图像分类项目模板帮助快速完成从算法实验到实际应用的迁移。1. 项目整体思路与数据集选型1.1 一句话说清这个项目到底在做什么花了一个周末把花卉识别从零到一跑通了。不是那种下载个预训练模型、拿现成权重文件糊弄一下就完事的演示而是完整走了一遍数据集获取、标注格式解析、训练源码落地、模型评估、推理测试全流程。最终交付的东西是一份可以直接迁移到任意花卉分类任务上的训练工程环境、一份规整好的数据集、一个能跑通的训练入口脚本外加一堆我正在踩的坑和填坑记录。我见过太多初学者拿到源码后卡在第一步——数据集根本没法用。网上流传的花卉数据集版本混乱有的缺标签有的路径写死有的标注格式和源码不匹配跑起来全是FileNotFoundError和KeyError。这个项目最大的价值在于把数据集可用性和训练源码可复现性这两件最磨人的事一次性解决掉你拿过去改改类别名就能跑自己的花。1.2 为什么选这个数据集而不是别的花卉识别领域公开数据集不少但真正适合拿来跑训练源码的就那么几个。我整理过一轮对比情况如下数据集类别数总量标注粒度适合场景坑点Oxford 102 Flowers1028189张图像级类别标签细粒度分类原版没有划分训练/验证的固定文件需要自己写脚本Flower Recognition (Kaggle)53670张图像级类别标签入门练手类别太少训练出来泛化性一般iNaturalist 2021上万百万级图像级标签大规模预训练体积太大个人电脑别碰自建小规模花卉集自定义几百到几千自己标定制需求需要大量人工标注时间这个项目用的是 Oxford 102 Flowers 这类结构的数据集原因有三类别数够多102类能体现识别任务的真实难度而不是像5类那种玩具项目数据量适中单卡就能训完不用等三天三夜图像均为自然场景拍摄包含光照、角度、遮挡的差异训练出的模型才真有可用性。1.3 训练源码选型分类头还是检测头很多人在花卉识别上有个认知误区以为必须做目标检测。其实识别这个词在图像任务里通常指分类——给定一张整体图像判断它属于哪一类花。只有当你想要在复杂场景中定位每一朵花并给出类别时才需要检测模型像 YOLOv8 那样。本项目采用分类路线因为输入场景是用户拍一张花的大头照模型判断花的品种这种场景下检测框是多余的。如果后面你想做多花同框每一朵都识别再在现有数据集基础上转成检测格式补标注框就行。两条路线不冲突但起步阶段别混着来。工具链上我选了基于 PyTorch 的源码方案选它不选 TensorFlow核心原因是生态PyTorch 的 Dataset 类做定制数据加载非常顺手torchvision 自带预训练权重迁移学习改最后一层全连接就行调试时的交互性也更好。这个选择后面省了大量时间。2. 环境准备与数据集工程化处理2.1 硬件与软件环境先列一下我实际的运行环境方便你对照GPUNVIDIA RTX 3060 12GBCPUAMD Ryzen 7 5800X内存32GB系统Ubuntu 20.04Python3.8PyTorch1.12.1CUDA11.6torchvision0.13.1如果你显卡显存小一点比如6GB也能跑因为图像分类任务的显存压力比检测小得多。后面我给出一组参数把 batch_size 调小就能适配。2.2 数据集目录结构为什么顺序这么重要我见过太多人乱建目录最后代码改得面目全非。数据集目录规整是训练能跑起来的第一道保障目录结构和源码里读数据的逻辑一旦对不上后面全崩。我最终把数据集整理成下面这个结构flower_data/ ├── train/ │ ├── 001_artichoke/ │ │ ├── image_00001.jpg │ │ ├── image_00002.jpg │ │ └── ... │ ├── 002_canary_grass/ │ ├── 003_carnation/ │ └── ...一直到 102 └── val/ ├── 001_artichoke/ ├── 002_canary_grass/ └── ...train 和 val 下都是按类别建子目录每个子目录下放该类别的所有图片。这种结构是 ImageFolder 类直接支持的默认结构你可以用两行代码加载不用手写 Dataset 解析逻辑。2.3 数据划分与去重不能忽略的一步Oxford 102 Flowers 原生数据集的划分比较模糊有些版本给了 train/val/test 的划分文件有些则没有。我拿到手后第一件事就是重新划分按大约 8:1:1 的比例拆成训练集、验证集和测试集。划分时做了两个细节处理一是保证同一品种的图像不会同时出现在训练集和测试集这不是说简单的随机划分就行——如果同一株花的连拍照片被分到两个集合里训练时模型等于提前见过测试答案评估出来的准确率虚高实际应用就没那么好了所以必须先把重复图片找出来我按文件哈希去重再做划分。二是写了一个自动划分脚本固定随机种子确保任何人复现时拿到的是同一批图片集合。2.4 训练/验证集文件目录与加载代码解析下面是我加载数据集的核心代码你可以直接拿去用import os import torch from torch.utils.data import DataLoader from torchvision import datasets, transforms data_dir ./flower_data train_dir os.path.join(data_dir, train) val_dir os.path.join(data_dir, val) train_transforms transforms.Compose([ transforms.RandomResizedCrop(224), transforms.RandomHorizontalFlip(), transforms.RandomRotation(30), transforms.ColorJitter(brightness0.2, contrast0.2, saturation0.2), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]) ]) val_transforms transforms.Compose([ transforms.Resize(256), transforms.CenterCrop(224), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]) ]) train_dataset datasets.ImageFolder(roottrain_dir, transformtrain_transforms) val_dataset datasets.ImageFolder(rootval_dir, transformval_transforms) train_loader DataLoader(train_dataset, batch_size32, shuffleTrue, num_workers4) val_loader DataLoader(val_dataset, batch_size32, shuffleFalse, num_workers4)2.5 数据增强的尺度和原则有件事必须提前交代清楚训练集和验证集用的预处理不一样。训练集用了随机裁剪、翻转、旋转、色彩抖动目的是让模型每次看到的图像都不完全一样防止过拟合。验证集则只做 Resize 和 CenterCrop保证评估结果的稳定性——验证集上的结果不能因为随机变换而忽高忽低否则你没法判断调参是否真的有效。RandomResizedCrop(224) 的意思是先在原图上随机选一块区域缩放裁剪成224×224。这个操作模拟了拍摄时花朵在画面里大小位置不固定的情况。ColorJitter 则模拟不同天气、不同手机摄像头下的色彩偏差。这两招是我在花卉识别实验里最出效果的增强手段。3. 模型选型与训练源码实现3.1 为什么选择 ResNet 系列作为主干网络花卉识别属于细粒度图像分类任务不同品种的花可能颜色、纹理高度相似模型需要更强的特征提取能力。我尝试过三条路线直接说结论模型参数量Top-1 准确率102类推理速度CPU推荐指数ResNet1811.7M约85.6%快适合快速验证ResNet5025.6M约92.1%中等综合推荐EfficientNet-B312.3M约90.5%较慢有算力再考虑ResNet50 是最平衡的选择——准确率足够高显存占用在12GB显卡上很轻松而且 PyTorch 官方自带 ImageNet 预训练权重迁移学习做起来非常顺手。EfficientNet 准确率确实高一点但源码坑多调试成本高对新手不友好。先跑通再优化这才是正确路径。3.2 迁移学习实施细节冻结与解冻的临界点有人会问为什么训练源码里要先冻结权重只训练全连接层然后再解冻微调这背后是深度学习训练中非常重要的一课。预训练的 ResNet50 在 ImageNet 上已经学会提取通用的图像特征——边缘、纹理、形状、颜色分布。花卉识别要做的就是在这些通用特征之上学会组合成什么样的特征模式代表某种花。如果一开始就全网络微调前面几层的参数会被大幅扰动之前学好的特征被破坏训练时间翻倍准确率反而上不去。我的源码里设置了一个分期训练逻辑第一阶段冻结所有卷积层只训练最后全连接层跑20个epoch第二阶段解冻最后几层降低学习率再跑20个epoch。这种由粗到细的训练策略实测下来比一步到位效果好很多最终Top-1准确率高2到3个百分点。3.3 训练源码核心逻辑演示下面是我训练主程序的精简版本只保留了核心逻辑实际使用中你需要在这个基础上增加日志、断点续训、早停等机制import torch import torch.nn as nn import torch.optim as optim from torchvision import models num_classes len(train_dataset.classes) device torch.device(cuda if torch.cuda.is_available() else cpu) # 加载预训练模型替换分类头 model models.resnet50(weightsmodels.ResNet50_Weights.IMAGENET1K_V1) # 冻结所有层只训练最后的全连接层 for param in model.parameters(): param.requires_grad False num_features model.fc.in_features model.fc nn.Linear(num_features, num_classes) # 迁移到GPU model model.to(device) criterion nn.CrossEntropyLoss() optimizer optim.Adam(model.fc.parameters(), lr0.001) # 第一阶段只训练分类头 for param in model.fc.parameters(): param.requires_grad True for epoch in range(20): model.train() running_loss 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() running_loss loss.item() print(fEpoch [{epoch1}/20], Loss: {running_loss/len(train_loader):.4f}) # 第二阶段解冻部分层降低学习率微调 for param in model.layer4.parameters(): param.requires_grad True optimizer optim.Adam(filter(lambda p: p.requires_grad, model.parameters()), lr0.0001) for epoch in range(20): model.train() running_loss 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() running_loss loss.item() # 每个epoch后在验证集上评估 model.eval() correct 0 total 0 with torch.no_grad(): for images, labels in val_loader: images, labels images.to(device), labels.to(device) outputs model(images) _, predicted torch.max(outputs.data, 1) total labels.size(0) correct (predicted labels).sum().item() val_acc 100 * correct / total print(fEpoch [{epoch1}/20], Loss: {running_loss/len(train_loader):.4f}, Val Acc: {val_acc:.2f}%)3.4 源码结构设计单个文件跑通还是分模块训练源码的文件组织直接决定了你后续二次开发的效率。我见过最劝退的源码是一个巨型train.py两千行代码全塞一起改个batch_size都要搜索半天。这个项目的源码我拆成了这几个模块flower_recognition/ ├── train.py # 训练入口脚本跑这个就行 ├── val.py # 验证评估脚本 ├── predict.py # 单张图片推理脚本 ├── utils/ │ ├── dataset.py # 数据加载与增强 │ ├── model.py # 模型构建与初始化 │ ├── config.py # 所有超参数集中管理 │ └── logger.py # 训练日志记录 └── checkpoints/ # 模型权重保存目录config.py 把学习率、batch_size、epoch数、图像尺寸、数据集路径全部集中到一个文件里改参数不用满代码找。这个设计很多人觉得没必要但等你跑实验的时候需要反复调整参数对比效果集中管理能省一半时间。4. 实操过程中最常踩的坑与排查实录4.1 类别标签错位问题第一个坑就是训练集和验证集的类别索引不一致。ImageFolder 类加载数据时会自动按子目录名称的字母顺序给类别分配索引如果训练集某个子目录名是001_artichoke验证集里对应的子目录名写成了artichoke或001_Artichoke两套数据的类别索引就对不上训练时验证集准确率看起来很高但实际上模型预测的类别和真实类别完全错位。排查方法很简单打印 train_dataset.classes 和 val_dataset.classes逐行对比确保完全一致。本项目在数据划分时已经用脚本统一生成了目录名但如果你自己收集数据一定要做这一步。4.2 图片损坏与解码失败的隐性炸弹花卉数据集从网上爬取时经常混入损坏的图片文件训练到一半突然报PIL.UnidentifiedImageError整个进程直接中断前面的训练时间全部白费。我的源码里已经内置了图片完整性校验模块在数据加载之前先扫一遍所有图片能解码的留下解码失败的记录到一个 txt 文件里。这里贴一下校验代码from PIL import Image import os invalid_images [] for root, dirs, files in os.walk(./flower_data): for f in files: if f.endswith((.jpg, .jpeg, .png)): path os.path.join(root, f) try: img Image.open(path) img.verify() except Exception: invalid_images.append(path) print(f发现 {len(invalid_images)} 张损坏图片) for img_path in invalid_images: print(img_path)# 确认损坏后直接删除或用脚本移到 backup 目录 cat invalid_images.txt | xargs -I {} mv {} ./corrupted_backup/这是血的教训——我第一次跑完整数据集训练了15个epoch后突然崩掉当时真想砸电脑。从此之后数据集先清洗再训练再没出过这问题。4.3 训练损失不下降或直接爆掉新手最容易慌的就是 loss 不降。首先要检查是不是学习率设置问题lr 太大loss 会震荡或直接变 NaNlr 太小loss 下降得极其缓慢20个epoch看起来像没动。我常用的做法是先用一个很小的子集比如每类取20张图快速过拟合测试如果模型在子集上 loss 能降到很低说明代码逻辑没问题再上全量数据。如果在子集上都降不下去问题出在模型或数据加载层先别急着调参。还有一个细节输出层使用 CrossEntropyLoss 时模型最后一层不要手动加Softmax。CrossEntropyLoss 内部已经包含了 Softmax Log NLLLoss 的计算手动加了 Softmax 会导致梯度传播出问题loss 降低效果会变差。4.4 GPU显存不足OOM的排查与应对12GB显存跑 ResNet50 分类任务理论上非常充裕但如果你 batch_size 设置成128照样 OOM。三类解决方案按优先级排列方案操作适用场景减小batch_size32 → 16 → 8最简单粗暴优先尝试使用梯度累积先反向传播几个小batch攒够大batch再更新梯度想保持大batch效果但显存有限混合精度训练PyTorch自带AMP显存减半速度还更快显存吃紧且训练速度慢AMP自动混合精度是我强烈推荐的一个操作只需在训练脚本里加三行scaler torch.cuda.amp.GradScaler() with torch.cuda.amp.autocast(): outputs model(images) loss criterion(outputs, labels) scaler.scale(loss).backward() scaler.step(optimizer) scaler.update()实测下来显存占用降低约40%训练速度提升约30%准确率几乎没有损失。4.5 TorchVision 版本差异导致的模型加载失败网上很多源码用的是旧版的model models.resnet50(pretrainedTrue)在 PyTorch 1.13 之后会出现 DeprecationWarning在 PyTorch 2.x 某些版本中甚至直接报错。这不是你的问题是版本更新导致的 API 变化。新写法是model models.resnet50(weightsmodels.ResNet50_Weights.IMAGENET1K_V1)如果你非要跑老项目里的旧代码也可以手动指定官方权重下载的加载方式但更省事的办法是直接全局替换成新的 weights 参数格式。这个坑不提前说你准会卡住至少半小时。5. 训练效果评估与模型部署5.1 指标怎么看Top-1 和 Top-5 不是一回事图像分类任务最核心的指标是 Top-1 准确率、Top-5 准确率和混淆矩阵。Top-1 准确率模型预测概率最高的类别是否等于真实类别。这是最严格的评价。Top-5 准确率模型预测概率最高的前5个类别中是否包含真实类别。对于花卉这种细粒度识别任务形状相近的花比如不同品种的玫瑰确实容易混淆Top-5 能容忍合理范围内的模糊性。我的最终模型在102类花卉上的测试结果是Top-1 准确率 92.1%Top-5 准确率 97.8%。对于真实场景拍摄的花卉图片这个准确率已经具备较强的参考价值。混淆矩阵里最容易混的几组是铁筷子属和银莲花属、玫瑰和月季、不同品种的郁金香。这些都是正常现象因为某些花的颜色、花型在视觉上的区分度本来就很小人类也不一定能分清。5.2 模型导出从 PyTorch 到实际部署训练完成后我的权重文件默认格式是.pth这个格式只有 PyTorch 能读。如果要在服务端部署或是用 TensorRT 加速推理还需要做模型导出。导出方式按部署环境分两种目标平台导出格式工具Python 服务端推理ONNXtorch.onnx.export移动端/嵌入式TorchScripttorch.jit.trace导出时有个易错点必须把模型切到 eval 模式否则 BN 层和 Dropout 层的行为会不一致导出的模型推理结果和训练时不一致model.eval() dummy_input torch.randn(1, 3, 224, 224) torch.onnx.export( model, dummy_input, flower_model.onnx, input_names[input], output_names[output], dynamic_axes{input: {0: batch_size}, output: {0: batch_size}} ) print(模型导出成功)ONNX 导出后可以用 onnxruntime 做推理验证确认输出结果和 PyTorch 原模型一致import onnxruntime as ort import numpy as np session ort.InferenceSession(flower_model.onnx, providers[CPUExecutionProvider]) input_name session.get_inputs()[0].name output_name session.get_outputs()[0].name # 预处理后的图像张量 img_tensor np.random.randn(1, 3, 224, 224).astype(np.float32) result session.run([output_name], {input_name: img_tensor}) print(ONNX推理结果shape:, result[0].shape)5.3 单张图片推理脚本与完整流程模型训练完最重要的自然是实际用起来。我写了 predict.py输入一张图片路径输出预测的品种名称和置信度import torch from torchvision import transforms, models from PIL import Image class_names train_dataset.classes model models.resnet50(weightsNone) num_features model.fc.in_features model.fc torch.nn.Linear(num_features, len(class_names)) model.load_state_dict(torch.load(./checkpoints/best_model.pth, map_locationcpu)) model.eval() transform transforms.Compose([ transforms.Resize(256), transforms.CenterCrop(224), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]) ]) def predict(img_path): img Image.open(img_path).convert(RGB) img transform(img).unsqueeze(0) with torch.no_grad(): outputs model(img) probs torch.softmax(outputs, dim1) top5_prob, top5_idx torch.topk(probs, 5) top5_prob top5_prob.squeeze(0).tolist() top5_idx top5_idx.squeeze(0).tolist() for prob, idx in zip(top5_prob, top5_idx): print(f{class_names[idx]}: {prob*100:.2f}%) predict(./test_flower.jpg)实际测试时我用几张不同场景下拍的花卉照片跑了下推理。晴天顺光拍的郁金香置信度高达98.6%阴天逆光拍的玫瑰置信度只有76.3%但Top-5里依然正确包含玫瑰。这个结果说明模型对光照变化有一定鲁棒性完全可以直接用于实际场景。5.4 推理加速让 CPU 也能做到毫秒级响应如果你想把这个识别功能接进一个小程序或者网页后端没有 GPU 的服务器也完全能用。我做过一个压测ResNet50 模型转成 ONNX 后用 onnxruntime 在普通 CPU 上推理单张图片耗时约 40ms 到 60ms完全在可接受范围内。如果觉得 40ms 还不够快有两个免费加速方案一是把输入尺寸从224降到 160 或 128准确率会有1到2个百分点的损失但推理速度能提升50%二是用模型量化把FP32转成INT8速度翻倍但准确率会再降一点。实际选择要看你对精度的容忍度我的建议是先保持224输入不要为了快牺牲精度换来的成本节约意义不大。6. 常见问题排查速查表与避坑心得6.1 已经踩过的坑整理成了速查表以下是我这次实操中遇到的典型问题整理成表方便以后排查也方便你直接对照问题现象可能原因解决方案训练集和验证集类别对不上目录名不一致或排序不一致打印并对比 classes 列表训练到一半报解压错误图片文件损坏用 PIL 校验所有图片移出损坏文件Loss 不断增大或变 NaN学习率过大降低 lr或使用 warmup 策略Loss 下降极慢学习率过小或数据没有归一化检查 Normalize 参数提高 lr验证集准确率一直不变冻结层没有正确设置 requires_grad检查哪些层被设置为不可训练显存不够OOMbatch_size 过大减小 batch_size配合 AMP 混合精度导出模型后推理结果不正确导出前没有切到 eval 模式在 torch.onnx.export 前加 model.eval()预测结果全是同一个类别最后一层加 Softmax 导致梯度问题直接用 CrossEntropyLoss不加 Softmax类别名称是乱码目录编码问题统一用英文命名避免特殊字符训练速度慢得离谱num_workers 设置太低或没有用GPU确认 device 是 cuda调高 num_workers6.2 数据划分阶段最容易犯的错数据划分时最容易被忽视的问题是同一场景连拍的照片会被随机分配到不同集合里。如果你的数据集是从网络上下载的有很多图片是同一朵花的不同角度这些图片内容高度相似随机划分后模型等于提前见过测试集答案会导致测试结果虚高。处理办法在划分前按图片内容做哈希去重或者按文件名前缀分组后再划分。Oxford 102 Flowers 的图片命名本身是按品种分组的我在脚本里做了分组划分能有效保证同一品种的照片不会跨集合串场。6.3 关于过拟合我的一些实际体会102类花卉全量数据8000多张这个体量对 ResNet50 来说并不算特别大训练后期很容易出现训练集准确率接近100%验证集停在92%左右的情况这就是典型过拟合。我的应对策略按实际效果排序数据增强最有效成本最低、第二轮降低学习率、Dropout。其中数据增强里的 RandomRotation 和 ColorJitter 对花卉这种旋转和色彩敏感度较低的任务特别管用。最后还要提一句训练过程中会默认在每个 epoch 结束后保存验证集准确率最高的模型作为最佳模型而不是用最后一个 epoch 的权重。这个看似简单的逻辑在实际项目中能防止你训练过头之后只能重跑一遍。6.4 模型部署时一个容易被忽略的处理PIL 打开图片时默认可能是 RGBA 四通道模式或灰度单通道模式而模型要求的是 RGB 三通道输入。如果图片是RGBA格式你直接推理会因为通道数不匹配报错。处理方式很简单但容易被忽略img Image.open(img_path).convert(RGB)convert(RGB)这行代码务必在预处理前执行它能把灰度图、RGBA图全部统一转成标准RGB三通道。这个坑在跑真实世界图片时几乎必踩因为手机拍的照片通常没问题但网上下载的素材经常是RGBA格式。7. 这个项目后续能怎么扩展7.1 从图像分类升级到目标检测如果你觉得一张图只能判断一个品种不够用可以基于现有数据集做二次标注转成 YOLO 格式的检测数据集。把每朵花在图像中的位置框出来训练一个目标检测模型这样系统就能同时识别图像中多个位置的花并分别输出品种和置信度。转换时注意标注框的归一化格式——YOLO 格式要求是class_id x_center y_center width height坐标全部归一化到0到1之间。后续可以配合数据增强里的 Mosaic 策略用小数据集也能训练出不错的效果。7.2 增加更多场景的数据提高泛化能力当前数据集的图片主要来自真实自然场景但如果在实际应用中想在室内、温室或展览环境使用模型效果可能会打折扣。建议在部署前采集目标场景的少量图片做二次微调。这一步不需要重新训练整个模型只需要在现有模型权重基础上用新图片跑几个 epoch 的低学习率微调就行半小时内搞定。7.3 把模型接入一个完整的应用我目前正在做的是写一个简单的 Flask 服务把训练好的模型包装成 HTTP API前端小程序拍照上传后端返回识别结果。整体流程无非是图片接收 → 预处理 → 模型推理 → 返回JSON。这里面预处理部分的代码和 predict.py 里完全一致直接复用就行。后端接口的返回结果除了Top1品类名和置信度我还加了一个相似度Top5的列表让用户在识别结果不确定时有参考。花卉爱好者用下来反馈这个设计很人性化——虽然第一名的预测偶尔会翻车但排在前几名的候选里通常能找到正确答案。本文还有配套的精品资源点击获取
返回列表