首页
/ 在移动端运行 TF2 目标检测模型:TensorFlow Lite 转换与 Android 部署实战指南

在移动端运行 TF2 目标检测模型:TensorFlow Lite 转换与 Android 部署实战指南

2026-09-06 18:59:01作者:牧宁李

本指南以 TensorFlow Object Detection API(当前仓库 research/object_detection)中的官方文档 running_on_mobile_tf2.md 为骨架,讲解如何将 TF2 Detection Zoo 中符合要求的 SSD 目标检测模型导出为可供 TensorFlow Lite Converter 消费的中间 SavedModel,再转换成 .tflite 模型,并最终集成进 Android 应用完成端侧推理。读者将掌握「模型导出 → TFLite 量化转换 → 元数据打包 → Android 集成与真机验证」的完整链路,并理解 TFLite_Detection_PostProcess 自定义算子被注入的底层原理。

适用范围与前提:SSD 架构是端侧主力

TensorFlow Lite 是 TensorFlow 面向移动端与嵌入式设备的轻量级推理方案,通过量化 kernel 与定点运算获得低延迟和小体积模型。不过,本文档讨论的框类(boxes-based)检测模型转换目前仅支持 SSD 元架构(meta-architecture),且排除 EfficientDet。具体约束如下:

  • SSD 系模型(含 MobileNetV1/V2/V3、ResNet、Inception 等 backbone 组合)可直接走本文的转换流程;
  • EfficientDet 不在本工具链支持范围内,官方建议改用 TFLite Model Maker 的目标检测库完成转换;
  • CenterNet 支持仍处于实验阶段,相关流程与限制请直接参考仓库内的 centernet_on_device.ipynb

此外,模型在 mobile 上能否达到理想效果,与它在 Detection Zoo 中的原始结构强相关。可对照 tf2_detection_zoo.md 挑选端侧友好的 SSD checkpoint(例如以 ssd_mobilenet_* 开头的配置)作为转换输入。如果你想走一条「从微调到 TFLite」的端到端 Python 路线,仓库中还提供了 eager_few_shot_od_training_tflite.ipynb 以及图文并茂的 convert_odt_model_to_TFLite.ipynb 供直接运行。

转换后模型的输入 / 输出契约

转换产物的接口被严格固定,方便下游按统一约定解析。对于 SSD 模型,输出一个浮点输入与四个输出张量:

One input:
  image: a float32 tensor of shape [1, height, width, 3] containing the
  *normalized* input image.
  NOTE: 归一化逻辑定义在各 feature extractor 类内(见下文源码分析)

Four Outputs:
  detection_boxes:    float32 [1, num_boxes, 4]       框坐标
  detection_classes:  float32 [1, num_boxes]           类别索引
  detection_scores:   float32 [1, num_boxes]           类别得分
  num_boxes:          float32 标量(size 1)            有效检测框数量

需要强调的是输入必须是归一化后的图像,且尺寸与训练时 pipeline 中配置的 fixed_shape_resizer 保持一致——因为导出脚本会将输入固定为静态单 batch 形状 [1, height, width, channels]。归一化方式由各特征提取器的 preprocess 实现决定,例如 ssd_mobilenet_v1_feature_extractor.py 明确注释了「Maps pixel values to the range [-1, 1]」,即以 [0, 255] 的原始像素除以 127.5 再减 1 得到 [-1, 1] 区间。这一点在后端读取输入图像前必须手动对齐,否则检测质量会明显劣化。

三步完成模型转换

Step 1:导出可供 TFLite 转换的中间 SavedModel

转换第一步由仓库脚本 export_tflite_graph_tf2.py 完成。它读入 pipeline 配置与训练 checkpoint,生成一个中间 SavedModel(通常位于 output_directory/saved_model),供后续 TFLite Converter 通过命令行或 Python API 消费。基本用法如下(假设在仓库 research/ 目录下执行):

# 在 tensorflow/models 仓库的 research/ 目录下执行
python object_detection/export_tflite_graph_tf2.py \
    --pipeline_config_path path/to/ssd_model/pipeline.config \
    --trained_checkpoint_dir path/to/ssd_model/checkpoint \
    --output_directory path/to/exported_model_directory

运行 python object_detection/export_tflite_graph_tf2.py --help 可查看全部参数。下表基于 export_tflite_graph_tf2.py 的 flags 定义整理,用于按需调整导出精度与速度:

参数 类型 默认值 说明
--pipeline_config_path string 必填 pipeline_pb2.TrainEvalPipelineConfig 文本配置文件路径
--trained_checkpoint_dir string 必填 训练好的 checkpoint 所在目录
--output_directory string 必填 输出目录(不存在会自动创建)
--config_override string '' 文本 proto 形式的配置覆盖,用于微调推理图(见下文示例)
--max_detections int 10 返回的最大检测框数量
--ssd_use_regular_nms bool False SSD 专用:后处理改用 Regular NMS 而非 Fast NMS
--centernet_include_keypoints bool False CenterNet 专用:是否一并导出关键点张量
--keypoint_label_map_path string None CenterNet 关键点任务的 label map 路径,会替换 pipeline 中的同名配置

