基于PyTorch Mobile与MobileNetV3的移动端花卉识别实践

发布时间:2026/8/31 17:57:27
基于PyTorch Mobile与MobileNetV3的移动端花卉识别实践 简介这是一套面向AI移动应用开发初学者与进阶者的实战项目资源聚焦于在Android终端上部署轻量级花卉图像识别功能解决用户通过手机拍照即时识别花卉种类的实际需求。资源包含完整可运行的Android Studio工程与PyTorch训练/部署全流程代码涵盖前端界面、图像采集与预处理、模型加载推理含已训练.pt模型、结果可视化等核心模块适合深度学习与移动端交叉领域实践者系统学习。压缩包共111个文件约13.49MB其中Python脚本14个用于数据集构建与模型训练Java源码7个实现Android端图像处理与模型调用JPG/PNG图片50张及XML标注文件支撑本地测试Gradle配置与Jar依赖库保障工程可直接编译运行。目前已有405人学习下载目录结构清晰分层含flower_recognition-master主工程、模型文件夹、数据集样本及详细README说明开箱即用且便于二次定制。1. 项目概述与技术架构1.1 这个项目到底做了什么花卉识别这个方向在计算机视觉里算是个经典又不算太难的分类任务。但一旦加上移动终端这个限定词事情就没那么简单了。我最初接到这个项目需求的时候对方说要做一个能在手机上跑的花卉分类器要包含完整的源码、数据集和训练好的模型我第一反应就是模型不能太大推理速度得够快还得保证准确率不能太拉胯。这个项目最终落地为Android Studio 开发移动端应用PyTorch 负责模型训练和导出花卉数据集选用牛津大学发布的 Flower 102 数据集并搭配预训练的 MobileNetV3 作为骨干网络做迁移学习。整套流程跑通之后App 在普通中端手机上单张图片推理耗时大约 80 到 120 毫秒Top-1 准确率在测试集上能到 92% 左右。如果你正打算做一个类似的移动端图像分类项目或者想把手头的 PyTorch 模型塞进 Android 手机里跑起来这篇文章应该能帮你省掉不少弯路。1.2 为什么是 Android Studio PyTorch 这个组合选 Android Studio 作为移动端开发环境是因为它在 Android 生态里没有对手自带的布局编辑器、模拟器、gradle 构建系统、性能分析工具都相当成熟。你不需要额外配置一堆乱七八糟的插件新建工程、调 UI、打日志、跑真机调试一条龙全在 IDE 里完成。PyTorch 负责的是模型训练侧。它在科研和工业界的使用率极高预训练模型库 torchvision 里直接能拉 MobileNetV3、ResNet50 这些现成的 backbone省去从零训练的时间成本。更重要的是PyTorch 官方提供了 PyTorch Android 库能把训练好的模型导出成 TorchScript 格式直接塞进 Android 工程里调用。这种训练用 PyTorch、部署用 PyTorch Mobile的统一技术栈避免了模型格式转换和算子兼容性的一堆麻烦。有人可能会问那为什么不直接用 TensorFlow Lite 或者 ONNX Runtime我的回答是如果你只是做一个中小规模的分类器PyTorch Mobile 的算子覆盖已经完全够用。而且项目维护起来心智负担小——训练脚本和移动端推理代码用的都是同一套 API 风格出了问题也好排查。2. 数据集选型与预处理细节2.1 为什么选 Flower 102 数据集在做花卉识别这类细粒度图像分类任务时数据集的选择直接决定项目的上限。我比较过几个常见的数据集Flower 102 包含 102 个类别的英国常见花卉每类 40 到 258 张图片总计 8189 张而 Oxford 102 之所以比很多自制数据集靠谱是因为它的类别划分是经过植物学专家确认的不同类别之间在颜色、形态、拍摄背景上存在较强的相似性比如虞美人和加州罂粟从外观上确实容易混淆非常适合用来检验分类器的细粒度识别能力。如果你不想用现成的公开数据集也可以自己爬图建数据集。但我必须劝你一句千万别贪图省事只用三五十张图一类那模型的泛化能力会差到让你怀疑人生。Flower 102 这个数据规模再搭配数据增强和迁移学习是性价比最高的组合。2.2 数据清洗和类别平衡处理下载下来的数据集不能直接丢给模型训练。原始 Flower 102 的图片分辨率不统一背景复杂度也各不相同有的图片甚至有明显的遮挡和模糊。我在做预处理时主要做了这么几件事第一统一尺寸。所有图片缩放到 224x224 像素。为什么是 224因为 MobileNetV3 的输入张量设计就是 224x224x3这是 ImageNet 分类任务的标准输入尺寸模型的全局平均池化层和全连接层都依赖这个尺寸。你如果改成 192 或 256 也可以但要重新微调模型收益并不大。第二过滤损坏图片。用 PIL 打开每一张图片捕获异常并记录文件路径遇到读取失败就剔除。这个操作看似基础但原始数据集里真的存在个别损坏文件你训练到一半因为一张坏图崩掉 session那心情真的是原地爆炸。第三类别平衡检查。Flower 102 各类别图片数量比较均衡基本都在 40 到 80 张左右不需要额外做重采样。但如果你用的是自己的数据集务必统计一下每个类别的样本数。若有类别比别的少一半以上建议用随机裁剪、水平翻转、色彩抖动这类的在线增强方式在训练时动态补充样本而不是粗暴地复制粘贴原始图片。2.3 数据增强策略与归一化参数数据增强是防止过拟合的重中之重。Flower 102 每类只有几十张原始图如果直接训练一个深层模型几乎必然过拟合。我在训练时这样做随机水平翻转概率 0.5随机旋转范围 ±15 度随机仿射变换scale 范围 0.8 到 1.2随机颜色抖动brightness 和 saturation 的扰动 ±20%RandomResizedCrop随机裁剪后缩放回 224x224归一化时使用 ImageNet 的均值和标准差mean [0.485, 0.456, 0.406]std [0.229, 0.224, 0.225]这里的归一化必须跟预训练模型的统计值一致否则 MobileNetV3 迁移学习的效果会大打折扣。很多人刚开始用预训练模型时都会忽略这一点直接把像素值除以 255 就丢进网络结果精度死活上不去还以为是模型的问题。验证集和测试集只做缩放和归一化不做任何随机增强保证评测结果的可复现性。3. 模型训练与性能优化3.1 骨干网络选型为什么是 MobileNetV3模型结构的选择核心要回答一个问题在移动端有限的计算资源下如何在精度和速度之间取得平衡。我对比了三个候选ResNet50、MobileNetV2、MobileNetV3。ResNet50 的 Top-1 精度确实最高但在手机上跑一次前向推理需要 300 到 500 毫秒而且模型文件超过 90MB安装包体积直接膨胀这个代价在移动场景下不太能接受。MobileNetV2 是折中之选但 MobileNetV3 在相同计算量下精度更优引入了神经网络结构搜索NAS和基于 swish 激活函数的改进在 ImageNet 上的表现比 V2 提升约 3% 到 5%模型体积却几乎一样。最终选择的是 MobileNetV3-Large输入尺寸 224x224、宽度乘子 1.0。这个配置在 CPU 上的推理耗时约 100 毫秒级别模型文件经过量化后只有 16MB 左右压缩到 APK 里的增量不大。对了如果你对识别速度的要求更高可以把宽度乘子调成 0.75 甚至 0.5模型会更小更快但精度会有 1 到 2 个百分点的下滑。以我实测下来的经验宽度乘子 1.0 是平衡点。3.2 训练超参数的设置逻辑迁移学习的套路是加载 ImageNet 预训练权重替换最后一层全连接分类器。具体来说MobileNetV3 原来的全连接层输出 1000 类我要改成 102 类。PyTorch 代码里这样写import torch import torchvision.models as models from torch import nn model models.mobilenet_v3_large(weightsmodels.MobileNet_V3_Large_Weights.IMAGENET1K_V1) num_features model.classifier[-1].in_features model.classifier[-1] nn.Linear(num_features, 102)接下来是训练参数。我用的是 AdamW 优化器初始学习率 3e-4批量大小 64训练 40 个 epoch。前 5 个 epoch 做 warmup将学习率从 1e-4 线性增加到 3e-4然后使用余弦退火调度器逐步衰减到 1e-5。权重衰减系数设为 1e-4。为什么用 AdamW 而不是传统的 SGD因为 AdamW 对学习率的敏感度更低在迁移学习场景下收敛更快。你完全不需要手动调整每个 layer 的学习率倍率一把梭就能跑到不错的精度。SGD 当然也能用但需要更细致的调参比如 momentum、weight decay 的搭配对新手不友好。训练过程中我还在全连接层前加了一个 dropout概率 0.2。MobileNetV3 本身自带 dropout但默认值偏保守在数据量不大的情况下额外加一点有助于抑制过拟合。训练时监控训练集和验证集的 loss如果训练 loss 持续下降但验证 loss 开始回升就把 best checkpoint 以验证集精度为准保留下来。3.3 精度评测与混淆分析最终在测试集上Top-1 准确率 91.8%Top-5 准确率 98.6%模型的参数量约 4.2M计算量约 218M FLOPs。这些数字在同级别移动端分类模型里属于正常偏上的水平。但我建议你练完模型不要只看整体准确率一定要看一下混淆矩阵。我跑了一遍发现大多数识别错误集中在几个形态相近的类别上雏菊daisy和蒲公英dandelion经常互相认错因为两者的花瓣形状和颜色都偏白色系秋麒麟草goldenrod和金盏花marigold也有明显混淆。这属于细粒度分类的典型难题解决思路主要有两个方向一是收集更多困难样本做难例挖掘二是引入注意力机制让模型更关注花蕊区域的特征而不是背景纹理。做毕设或工程落地时往这两个方向做优化是比较合理的延伸点。3.4 模型量化与转换导出训练完成后模型需要转成 Android 能跑的格式。PyTorch 官方推荐的是 TorchScript它相当于一个可序列化、可被 C/Java 调用的模型中间表示。导出代码如下model.eval() example_input torch.randn(1, 3, 224, 224) traced_model torch.jit.trace(model, example_input) traced_model.save(flower_mobilenetv3.pt)这里有个坑要特别注意torch.jit.trace是基于示例输入运行的如果你的模型里含有依赖输入数据形状或者数据分布的分支控制逻辑比如某些动态维度的操作trace 会漏掉这些路径。好在 MobileNetV3 是纯卷积结构不涉及动态控制流trace 完全够用。如果你后续换成一个带自注意力或者动态 mask 的模型建议改用torch.jit.script但 script 对代码写法要求更高很多 PyTorch 的 Pythonic 特性都不支持。导出后建议再做一步量化。我尝试了两种方案训练后动态量化post-training dynamic quantization和训练后静态量化post-training static quantization。动态量化只针对全连接层做 int8 转换精度损失小但加速有限静态量化需要校准数据集把权重和激活都量化到 int8推理速度提升更明显。在 Flower 102 这个任务上静态量化后 Top-1 准确率掉到 89.4%损失 2.4 个百分点换来的是推理速度从 100 毫秒降到 60 毫秒左右。考虑到移动端体验速度优先最终选择了量化后的模型。4. Android 端集成与实机调优4.1 新建工程与依赖配置在 Android Studio 里新建工程时选择 Empty Views Activity语言用 Kotlin最低支持的 API 级别设为 24Android 7.0。API 24 以下设备的市场占比已经很低没必要为此牺牲可用的 API 特性。在模块的 build.gradle 里加入 PyTorch Mobile 依赖dependencies { implementation org.pytorch:pytorch_android_lite:2.1.0 implementation org.pytorch:pytorch_android_torchvision_lite:2.1.0 }使用 lite 版本而非完整版可以减小 APK 体积。PyTorch Android 的完整版会包含大量 CPU 算子而 lite 版本裁剪了部分不常用的算子对 MobileNetV3 这种纯卷积模型来说完全够用。两个库加起来大约增加 8MB 体积可以接受。模型文件放在app/src/main/assets/目录下Android 的 AssetManager 可以直接读取这个路径下的文件不需要额外申请存储权限。4.2 图片加载与预处理流程移动端图像分类最容易被忽略的是预处理流程。很多人在电脑上用 PIL 或 OpenCV 处理图片习惯了到了 Android 上直接用 Bitmap 解码结果色域、缩放方式不对模型精度崩得稀里哗啦。正确的做法是用 BitmapFactory 解码图片后统一缩放到 224x224。这里有一个关键细节BitmapFactory 解码 JPEG 时可能会把图片的方向信息丢掉导致横向拍摄的照片旋转 90 度。需要在解码时读取 EXIF 中的 orientation 字段做旋转校正。接下来把 Bitmap 转成 FloatArray并做归一化。这里要注意 Android Bitmap 默认的颜色通道顺序是 RGBA而 PyTorch 模型输入是 RGB需要显式转换。同时 Bitmap 的像素值范围是 0 到 255要先除以 255 再应用 ImageNet 的均值和标准差。fun bitmapToFloatArray(bitmap: Bitmap): FloatArray { val width bitmap.width val height bitmap.height val pixels IntArray(width * height) bitmap.getPixels(pixels, 0, width, 0, 0, width, height) val floatArray FloatArray(3 * width * height) val mean floatArrayOf(0.485f, 0.456f, 0.406f) val std floatArrayOf(0.229f, 0.224f, 0.225f) for (i in pixels.indices) { val pixel pixels[i] val r ((pixel shr 16) and 0xFF) / 255.0f val g ((pixel shr 8) and 0xFF) / 255.0f val b (pixel and 0xFF) / 255.0f floatArray[i] (r - mean[0]) / std[0] floatArray[floatArray.size / 3 i] (g - mean[1]) / std[1] floatArray[2 * floatArray.size / 3 i] (b - mean[2]) / std[2] } return floatArray }注意这里的存储顺序PyTorch 的 Tensor 布局是 CHWChannel First所以 floatArray 要先把所有像素的 R 通道存完再存 G 通道最后存 B 通道。如果你按 NHWC 的顺序存模型输出会是完全乱掉的预测结果。4.3 模型加载与推理代码实现PyTorch Android 的推理代码非常简洁val module Module.load(assetFilePath(context, flower_mobilenetv3_quantized.pt)) val inputTensor Tensor.fromBlob(floatArray, longArrayOf(1, 3, 224, 224)) val outputTensor module.forward(IValue.from(inputTensor)).toTensor() val scores outputTensor.dataAsFloatArrayModule.load负责从 assets 目录加载模型文件Tensor.fromBlob把 FloatArray 包装成指定 shape 的 Tensorforward执行推理。返回的 scores 是 102 个浮点数对应 102 个类别的置信度。通过 softmax 转成概率分布后取最大值对应的类别索引再映射到花卉名称即可。这里有个经验点Model 的加载和初始化放在 Activity 的 onCreate 或 ViewModel 的 init 里只加载一次。不要把 load 放到点击事件里每次执行那会带来肉眼可见的卡顿。Thread.sleep 级别的延迟很容易让用户直接卸载你的 App。4.4 推理速度优化与内存管理如果模型在低端手机上推理耗时太长可以从几个方向优化线程数和内存占用。PyTorch Mobile 默认使用单线程推理但你可以手动指定线程数PyTorchAndroid.setNumThreads(4)设置成 4 线程后推理时间大约能缩短 30% 到 40%。不过线程数不是越大越好手机 CPU 的核心调度策略很复杂设置 8 线程反而会因为线程切换开销导致性能回退。内存方面Bitmap 解码后的图片占用较大一张 4000x3000 的照片解码成 ARGB_8888 格式需要 48MB 内存。使用BitmapFactory.Options的inSampleSize做采样压缩比如采样率为 4图片降到 1000x750内存只需 3MB。然后再缩放至 224x224损失的信息可以忽略不计。另外一个容易被忽视的点是相机预览的实时识别。如果做的是实时取景识别建议不要每帧都跑模型而是每秒采样 3 到 5 帧做推理配合一个简单的帧差判断画面静止时不触发推理既能降低功耗又能减少 CPU 占用。5. 完整实操流程从零到一跑通这个项目5.1 环境准备与依赖安装先把基础环境捋一遍。开发机我用的 Ubuntu 20.04 CUDA 11.8如果你没有 NVIDIA 显卡纯 CPU 训练也可以跑只是时间会慢很多40 个 epoch 可能要从 30 分钟变成 6 个小时。Python 环境建议用 Anaconda 管理conda create -n flower python3.9 conda activate flower pip install torch2.0.1 torchvision0.15.2 pip install numpy opencv-python pillow matplotlibAndroid 侧的开发环境是 Android Studio 最新版SDK 安装 API 33 和 build-tools 33.0.1。IDE 装好后记得在 SDK Manager 里安装 Android SDK Command-line Tools后面可能用到 adb 调试和 logcat 查看命令行日志。5.2 数据集下载与目录整理Flower 102 数据集的下载需要从 Oxford 官网获取图片包和标签文件。下载完成后把图片放到dataset/jpg/目录下三个标注文件labels.txt、setid.mat、imagelabels.mat放到dataset/根目录。这里建议写一个脚本把原始文件按训练集、验证集、测试集拆分到独立目录同时生成一个classes.txt文件记录 102 个类别的名称。Android 端要显示中文花名可以在后续做一个英文到中文的映射表手工维护或者直接调用在线翻译 API 批量翻译。5.3 训练脚本与模型导出训练脚本的整体流程是定义数据集和数据增强、加载预训练模型、替换分类头、设置优化器、进入训练循环、验证并保存最优模型。核心代码部分前面已经给出这里补充两个容易被忽略的细节。第一个是随机种子固定。所有涉及随机的地方都要固定 seed包括 PyTorch、NumPy、Python 的 random否则你每次训练出来的结果会有波动复现困难。import random import numpy as np import torch seed 42 random.seed(seed) np.random.seed(seed) torch.manual_seed(seed) torch.cuda.manual_seed_all(seed)第二个是早停策略。当验证集精度连续 8 个 epoch 没有提升时停止训练并加载已保存的最佳模型。这个策略能省掉不少纯浪费时间的训练循环尤其在 40 个 epoch 的设定下通常到第 25 到 30 个 epoch 精度就已经趋于饱和。5.4 Android 工程搭建与页面设计移动端 App 的界面保持极简即可核心功能就两个选图识别和拍照识别。我用一个主页面上面是 ImageView 展示预览图下面是从相册选择和拍照识别两个按钮底部是结果展示区域TextView 显示花卉名称和置信度。拍照功能通过 CameraX 实现。调用takePicture的流程比老旧的 Camera2 API 省心很多生命周期和权限处理都封装好了。在 AndroidManifest.xml 里声明 CAMERA 权限运行时再动态申请一次。相册选择通过 ActivityResultContracts.GetContent 这个契约实现选择图片后拿到的是 content:// URI需要通过 ContentResolver 转成 Bitmap这个过程建议放到后台线程执行防止大图解码阻塞 UI。5.5 真机调试与性能数据实测我用一台骁龙 778G 的终端和一台骁龙 865 的开发机分别做了测试。实测结果如下设备量化前推理耗时量化后推理耗时Top-1 准确率骁龙 778G135ms82ms91.3%骁龙 86598ms58ms91.6%功耗方面连续运行 5 分钟推理电池温度上升约 3 到 4 摄氏度属于正常范围。APK 体积从原始的 85MB 压缩到 32MB其中模型文件和 PyTorch 库占了绝大部分。6. 踩坑实录与问题排查清单6.1 模型加载路径错误导致崩溃Module.load的路径必须是 assets 目录下的相对路径不能传绝对路径。我在初次集成时试过传/data/user/0/...这种文件系统的绝对路径结果直接抛 FileNotFoundException。正确做法是使用 AssetManager 打开 assets 目录下的文件复制到临时文件后在从该路径加载。官方提供了assetFilePath的辅助方法直接调用即可。如果模型文件放在 assets 子目录里比如assets/models/flower.pt代码里要传models/flower.pt这个相对路径。6.2 推理结果全是同一类别的问题这是一个高概率踩坑点。我在训练好的模型接入 Android 后发现无论输入什么花输出的结果都是同一个类别。排查了半天最后发现是预处理时 FloatArray 的通道顺序错了——把 RGBA 当 RGB 用导致所有图片的 B 通道值变成了 A 通道的透明值。另一个常见原因是归一化参数不正确。有些图片加载库会把像素值从 0 到 255 自动缩放到 0 到 1你再手动除以 255数值范围就乱了。解决办法是在代码注释里写清楚这里的输入范围是 0 到 255防止自己忘了。6.3 集成后 APK 体积暴增PyTorch Android 的完整版依赖会引入不少无用的算子导致 APK 体积变得非常大。这时检查一下你的 build.gradle确认用的是pytorch_android_lite而不是pytorch_android。另外在打包配置中开启 abiFilters只保留 arm64-v8a也可以显著减小体积。android { defaultConfig { ndk { abiFilters arm64-v8a } } }绝大多数现代 Android 手机都是 arm64 架构只保留这一种架构是合理的选择。6.4 低端机型推理耗时明显增加如果你发现推理耗时在低端机上飙到 400 毫秒以上最优先检查的是线程数设置。另外确认系统是否开启了省电模式它会主动限制 CPU 频率。还有一个优化方向是模型输入的图片分辨率如果从 224x224 降到 192x192推理耗时大约能降低 20%但精度也会有 1 个百分点左右的下滑需要做权衡。6.5 问题排查速查表现象可能原因排查方向加载模型崩溃模型路径错误 / 模型文件损坏检查 .pt 文件大小、assets 路径输出全为同一类别通道顺序错误 / 归一化错误打印预处理后的 floatArray 前 10 个值推理耗时过长未设置多线程 / 图片过大设置 setNumThreads、inSampleSizeAPK 体积过大PyTorch 完整库 / 多架构未裁剪换成 lite 版本、abiFilters精度远低于训练结果预处理与训练不一致对比 PC 端推理代码与移动端的预处理逻辑7. 后续可扩展的方向这个项目做完之后还有几个值得往下走的延展点这里一次性分享给你。第一个是替换更轻量的骨干网络。MobileNetV3 只是起点现在还有 EfficientNet-Lite、GhostNet、MobileOne 这些新结构在同等 FLOPs 下精度更高。如果你是跟着最新论文走可以把骨干网络换成 MobileOne-S1模型体积更小速度也更快但需要你重新做一轮训练调参。第二个是尝试知识蒸馏。先用大模型 ResNet50 训练一个高精度 Teacher再用 MobileNetV3 做 Student蒸馏后的 MobileNetV3 往往能比直接训练的版本提升 2 到 3 个百分点。这个方案在移动端部署中非常实用也是很多工业级应用的标配做法。第三个是加入图像分割做精细化识别。花卉识别如果遇到一束花里有多个种类单纯分类器只能给出一个整体标签。如果你的应用场景涉及花束、花坛这类复杂画面考虑到后处理阶段叠加语义分割网络先分割出每朵花的区域再逐个分类这会是体验更好的落地方案。当然计算量会翻倍需要根据设备性能做取舍。第四个是接入更丰富的花卉名称展示。目前 App 里只显示了英文名和置信度如果你做的是中文用户市场推荐内置一份中文名数据库再配合花的产地、花期、养护要点等信息切到科普工具这个定位实用价值会高很多。写在最后的几点经验这个项目从开始到完全跑通整体耗时一周左右。如果只做训练和 PC 端验证其实两天就能搞定真正花时间的是 Android 端集成时的各种小坑。我的个人体感是移动端 AI 应用的最大成本不在模型训练而在工程化落地——预处理、线程调度、内存管理、ABI 裁剪每一个环节都可能让你的模型在真机上表现打折。如果你打算抄这个项目来练手我建议别着急改代码先把整条链路跑通确保在真机上能看到预测结果然后再逐步替换数据集、调整模型结构。否则的话改一个变量出现 bug你根本分不清问题是出在训练侧还是部署侧。先用最稳定的基线方案跑通全流程再追求优化这才是最省时间的方式。另外一个小技巧训练和部署两边最好各留一份日志输出工具。训练时记录每轮的 loss 和 acc部署时在关键位置打印预处理后的 tensor 数值和推理耗时两边一对照很多诡异问题能迅速定位。这个习惯我保持了很多年几乎每次项目翻车都是靠它救命。本文还有配套的精品资源点击获取