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

如何将基于Google Object Detection API的SSD冻结图转为TensorFlow.js格式?

把Object Detection API的冻结SSD模型转成TensorFlow.js格式的完整步骤

我之前也踩过这个一模一样的坑——用Google Object Detection API训完自定义SSD模型,拿到冻结的.pb文件后,发现TF.js转换器只认SavedModel格式,官网的示例根本没法直接套用。其实解决起来分两步走:先把冻结图转成SavedModel,再转成TF.js的Web友好格式,下面给你一步步拆解清楚:

第一步:将冻结推理图转换为SavedModel格式

冻结图(.pb)本质是序列化的计算图,但缺少SavedModel要求的签名信息(也就是明确的输入输出映射定义),所以得写个小脚本补全这个关键信息。

1. 先确认模型的输入输出节点名称

SSD模型的常规命名是:

  • 输入节点:image_tensor:0
  • 输出节点:detection_boxes:0、detection_scores:0、detection_classes:0、num_detections:0

如果不确定自己的模型节点名,可以用这段代码打印所有节点名称(如果是TF1.x训练的模型,建议用tf.compat.v1兼容模式):

import tensorflow as tf

# 加载冻结图
graph_def = tf.compat.v1.GraphDef()
with open('./frozen_inference_graph.pb', 'rb') as f:
    graph_def.ParseFromString(f.read())

# 遍历打印所有节点名
for node in graph_def.node:
    print(node.name)

2. 编写转换脚本生成SavedModel

创建一个freeze_to_savedmodel.py文件,填入以下内容:

import tensorflow as tf
from tensorflow.python.saved_model import signature_constants
from tensorflow.python.saved_model import tag_constants

# 配置路径,根据你的实际路径修改
FROZEN_GRAPH_PATH = './frozen_inference_graph.pb'
SAVED_MODEL_DIR = './saved_model'

# 加载冻结图
graph = tf.Graph()
with graph.as_default():
    input_graph_def = tf.compat.v1.GraphDef()
    with open(FROZEN_GRAPH_PATH, 'rb') as f:
        input_graph_def.ParseFromString(f.read())
        tf.import_graph_def(input_graph_def, name='')

# 获取输入输出张量
image_tensor = graph.get_tensor_by_name('image_tensor:0')
detection_boxes = graph.get_tensor_by_name('detection_boxes:0')
detection_scores = graph.get_tensor_by_name('detection_scores:0')
detection_classes = graph.get_tensor_by_name('detection_classes:0')
num_detections = graph.get_tensor_by_name('num_detections:0')

# 导出SavedModel
with tf.compat.v1.Session(graph=graph) as sess:
    builder = tf.compat.v1.saved_model.builder.SavedModelBuilder(SAVED_MODEL_DIR)
    
    # 定义服务签名:明确输入输出映射
    signature = tf.compat.v1.saved_model.signature_def_utils.predict_signature_def(
        inputs={'image_tensor': image_tensor},
        outputs={
            'detection_boxes': detection_boxes,
            'detection_scores': detection_scores,
            'detection_classes': detection_classes,
            'num_detections': num_detections
        }
    )
    
    builder.add_meta_graph_and_variables(
        sess=sess,
        tags=[tag_constants.SERVING],
        signature_def_map={
            signature_constants.DEFAULT_SERVING_SIGNATURE_DEF_KEY: signature
        }
    )
    builder.save()

运行这个脚本后,./saved_model目录下就会生成标准的SavedModel格式文件。

第二步:将SavedModel转换为TensorFlow.js格式

这一步就简单了,用官方的tensorflowjs_converter工具即可。首先确保你已经安装了依赖:

pip install tensorflowjs

然后执行转换命令,注意指定输入形状(SSD模型通常用[1, None, None, 3]支持动态尺寸输入):

tensorflowjs_converter \
    --input_format=tf_saved_model \
    --output_node_names='detection_boxes,detection_scores,detection_classes,num_detections' \
    --signature_name=serving_default \
    --saved_model_tags=serve \
    ./saved_model \
    ./tfjs_model

执行完成后,./tfjs_model目录下会生成model.json和若干分片权重文件,这就是可以直接在Web项目中使用的TF.js格式模型了。

一些踩坑提示

  • 如果你的模型是用TF1.x训练的,转换时尽量用TF2.x的兼容模式(也就是脚本里的tf.compat.v1),避免版本不兼容报错;
  • 一定要确认输出节点名称和你冻结图里的完全一致,否则会出现找不到节点的错误;
  • 如果模型体积太大,转换时可以加上--quantize_uint8参数做8位量化,大幅减小模型体积,更适合Web部署。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.22 10:08:55