如何将TF2的saved_model.pb转为TF1的frozen_inference_graph.pb以适配OpenCV?
解决TF2 SSD模型转TF1兼容格式适配OpenCV的方法
一、完整的模型冻结与导出流程
之前的冻结步骤不完整,未明确输入输出节点信息,导致OpenCV无法正确解析模型。以下是标准转换流程:
1. 加载TF2模型并获取推理签名
import tensorflow as tf from tensorflow.python.framework.convert_to_constants import convert_variables_to_constants_v2 # 加载TF2 SavedModel格式模型 model = tf.saved_model.load("path/to/your/tf2_saved_model") infer_signature = model.signatures["serving_default"]
2. 固定输入维度并生成冻结图
TF2模型常含动态维度(如batch_size为None),需固定为单张推理的维度(batch_size=1),避免OpenCV解析失败:
# 获取输入张量的原始形状,替换动态batch_size为1 input_shape = infer_signature.inputs[0].shape.as_list() input_shape[0] = 1 # 固定batch_size为1 input_spec = tf.TensorSpec(input_shape, infer_signature.inputs[0].dtype, name="input_tensor") # 生成可导出的concrete function concrete_func = infer_signature.get_concrete_function(input_spec) # 冻结图,将变量转为常量 frozen_func = convert_variables_to_constants_v2(concrete_func) frozen_graph = frozen_func.graph # 保存冻结后的.pb模型文件 with tf.io.gfile.GFile("frozen_ssd_model.pb", "wb") as f: f.write(frozen_graph.as_graph_def().SerializeToString())
3. 生成OpenCV所需的.pbtxt配置文件
OpenCV需要明确的输入输出节点定义,先打印节点名称:
# 打印输入节点名称 print("输入节点名称:", frozen_graph.get_operations()[0].name) # 打印检测相关输出节点(通常包含detection_boxes、scores等关键词) for op in frozen_graph.get_operations(): if op.type == "Identity" and any(key in op.name for key in ["detection_boxes", "detection_scores", "detection_classes", "num_detections"]): print("输出节点:", op.name)
根据打印结果创建frozen_ssd_model.pbtxt文件,示例内容如下(替换为你实际的节点名称):
model { node { name: "input_tensor" op: "Placeholder" attr { key: "dtype" value { type: DT_FLOAT } } attr { key: "shape" value { shape { dim { size: 1 } dim { size: 640 } dim { size: 640 } dim { size: 3 } } } } } node { name: "StatefulPartitionedCall/detection_boxes" op: "Identity" input: "StatefulPartitionedCall/Postprocessor/ExpandDims" } node { name: "StatefulPartitionedCall/detection_scores" op: "Identity" input: "StatefulPartitionedCall/Postprocessor/ExpandDims_1" } node { name: "StatefulPartitionedCall/detection_classes" op: "Identity" input: "StatefulPartitionedCall/Postprocessor/Cast" } node { name: "StatefulPartitionedCall/num_detections" op: "Identity" input: "StatefulPartitionedCall/Postprocessor/ExpandDims_2" } input: "input_tensor" output: "StatefulPartitionedCall/detection_boxes" output: "StatefulPartitionedCall/detection_scores" output: "StatefulPartitionedCall/detection_classes" output: "StatefulPartitionedCall/num_detections" }
二、OpenCV加载模型的修正代码
之前的报错核心是模型加载不完整或预处理不匹配,以下是正确的推理代码:
import cv2 import numpy as np # 加载冻结模型与配置文件 net = cv2.dnn.readNetFromTensorflow("frozen_ssd_model.pb", "frozen_ssd_model.pbtxt") # 预处理图像(匹配MobileNet V2训练时的归一化规则) cropped_img = cv2.imread("test_image.jpg") blob = cv2.dnn.blobFromImage( cropped_img, scalefactor=1.0/127.5, # 对应训练时的(image/127.5)-1 size=(640, 640), mean=(127.5, 127.5, 127.5), swapRB=True, crop=False ) net.setInput(blob) # 显式指定输出节点,避免OpenCV自动推断错误 output_layers = [ "StatefulPartitionedCall/detection_boxes", "StatefulPartitionedCall/detection_scores", "StatefulPartitionedCall/detection_classes", "StatefulPartitionedCall/num_detections" ] detections = net.forward(output_layers) # 解析检测结果 boxes = detections[0][0] scores = detections[1][0] classes = detections[2][0] num_detections = int(detections[3][0][0]) # 过滤高置信度结果并绘制 conf_threshold = 0.5 h, w = cropped_img.shape[:2] for i in range(num_detections): if scores[i] > conf_threshold: # 将归一化坐标转换为图像像素坐标 x1 = int(boxes[i][1] * w) y1 = int(boxes[i][0] * h) x2 = int(boxes[i][3] * w) y2 = int(boxes[i][2] * h) # 绘制框与标签 cv2.rectangle(cropped_img, (x1, y1), (x2, y2), (0, 255, 0), 2) label = f"Class: {int(classes[i])}, Conf: {scores[i]:.2f}" cv2.putText(cropped_img, label, (x1, y1-10), cv2.FONT_HERSHEY_SIMPLEX, 0.5, (0,255,0), 2) cv2.imshow("Detections", cropped_img) cv2.waitKey(0) cv2.destroyAllWindows()
三、常见问题排查
- 模型加载失败:检查.pb和.pbtxt文件路径是否正确,确保.pbtxt中的输入输出节点名称与冻结图完全一致。
- 预处理不匹配:MobileNet V2训练时通常使用
(image/127.5)-1的归一化规则,若使用错误的mean和scalefactor会导致模型输出异常。 - 动态维度问题:冻结图时必须固定batch_size为1,OpenCV的dnn模块对动态维度支持有限。
内容的提问来源于stack exchange,提问作者Slava Zhulanov
相关产品推荐
相关产品推荐

