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

TensorFlow预训练Faster R-CNN模型加载与检测结果异常求助

我来帮你一步步梳理这些问题,解决目标检测模型加载和结果异常的困惑:

一、先说说你最初的模型加载问题

你原来的代码有两个关键问题:

  1. 你解析了graph_def但没有把它导入到TensorFlow的Graph对象中,直接调用graph.get_tensor_by_name会找不到对应的节点,运行时肯定会报错。
  2. 你没必要加载model.ckpt.data文件——如果用frozen_inference_graph.pb,这个文件已经包含了所有冻结的权重,不需要单独加载ckpt文件;如果用saved_model目录,直接加载整个saved_model即可。
二、如何确定模型的输入输出层?

有几个简单的方法可以快速确认:

  • 用saved_model_cli命令行工具:这是最直接的方式,在终端运行:
    saved_model_cli show --dir ./faster_rcnn_inception_resnet_v2_atrous_lowproposals_oid_2018_01_28/saved_model --all
    
    它会输出所有签名的输入输出张量名称、形状和类型,比如你会看到输入是inputs(形状[1, None, None, 3]),输出包括detection_boxes、detection_scores、detection_classes、num_detections。
  • 打印图中的节点:如果用冻结图,你可以在导入graph_def后,打印所有Placeholder和输出相关的节点:
    with detection_graph.as_default():
        for n in detection_graph.as_graph_def().node:
            if n.op == 'Placeholder':
                print(f"输入节点:{n.name}")
            if 'detection' in n.name or 'num_detections' in n.name:
                print(f"输出节点:{n.name}")
    
  • 用TensorBoard可视化:把冻结图导入TensorBoard,就能直观看到整个图的结构和节点连接:
    tensorboard --logdir=./graphs --port=6006
    
    (需要先把pb文件转成TensorBoard能识别的格式,用tf.summary.FileWriter写入即可)
三、正确获取目标类别、置信度和检测框的代码示例

下面给你两种可靠的实现方式,对应你手里的两种模型文件:

方式一:使用frozen_inference_graph.pb

import tensorflow as tf
import cv2
import numpy as np

model_folder = "faster_rcnn_inception_resnet_v2_atrous_lowproposals_oid_2018_01_28"
model_graph_file = model_folder + "/frozen_inference_graph.pb"

# 加载冻结图到Graph对象
detection_graph = tf.Graph()
with detection_graph.as_default():
    od_graph_def = tf.GraphDef()
    with tf.gfile.GFile(model_graph_file, 'rb') as fid:
        serialized_graph = fid.read()
        od_graph_def.ParseFromString(serialized_graph)
        tf.import_graph_def(od_graph_def, name='')  # 关键:把graph_def导入到当前图

# 获取输入输出张量并运行推理
with detection_graph.as_default():
    with tf.Session(graph=detection_graph) as sess:
        # 输入张量:模型接受[1, None, None, 3]的RGB图片
        image_tensor = detection_graph.get_tensor_by_name('image_tensor:0')
        # 输出张量集合
        detection_boxes = detection_graph.get_tensor_by_name('detection_boxes:0')
        detection_scores = detection_graph.get_tensor_by_name('detection_scores:0')
        detection_classes = detection_graph.get_tensor_by_name('detection_classes:0')
        num_detections = detection_graph.get_tensor_by_name('num_detections:0')

        # 图片预处理:cv2读入的是BGR,转成RGB;增加batch维度
        image = cv2.imread('resources/my_image.jpg')
        image_rgb = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)
        image_expanded = np.expand_dims(image_rgb, axis=0)

        # 执行推理
        (boxes, scores, classes, num) = sess.run(
            [detection_boxes, detection_scores, detection_classes, num_detections],
            feed_dict={image_tensor: image_expanded})

        # 处理结果:去掉多余的batch维度,过滤低置信度结果
        boxes = np.squeeze(boxes)
        scores = np.squeeze(scores)
        classes = np.squeeze(classes)
        valid_detections = int(num[0])

        # 只保留置信度>0.5的有效结果
        threshold = 0.5
        print(f"找到{sum(scores[:valid_detections] > threshold)}个有效目标:")
        for i in range(valid_detections):
            if scores[i] > threshold:
                # 检测框格式是[y_min, x_min, y_max, x_max],范围0~1,对应图片的相对坐标
                box = boxes[i]
                print(f"类别ID:{classes[i]},置信度:{scores[i]:.4f},检测框:{box}")

