ARTICLE DETAIL

资讯详情

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

基于迁移学习的花卉分类图像识别工具:Python实现

基于迁移学习的花卉分类图像识别工具:Python实现 简介面向Python机器学习与计算机视觉学习者这份花卉图像分类识别项目资源基于17种花卉、每类80张图片的数据集旨在帮助读者掌握从图像预处理、特征提取、模型构建到训练评估与部署的完整工程流程。包内共2755个文件压缩包大小约251.53MB其中包含2720张花卉jpg图片作为训练与测试数据另有3个Python脚本负责数据处理和模型训练还有npy特征文件、PNG示意图、markdown说明文档及TeX格式论文条目能够完整反映项目从数据到模型输出的实现脉络。目前已有927人学习下载适合作为入门图像分类与CNN迁移学习的实战参考。资源中保留了数据集划分、VGG16等预训练网络特征提取、全连接分类器构建以及准确率评估调优等关键环节的代码与文档说明解压后即可查看清晰目录结构并运行验证能够帮助读者快速上手计算机视觉项目的工程实践。1. 花卉分类图像识别工具Python 代码下载之后要先想清楚的事把手机镜头对准一朵路边花程序在几百毫秒内返回「菊科·蒲公英」这就是花卉分类图像识别工具干的事。它比猫狗二分类难得多同一种花在不同花期、光照和拍摄角度下差异巨大而菊花和向日葵这类近缘属在形状上高度接近天然是图像识别里“类内差异大、类间差异小”的极端样本。用 Python 实现它的最短路径不是从零训 CNN而是迁移学习拿 ImageNet 预训练权重替换分类头在花卉数据上微调。下面按模型选型、数据整理、训练代码、参数调优、导出推理的顺序给出一套下载后改路径就能跑的工具代码以及跑不通时先查哪里的经验。2. 花卉分类的模型选型从传统特征到最新的图像识别模型2.1 传统图像识别特征为什么在花卉上不灵早年的花卉分类工具普遍是流水线先把花朵从背景里分割出来提取颜色直方图和 SIFT 特征再丢给 SVM。这条路线在单一背景的标本图上效果尚可一到自然场景就崩因为分割步骤把花瓣和背景糊在一起错误一路传导到分类器。特征是人手工设计的覆盖不了“同一朵花不同角度、不同光照”这种真正的难点这类方案在 102 类规模的数据集上很难过 60% 的准确率。深度学习图像识别的做法是让网络自己学特征浅层学边缘和纹理深层学花瓣排列与花型结构分类头只做最后一跳。Python 生态里 torchvision 和 timm 把预训练权重、数据增强、训练循环都封装好了这也是为什么现在做花卉分类几乎没人再手写特征提取器。需要提醒的是别把花卉分类和 tesseract.exe 那类 OCR 图像识别工具混为一谈OCR 的输出是文字序列模型做的是文字定位加序列识别花卉分类是整图判别输出一个类别分布。任务结构不同选型思路不能互相套用。2.2 主干网络选型算力决定下限数据量决定上限迁移学习里骨干网络的选择本质是「参数量—精度—部署成本」三角。下表是我在实际项目里常用的对照ImageNet Top-1 为官方报告的近似值主干网络参数量ImageNet Top-1适合场景ResNet5025.6M约 76%通用 baselineCPU 也能推理MobileNetV3-Large5.4M约 75%边缘设备、移动端部署EfficientNet-B312M约 82%训练数据够、追求精度性价比ConvNeXt-T28.6M约 82%较新的 CNN 架构精度优先ViT-B/1686M约 82%84%数据量大且有 GPU 集群时再上102 类的花卉分类我的默认选择是 ResNet50 起步它足够当 baseline显存占用小调参空间大遇到问题好排查。只有当验证集精度明确卡在 80% 以下、且确认不是数据问题时才换 EfficientNet-B3 或 ConvNeXt。别一上来就上 ViTViT 在小数据集上迁移效果反而不如 CNN除非你手里有 ImageNet-21K 级别的大预训练权重。这里有个容易被忽略的点迁移学习的上限由数据量决定而不是由网络大小决定。预训练权重提供的是通用视觉特征“菊花和向日葵的区别”这类细粒度知识必须从你自己的花卉数据里学。每类只有 50 张图时换更大的网络只会加速过拟合不会带来精度提升。2.3 数据集Oxford 102 与自建数据集的目录规范公开数据集最常用的是 Oxford Flowers 102102 类花卉共 8189 张图每类 40258 张图像带有明显的尺度和姿态变化是图像识别论文的标准 benchmark 之一。它自带训练、验证、测试划分文件下载解压后需要按 ImageFolder 约定重新组织目录每个类别一个子目录目录名就是类别名。我一般先把图片和划分文件归位再交给 torchvision 的 ImageFolder 读取。假设train.txt每行是「图片相对路径 类别索引」组织脚本长这样import os, shutil for line in open(train.txt): img_path, label line.strip().split() label int(label) # 0~101 dst_dir fdata/train/{label:03d} # 数字补零当目录名保证排序稳定 os.makedirs(dst_dir, exist_okTrue) shutil.copy(img_path, os.path.join(dst_dir, os.path.basename(img_path))) # valid.txt 同理目标目录换成 data/valid这段代码的关键是类名目录用{label:03d}补零ImageFolder 按目录名的字典序决定类别索引2会排在10前面补零之后 102 个目录顺序才稳定训练和推理时的索引映射才不会错位。自建数据集时经验值是每类至少 100 张原始图加上数据增强才够微调如果每类只有二三十张先别谈训练回去补数据或者换更保守的增强策略。评估预期也要摆正猫狗二分类随便一个模型都能上 95%花卉 102 分类能到 80% 以上就已经是有实用价值的工具了。3. 用 PyTorch 写好花卉分类图像识别的最小训练代码3.1 环境与目录结构下载代码后先跑通的骨架代码下载到本地后第一件事不是训练是把环境对齐。Python 3.9 以上即可依赖只有三件套用pip install torch torchvision tensorboard一次装完。torchvision 自带 ResNet50 等网络的预训练权重接口不需要额外装 timm。如果用 vscode 跑python 环境配置里最容易出问题的是解释器选错在命令面板执行 Python: Select Interpreter指到刚才执行过 pip install 的那个环境再跑脚本才不报 ModuleNotFoundError。项目骨架保持最小flower-classifier/ ├── data/ │ ├── train/class_id/*.jpg │ └── valid/class_id/*.jpg ├── train.py ├── infer.py └── requirements.txt3.2 训练脚本迁移学习加自动混合精度的标准写法import argparse, torch import torchvision from torch import nn from torch.utils.data import DataLoader from torchvision import datasets, transforms def make_transforms(size224, trainTrue): if train: return transforms.Compose([ transforms.RandomResizedCrop(size, scale(0.6, 1.0)), # 模拟花在画面中的占比变化 transforms.RandomHorizontalFlip(), transforms.ColorJitter(0.2, 0.2, 0.2), # 光照扰动 transforms.ToTensor(), transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]), # ImageNet 统计量迁移学习必须用 ]) return transforms.Compose([ transforms.Resize(int(size * 8 / 7)), transforms.CenterCrop(size), transforms.ToTensor(), transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]), ]) def build_model(num_classes): model torchvision.models.resnet50( weightstorchvision.models.ResNet50_Weights.IMAGENET1K_V2) model.fc nn.Linear(model.fc.in_features, num_classes) # 只换分类头卷积层保留预训练权重 return model def main(): ap argparse.ArgumentParser() ap.add_argument(--data, defaultdata) ap.add_argument(--epochs, typeint, default30) ap.add_argument(--batch-size, typeint, default32) ap.add_argument(--lr, typefloat, default3e-4) ap.add_argument(--size, typeint, default224) args ap.parse_args() train_ds datasets.ImageFolder(f{args.data}/train, make_transforms(args.size, True)) valid_ds datasets.ImageFolder(f{args.data}/valid, make_transforms(args.size, False)) train_loader DataLoader(train_ds, batch_sizeargs.batch_size, shuffleTrue, num_workers4) valid_loader DataLoader(valid_ds, batch_sizeargs.batch_size, shuffleFalse, num_workers4) device cuda if torch.cuda.is_available() else cpu model build_model(len(train_ds.classes)).to(device) criterion nn.CrossEntropyLoss(label_smoothing0.1) optimizer torch.optim.AdamW(model.parameters(), lrargs.lr, weight_decay1e-4) scaler torch.cuda.amp.GradScaler() for epoch in range(args.epochs): model.train() for images, labels in train_loader: images, labels images.to(device), labels.to(device) optimizer.zero_grad() with torch.cuda.amp.autocast(): loss criterion(model(images), labels) scaler.scale(loss).backward() scaler.step(optimizer) # AMP 梯度缩放显存减半 scaler.update() model.eval() # 验证必须关 dropout冻结 BN 统计 correct total 0 with torch.no_grad(): for images, labels in valid_loader: images, labels images.to(device), labels.to(device) preds model(images).argmax(1) correct (preds labels).sum().item() total labels.size(0) print(fepoch {epoch}: valid acc {100.0 * correct / total:.2f}%) torch.save({model: model.state_dict(), classes: train_ds.classes}, last.pt) if __name__ __main__: main()这段代码把整条训练流程压到了核心逻辑。三个值得注意的位置一是model.fc nn.Linear(...)只替换 ResNet 最后的全连接层前面的卷积层保留 ImageNet 权重这是迁移学习的标志性动作千万别把整个模型随机初始化二是label_smoothing0.1让模型不那么自信对花卉这种类间相似度高的任务能稳定提升泛化三是验证阶段必须model.eval()否则 BN 层会继续用当前 batch 统计量更新验证精度会虚高或抖动。3.3 关键训练参数lr、batch_size、image_size 怎么定直接给一组可抄的参数表参数推荐值调参方向image_size224小数据降到 192别低于 160batch_size32显存不足降到 16配合梯度累积lr (AdamW)3e-4只训分类头可以提到 1e-3epochs3050以早停为准不要死跑固定轮数weight_decay1e-45e-4越大越抑制过拟合label_smoothing0.1类间混淆严重时加到 0.2batch_size 是隐性调节器learning rate 要跟着 batch 走batch 翻倍lr 也应近似翻倍。32 的 batch 在 8GB 显存上跑 ResNet50 的 224 输入没问题如果卡只有 4GB把 batch 降到 16训练代码里的autocast和GradScaler已经帮你开了自动混合精度。还不够就做梯度累积每两步更新一次等效把 batch 扩到 64optimizer.zero_grad() scaler.scale(loss).backward() if (step 1) % 2 0: # 每 2 步更新一次等效 batch 翻倍 scaler.step(optimizer) scaler.update() optimizer.zero_grad()提示迁移学习最常见的学习率误区是沿用从头训练时的 1e-2。预训练权重的特征已经很好lr 过大会在头几个 epoch 把学好的权重冲坏表现就是验证精度早期猛涨之后一路崩。拿不准就从 3e-4 开始只降不升。4. 训练实战花卉分类图像识别的增强、混淆与过拟合4.1 数据增强小样本花卉的默认配置花卉数据集的普遍问题是量少且场景单一最有效的三个增强操作是 RandomResizedCrop、ColorJitter 和 RandomErasingtransforms.RandomResizedCrop(224, scale(0.5, 1.0)), # 花在画面中的占比从一半到全幅 transforms.RandomHorizontalFlip(p0.5), transforms.ColorJitter(brightness0.3, contrast0.3, saturation0.3, hue0.05), # hue 别调太大花色会失真 transforms.RandomErasing(p0.25, scale(0.02, 0.1)), # 模拟叶片遮挡逼模型看整体结构亮度对比度可以放开色调要克制。花卉分类里颜色是强判别特征比如白花和黄花的分界往往只靠色相hue 抖动过大等于把判别信息抹掉了。RandomErasing 的作用是逼模型不要只盯花朵中心的色块而去学茎叶、花型的组合结构这对遮挡场景下的识别提升明显。一个简单的验证方法是固定随机种子分别训练两版一版带 RandomErasing、一版不带对比验证集精度增不增强用数据说话。4.2 相似类混淆菊花和向日葵怎么分102 类里总有一批容易混淆的对子比如不同品种的菊科植物。判断哪些类在打架要在验证集上打混淆矩阵import torch cm torch.zeros(102, 102, dtypetorch.long) with torch.no_grad(): for images, labels in valid_loader: preds model(images).argmax(1) for t, p in zip(labels, preds): cm[t, p] 1 for i in range(102): top cm[i].argsort(descendingTrue)[1] # 第 0 名是自己看第 1 名 if cm[i, top] cm[i].sum() * 0.1: # 错认占比超过 10% 就值得记录 print(fclass {i} - {top}: {cm[i, top]})看到行方向的集中错认有两条路。如果发现的是「菊 A 被认成菊 B」这类近缘属混淆说明模型在学整体颜色和形状没抓到种间差异此时把颜色扰动调小、增加每类样本量比换模型更有效。如果错误分散在各处那更像数据量不足导致的欠拟合综合调高 epochs 并适当缩小 lr。另一种工程化做法是层级分类先训一个粗粒度模型区分到科再对混淆集中的科训细粒度分类器代价是要维护两套模型收益是每个子任务都更简单。4.3 过拟合的三个信号与早停实现花卉分类因为类多样本少过拟合来得很早。我判断过拟合只看三个信号信号表现对策train/valid 精度差拉大train 98%、valid 83%且持续扩大加大增强力度、提高 weight_decayvalid loss 反弹曲线先降后升形成 U 形提前截断或下一轮减小 lr单类退化个别类别精度突降为 0查该类样本数是否过少考虑合并相似类早停是省时间的硬道理固定 epoch 数训练是新手最常见的浪费时间方式。标准实现是记录最佳验证精度并计数连续多轮不刷新就停best_acc, bad_epochs 0.0, 0 patience 8 for epoch in range(max_epochs): valid_acc run_valid(..., ...) if valid_acc best_acc: best_acc valid_acc bad_epochs 0 torch.save({model: model.state_dict(), classes: classes}, best.pt) # 只存最好的一版别存 last else: bad_epochs 1 if bad_epochs patience: print(fearly stop at epoch {epoch}, best {best_acc:.2f}%) break提示保存的 checkpoint 里一定带上classes否则推理阶段你得靠猜来还原类别名。顺带把预处理方式写进注释三周后回来看代码的人很可能就是你自己会感谢这个习惯。5. 模型导出与单图推理把训练好的花卉分类工具跑在命令行5.1 导出 TorchScript 并写单图推理脚本训练结束拿到best.pt之后我一般会导出成 TorchScript脱离训练代码也能跑省掉推理时重建模型结构的麻烦import torch, torchvision model torchvision.models.resnet50() model.fc torch.nn.Linear(model.fc.in_features, 102) ckpt torch.load(best.pt, map_locationcpu) model.load_state_dict(ckpt[model]) model.eval() dummy torch.randn(1, 3, 224, 224) ts torch.jit.trace(model, dummy) # 用假输入固化计算图 ts.save(flower_model.pt) # 推理文件不再依赖 torchvision推理脚本 infer.py 的核心是保证预处理和训练侧完全一致import json, torch from PIL import Image from torchvision import transforms transform transforms.Compose([ transforms.Resize(int(224 * 8 / 7)), transforms.CenterCrop(224), transforms.ToTensor(), transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]), ]) classes json.load(open(classes.json)) # 训练时同步导出{0: daisy, ...} ts torch.jit.load(flower_model.pt) ts.eval() img Image.open(test.jpg).convert(RGB) # RGBA/灰度图要先转 RGB with torch.no_grad(): probs torch.softmax(ts(transform(img).unsqueeze(0)), dim1)[0] top3 torch.topk(probs, 3) for idx, score in zip(top3.indices, top3.values): print(f{classes[str(idx.item())]}: {score.item():.3f})softmax把 logits 转成概率后取 top-3是为了应对模型对相似花的低置信度如果最高分只有 0.4工具的正确行为是输出「疑似」而不是斩钉截铁的单一答案。这个交互设计对用户体验的影响比模型精度再提升 1% 更明显。5.2 验证指标与两个容易被忽略的坑导出前做一次全量验证集评测记录 top-1、top-5 精度和混淆最严重的五个类别对并保存每类的单类精度。只有训练、导出、推理三套路径验证一致才算真正「代码下载到本地就能用」。第一个坑是推理预处理与训练不一致。训练用 RandomResizedCrop推理直接Resize(224, 224)把图压变形花的长宽比一失真精度掉五个百分点以上很正常。上面的代码用Resize CenterCrop严格对齐验证集 transform这是底线。第二个坑是类别索引映射丢失。ImageFolder 按目录名排序生成索引重排目录或换机器后索引就会错位。解决方法是训练结束时把train_ds.classes以 json 形式落盘推理只认这份映射绝不自己推断。上面的classes.json就是这个用途TorchScript 模型文件本身不携带类别名这一步漏掉训练阶段的所有成果都会在推理时变成一串对不上的数字。本文还有配套的精品资源点击获取
返回列表