脚本内部将 config_overridepipeline_config_path 指向的配置合并(先解析后者,再 MergeFrom 前者),然后调用核心库 export_tflite_graph_lib_tf2.export_tflite_model(...) 真正完成导出。config_override 适合在不改动训练配置的前提下调整推理图行为,比如把 NMS 阈值调得更宽松以获得更多候选框:

python object_detection/export_tflite_graph_tf2.py \
    --pipeline_config_path path/to/ssd_model/pipeline.config \
    --trained_checkpoint_dir path/to/ssd_model/checkpoint \
    --output_directory path/to/exported_model_directory \
    --config_override "
        model{
        ssd{
          post_processing {
            batch_non_max_suppression {
                    score_threshold: 0.0
                    iou_threshold: 0.5
            }
         }
      }
   }
   "

导出背后的源码原理(可选精读)

理解 export_tflite_graph_lib_tf2.py 有助于排查问题:

  • 脚本按 pipeline 中 model 字段属于 ssd 还是 center_net 分别构造 SSDModuleCenterNetModule,其余架构会直接抛出 ValueError(见 export_tflite_graph_lib_tf2.py)。
  • 只支持 fixed_shape_resizerSSDModule._process_config 中若检测到其他 resizer 会报错;读取到的 height/width 以及灰度开关决定了输入形状,最终通过 input_shape() 返回 [1, height, width, channels](见 export_tflite_graph_lib_tf2.py)。
  • TF 图本身不执行真正的 NMS:因为 TF 侧没有与 TFLite 自定义后处理算子等价的实现,库中用一个 tf.function 包裹的空壳复合函数并打上 experimental_implements 签名 TFLite_Detection_PostProcess(见 export_tflite_graph_lib_tf2.py)。签名中携带 max_detectionsmax_classes_per_detection=1use_regular_nmsnms_score_thresholdnms_iou_thresholdx/y/h/w_scalenum_classes 等属性,TFLite 转换器通过 MLIR 的 legalization 将其改写为真正运行在端上的自定义 NMS 算子。
  • anchor 被固化为常量get_const_center_size_encoded_anchors(见 export_tflite_graph_lib_tf2.py)把来自 predicted_tensors['anchors'] 的中心尺寸编码 anchor 转成 shape 为 [num_anchors, 4]float32 常量节点,避免推理期动态计算。
  • 上述实现细节均有单元测试背书,见 export_tflite_graph_lib_tf2_test.py:其中既有针对 postprocess_implements_signature 生成的签名串的断言(如 test_postprocess_implements_signature),也有验证导出产物确实是 SavedModel、并对导出模型执行推理断言输出个数(SSD 为 4 个张量)的用例(如 test_exported_model_inference)。

Step 2:使用 TFLite Converter 转换为 .tflite 模型

拿到中间 SavedModel 后,用 TensorFlow Lite Converter 完成最终转换。注意 Python API 必须使用 from_saved_model 接口。最小可运行片段如下:

import tensorflow as tf

converter = tf.lite.TFLiteConverter.from_saved_model(
    'path/to/exported_model_directory/saved_model')
tflite_model = converter.convert()

with open('detect.tflite', 'wb') as f:
    f.write(tflite_model)

后训练量化(Post-training Quantization)

为了获得更小的体积与更快的推理,可叠加后训练量化。该能力仅在 Python API 中可用,且建议配合代表性数据集(representative dataset)做全整数量化。在原文档给出代码的基础上,补充说明其中每个选项的作用:

converter.optimizations = [tf.lite.Optimize.DEFAULT]
converter.target_spec.supported_ops = [tf.lite.OpsSet.TFLITE_BUILTINS_INT8,
                                       tf.lite.OpsSet.TFLITE_BUILTINS]
converter.representative_dataset = <...>  # 使用代表性数据集的生成器
  • optimizations = [tf.lite.Optimize.DEFAULT]:启用默认的量化优化策略;
  • target_spec.supported_ops = [TFLITE_BUILTINS_INT8, TFLITE_BUILTINS]:要求优先选用 INT8 内置算子,同时保留浮点内置算子作为兜底,从而把量化不适用的部分算子留在浮点;
  • representative_dataset:必须提供一个能产出若干张「具有代表性」输入样本的生成器,供转换器统计激活值范围,是全整数量化获取良好精度损失控制的关键。若你的模型是全量化模型,Android 端会以 TF_OD_API_IS_QUANTIZED = true 来驱动对应预处理分支(见下文)。

Step 3:为模型附加 Metadata 并打包标签文件

