You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

导出后如何正确加载TensorFlow Object Detection模型执行检测

TensorFlow目标检测导出模型正确加载推理方法

你加载得到的结果异常,是导出的SavedModel的输入要求、输出格式和训练阶段加载checkpoint的推理逻辑存在差异,按以下步骤操作即可正常使用:

正确操作流程

1. 确认加载路径

你导出的模型结构中,实际可加载的SavedModel存放在my_model/saved_model/目录下,加载时需指定到该子目录(目录内需包含saved_model.pb文件和variables文件夹),不要传入my_model根目录。

2. 完整加载+推理代码示例

import tensorflow as tf
import cv2
import numpy as np

# 1. 加载导出模型
PATH_TO_SAVED_MODEL = "./my_model/saved_model"
detect_fn = tf.saved_model.load(PATH_TO_SAVED_MODEL)

# 2. 图片预处理(不需要自定义归一化)
image_path = "test.jpg"
# 读取图片并转RGB格式(OpenCV默认读取为BGR)
img = cv2.imread(image_path)
img_rgb = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)
# 转为张量并添加batch维度,直接传入0-255范围的uint8数据即可
input_tensor = tf.convert_to_tensor(img_rgb, dtype=tf.uint8)
input_tensor = input_tensor[tf.newaxis, ...]

# 3. 执行推理
detections = detect_fn(input_tensor)

# 4. 推理结果后处理
# 去掉batch维度,提取有效检测结果
num_detections = int(detections.pop('num_detections'))
detections = {k: v[0, :num_detections].numpy() for k, v in detections.items()}
detections['num_detections'] = num_detections

# 将归一化边界框转为实际像素坐标
h, w = img.shape[:2]
detections['detection_boxes'] = (detections['detection_boxes'] * [h, w, h, w]).astype(np.int32)
# detection_classes默认从1开始计数,可直接对应你的label map的id值

常见异常原因排查

  • 输入了归一化到01范围的float类型张量:导出模型内置了归一化逻辑,仅接受0255范围的uint8输入
  • 未转换边界框坐标:模型输出的detection_boxes为[ymin, xmin, ymax, xmax]格式的归一化值,需要乘以图片实际高宽才能得到像素坐标
  • 通道顺序错误:没有将OpenCV读取的BGR格式转为RGB,会导致检测精度大幅下降
  • 类别索引混淆:输出的detection_classes是从1开始的id,不是从0开始的数组索引,需要和训练时的label map对应取值

内容的提问来源于stack exchange,提问作者Matheus Correia

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.09.26 13:36:05