TensorFlow预训练Faster R-CNN模型加载与检测结果异常求助
我来帮你一步步梳理这些问题,解决目标检测模型加载和结果异常的困惑:
一、先说说你最初的模型加载问题
你原来的代码有两个关键问题:
- 你解析了
graph_def但没有把它导入到TensorFlow的Graph对象中,直接调用graph.get_tensor_by_name会找不到对应的节点,运行时肯定会报错。 - 你没必要加载
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 --allinputs(形状[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,就能直观看到整个图的结构和节点连接:
(需要先把pb文件转成TensorBoard能识别的格式,用tensorboard --logdir=./graphs --port=6006tf.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]}")
四、关于你调整代码后结果差异大的问题
你现在的输出看起来有有效结果,但和预期差异大,大概率是这几个原因:
- 图片格式没转换:你可能没把cv2读的BGR图片转成RGB,预训练模型都是基于RGB数据训练的,BGR输入会导致检测结果异常(比如置信度偏低、类别识别错误)。
- 类别ID的含义误解:你用的是OID(Open Images Dataset)的预训练模型,它的类别ID和COCO数据集完全不同,比如33、68都是OID里的特定类别,你需要下载对应模型的
label_map.pbtxt文件,把ID映射成实际的类别名称(比如“Dog”“Car”)。 - 没过滤低置信度结果:模型默认输出20个检测框,其中大部分置信度极低(甚至为0),你需要用阈值过滤后,只看置信度>0.5的结果,那些0值的框都是无效的,可以忽略。
- 检测框坐标的含义:模型输出的检测框是
[y_min, x_min, y_max, x_max],是相对于图片尺寸的归一化坐标(范围0~1),如果要转成像素坐标,需要乘以图片的高度和宽度。
你可以按照上面的代码调整预处理步骤,再过滤低置信度结果,应该就能得到和预期一致的检测结果了。
内容的提问来源于stack exchange,提问作者Nakeuh
相关产品推荐
相关产品推荐

