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

如何使用TensorFlow 2的冻结推理图实现图像目标检测?

用TensorFlow 2加载冻结目标检测推理图并实现检测

你已经成功加载了冻结图,并且找到了关键的输出张量,这一步走对了!接下来咱们一步步完成图像检测:

1. 先搞懂张量名称的含义

你打印出的最后几个张量,就是目标检测模型的核心输出:

  • prefix/detection_boxes:0: 每个检测框的归一化坐标(格式是[y_min, x_min, y_max, x_max],数值范围0-1,对应图像的相对位置)
  • prefix/detection_scores:0: 每个检测结果的置信度(0-1之间,数值越高越可信)
  • prefix/detection_classes:0: 检测到的类别ID(和你训练时的类别映射对应)
  • prefix/num_detections:0: 模型输出的有效检测结果总数

而模型需要的输入张量通常是prefix/image_tensor:0(这是TensorFlow目标检测API的默认输入,你可以在之前打印的op列表里确认是否存在,大概率是有的)。

2. 编写完整的推理代码

基于你已有的加载代码,咱们扩展成完整的检测流程:

第一步:加载冻结图并获取张量引用

import tensorflow as tf
import numpy as np
from PIL import Image, ImageDraw

def load_graph(frozen_graph_filename):
    with tf.io.gfile.GFile(frozen_graph_filename, "rb") as f:
        graph_def = tf.compat.v1.GraphDef()
        graph_def.ParseFromString(f.read())
    with tf.Graph().as_default() as graph:
        tf.import_graph_def(graph_def, name="prefix")
    return graph

# 加载你的冻结推理图
g = load_graph("inference_graph/frozen_inference_graph.pb")

# 获取输入和输出张量的引用
input_tensor = g.get_tensor_by_name('prefix/image_tensor:0')
detection_boxes = g.get_tensor_by_name('prefix/detection_boxes:0')
detection_scores = g.get_tensor_by_name('prefix/detection_scores:0')
detection_classes = g.get_tensor_by_name('prefix/detection_classes:0')
num_detections = g.get_tensor_by_name('prefix/num_detections:0')

第二步:加载并预处理图像

模型要求输入是RGB格式、带batch维度的numpy数组(形状为[1, 高度, 宽度, 3]):

# 加载你的测试图像
img = Image.open("test_image.jpg")
# 转换成RGB格式的numpy数组(如果图像是RGBA格式会自动转成RGB)
img_np = np.array(img.convert("RGB"))
# 添加batch维度(模型要求输入必须带batch,哪怕只检测一张图)
input_img = np.expand_dims(img_np, axis=0)

第三步:运行推理

因为冻结图是TensorFlow 1.x格式的,在TF2里咱们用兼容的会话来执行推理:

with tf.compat.v1.Session(graph=g) as sess:
    # 传入输入图像,获取检测结果
    boxes, scores, classes, num = sess.run(
        [detection_boxes, detection_scores, detection_classes, num_detections],
        feed_dict={input_tensor: input_img}
    )

第四步:处理并可视化结果

推理结果带batch维度,先去掉它,再过滤掉低置信度的结果:

# 去掉多余的batch维度
boxes = np.squeeze(boxes)
scores = np.squeeze(scores)
classes = np.squeeze(classes).astype(np.int32)
num_detections = int(np.squeeze(num))

# 设定置信度阈值,只保留可信的检测结果(可以自己调整,比如0.3或0.6)
conf_threshold = 0.5
valid_indices = [i for i in range(num_detections) if scores[i] > conf_threshold]

# 在图像上绘制检测框和标签
draw = ImageDraw.Draw(img)
for i in valid_indices:
    # 把归一化坐标转换成图像像素坐标
    y_min, x_min, y_max, x_max = boxes[i]
    left = int(x_min * img.width)
    top = int(y_min * img.height)
    right = int(x_max * img.width)
    bottom = int(y_max * img.height)
    
    # 画红色矩形框
    draw.rectangle([left, top, right, bottom], outline="red", width=2)
    # 添加类别和置信度标签
    label = f"类别{classes[i]}: {scores[i]:.2f}"
    draw.text((left, top - 10), label, fill="red")

# 保存或显示结果图像
img.save("detected_result.jpg")
img.show()

3. 可能的小问题排查

  • 如果找不到prefix/image_tensor:0:回到你打印的op列表,搜索包含image或input的张量,替换成对应的名称即可。
  • 图像颜色异常:如果用OpenCV加载图像(默认是BGR格式),需要转换成RGB(img_np = img_np[..., ::-1])。
  • 检测结果太少/太多:调整conf_threshold阈值,数值越低检测到的结果越多,反之越少。

内容的提问来源于stack exchange,提问作者Hammad Ilyas

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.08 13:27:41