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

如何获取输入输出名称以将TensorFlow目标检测模型转为.pb文件

解决TensorFlow Object Detection API模型ckpt转pb的输入输出节点获取问题

别着急,其实不用先拿到.pb文件才能用summarize_graph工具,咱们直接从现有的ckpt文件就能搞定输入输出节点名称,给你分享几个实用的方法:

方法1:用TensorBoard可视化计算图

这是最直观的方式,步骤很简单:

  • 启动TensorBoard,指定你训练时的模型存储目录(就是保存ckpt文件的那个文件夹):
    tensorboard --logdir=你的训练模型目录路径
    
  • 打开浏览器访问 http://localhost:6006,切换到Graphs标签页
  • 找输入节点:TensorFlow Object Detection API的标准模型默认输入节点是image_tensor(对应输入图片的张量)
  • 找输出节点:核心输出一般是这四个:detection_boxes(检测框坐标)、detection_scores(置信度)、detection_classes(类别ID)、num_detections(有效检测框数量)

方法2:用Python代码直接遍历ckpt的计算图节点

如果不想开TensorBoard,写几行代码就能列出所有节点名称:

import tensorflow as tf

# 加载ckpt的元数据文件
saver = tf.train.import_meta_graph('model.ckpt-10000.meta')
with tf.Session() as sess:
    saver.restore(sess, 'model.ckpt-10000')
    # 遍历并打印所有操作节点的名称
    for op in sess.graph.get_operations():
        print(op.name)

运行后,你会看到一堆节点名称,从中筛选出类似方法1里提到的输入输出节点即可。如果是自定义的模型结构,可能会有细微差别,但标准模型基本都是那几个默认名。

方法3:直接查看训练用的pipeline.config文件

你训练模型时用的pipeline.config里其实也藏着线索:

  • 输入部分:配置里的input_reader模块默认指定的输入就是image_tensor
  • 输出部分:Object Detection API的模型都会输出那四个核心检测张量,所以直接用detection_boxes,detection_scores,detection_classes,num_detections作为输出节点名完全没问题

拿到节点名后,用freeze_graph.py生成pb文件

有了输入输出节点名,就可以执行freeze_graph命令了,示例如下:

python freeze_graph.py \
    --input_meta_graph=model.ckpt-10000.meta \
    --input_checkpoint=model.ckpt-10000 \
    --output_graph=frozen_model.pb \
    --output_node_names=detection_boxes,detection_scores,detection_classes,num_detections \
    --input_binary=true

注意这里--input_binary要设为true,因为.meta文件是二进制格式的。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.26 09:26:50