为便于在移动端直接使用,需要为 .tflite 模型写入 metadata,并把配套的标签(labels)文件一起打包进模型(pack associated files)。这样下游可以通过统一的 Task Library 自动解析类别名等信息,无需在业务代码里手工读取外部文本。若确需了解 Metadata 的写入细节与 image classification 侧的参考实现,仓库内没有内嵌该工具代码,请以官方 TensorFlow Lite 的 metadata writer 文档与示例工程为准;convert_odt_model_to_TFLite.ipynb 教程中也演示了 labels 与模型的打包配套过程。

在 Android 上运行检测模型

用 Task Library 的 ObjectDetector API 集成模型

推荐用 TensorFlow Lite Task Library 提供的 ObjectDetector API 把上一步产出的模型集成进 Android 应用,它会自动处理输入归一化、输出解析以及 metadata 读取。核心 Java 调用只有两段:

// 初始化
ObjectDetectorOptions options = ObjectDetectorOptions.builder().setMaxResults(1).build();
ObjectDetector objectDetector = ObjectDetector.createFromFileAndOptions(context, modelFile, options);

// 执行推理
List<Detection> results = objectDetector.detect(image);

setMaxResults(1) 表示每帧只保留得分最高的一个检测结果,可按需调大;modelFile 指向你的 detect.tflite 文件。由于模型已打包 labels(Step 3),Detection 结果中可以直接拿到类别名称。

用官方示例 App 真机验证

在没有自建工程的情况下,可以先用 TensorFlow Lite 官方示例工程中的目标检测 App(位于 tensorflow/examples 仓库的 lite/examples/object_detection,Android 子工程在 lite/examples/object_detection/android)做真机验证。该工程支持用 Android Studio 构建运行,需要安装支持 API >= 21 的 Android SDK 与 build tools。

  1. 把模型拷入 assets 目录。在 Android 工程根目录执行:
mkdir -p $TF_EXAMPLES/lite/examples/object_detection/android/app/src/main/assets
cp /tmp/tflite/detect.tflite \
  $TF_EXAMPLES/lite/examples/object_detection/android/app/src/main/assets

注意:labels 文件应当在上一步已打包进模型(而非常规地单独放在 assets 中),原文档强调这一点很重要。

  1. 修改 gradle 构建文件,避免资产被覆盖。打开 $TF_EXAMPLES/lite/examples/object_detection/android/app/build.gradle,注释掉自动下载模型的脚本行,防止构建时下载的默认模型覆盖你刚拷贝进去的文件:
// apply from: 'download_model.gradle'
  1. 核对模型与量化开关配置。只要模型命名为 detect.tflite 且位于 assets 根目录,示例 App 会自动加载;若使用自定义路径或文件名,则需编辑 DetectorActivity.java,定位 TF_OD_API_MODEL_FILETF_OD_API_LABELS_FILETF_OD_API_IS_QUANTIZED 三个常量。量化模型须将 TF_OD_API_IS_QUANTIZED 置为 true,浮点模型置为 false,例如量化模型对应的配置片段:
  private static final boolean TF_OD_API_IS_QUANTIZED = true;
  private static final String TF_OD_API_MODEL_FILE = "detect.tflite";
  private static final String TF_OD_API_LABELS_FILE = "labels_list.txt";

完成模型拷贝与 gradle 脚本调整后,即可按 Android Studio 常规流程构建并部署到设备。注意默认下载脚本被注释后,示例工程不再自行拉取模型,请确认 assets 目录中存在你的 .tflite(以及你的模型是否仍需要 labels 文件单独放置在 assets 根目录,这取决于你在 Step 3 中是否真正把标签打包进了模型)。

常见问题与提示

  • 「仅支持 SSD」的错误:若 pipeline 配置中的 model 字段既不是 ssd 也不是 center_netexport_tflite_graph_lib_tf2.py 会抛出 ValueError。此时请回到 Detection Zoo 选取 SSD 系模型,或改用 Model Maker 处理 EfficientDet。
  • 输入归一化不一致:SSD 端侧模型期望 [-1, 1] 归一化输入,而部分 CenterNet 导出则内含预处理、直接接收原始像素值。务必与目标模型的 preprocess 语义保持一致。
  • 推理输出顺序:SSD 导出最终对外暴露四个张量 detection_boxes / detection_classes / detection_scores / num_boxes;而 SSDModule.inference_fn 内部因 tf.function 会反转输入顺序,特意在返回前对结果做了一次反转(见 export_tflite_graph_lib_tf2.py)。集成时请以本文第一节给出的契约为准解析,不要直接套用 TF 训练图的字段顺序。
  • CenterNet 关键点导出:通过 --centernet_include_keypoints true--keypoint_label_map_path 可让导出结果额外包含关键点坐标与置信度(共 6 路输出),对应测试用例也验证了这一分支(见 export_tflite_graph_lib_tf2_test.py)。
  • TF1 迁移用户:仓库还保留了一份面向旧版流程的文档 running_on_mobile_tensorflowlite.md,其导图路径与本文的 TF2 方式不同,请按自身所使用框架版本(TF1 或 TF2)选择对应的转换脚本。
登录后查看全文
热门项目推荐
相关项目推荐