如何将基于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
相关产品推荐
相关产品推荐

