如何使用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
相关产品推荐
相关产品推荐

