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

如何将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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.27 13:32:53