方式二:使用saved_model目录(适配TF1.12)

import tensorflow as tf
import cv2
import numpy as np

model_folder = "faster_rcnn_inception_resnet_v2_atrous_lowproposals_oid_2018_01_28/saved_model"

with tf.Session() as sess:
    # 加载saved_model,注意TF1.12需要传标签列表
    model = tf.saved_model.loader.load(sess, [tf.saved_model.tag_constants.SERVING], model_folder)
    # 获取默认服务签名
    signature = model.signature_def['serving_default']

    # 获取输入输出张量名称
    input_tensor_name = signature.inputs['inputs'].name
    output_names = [
        signature.outputs['detection_boxes'].name,
        signature.outputs['detection_scores'].name,
        signature.outputs['detection_classes'].name,
        signature.outputs['num_detections'].name
    ]

    # 获取张量对象
    image_tensor = sess.graph.get_tensor_by_name(input_tensor_name)
    boxes_tensor = sess.graph.get_tensor_by_name(output_names[0])
    scores_tensor = sess.graph.get_tensor_by_name(output_names[1])
    classes_tensor = sess.graph.get_tensor_by_name(output_names[2])
    num_tensor = sess.graph.get_tensor_by_name(output_names[3])

    # 图片预处理和推理
    image = cv2.imread('resources/my_image.jpg')
    image_rgb = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)
    image_expanded = np.expand_dims(image_rgb, axis=0)

    boxes, scores, classes, num = sess.run(
        [boxes_tensor, scores_tensor, classes_tensor, num_tensor],
        feed_dict={image_tensor: image_expanded}
    )

    # 结果处理和上面一致
    boxes = np.squeeze(boxes)
    scores = np.squeeze(scores)
    classes = np.squeeze(classes)
    valid_detections = int(num[0])

    threshold = 0.5
    print(f"找到{sum(scores[:valid_detections] > threshold)}个有效目标:")
    for i in range(valid_detections):
        if scores[i] > threshold:
            print(f"类别ID:{classes[i]},置信度:{scores[i]:.4f},检测框:{boxes[i]}")
四、关于你调整代码后结果差异大的问题

你现在的输出看起来有有效结果,但和预期差异大,大概率是这几个原因:

  1. 图片格式没转换:你可能没把cv2读的BGR图片转成RGB,预训练模型都是基于RGB数据训练的,BGR输入会导致检测结果异常(比如置信度偏低、类别识别错误)。
  2. 类别ID的含义误解:你用的是OID(Open Images Dataset)的预训练模型,它的类别ID和COCO数据集完全不同,比如33、68都是OID里的特定类别,你需要下载对应模型的label_map.pbtxt文件,把ID映射成实际的类别名称(比如“Dog”“Car”)。
  3. 没过滤低置信度结果:模型默认输出20个检测框,其中大部分置信度极低(甚至为0),你需要用阈值过滤后,只看置信度>0.5的结果,那些0值的框都是无效的,可以忽略。
  4. 检测框坐标的含义:模型输出的检测框是[y_min, x_min, y_max, x_max],是相对于图片尺寸的归一化坐标(范围0~1),如果要转成像素坐标,需要乘以图片的高度和宽度。

你可以按照上面的代码调整预处理步骤,再过滤低置信度结果,应该就能得到和预期一致的检测结果了。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.06 22:07:46