如何获取输入输出名称以将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
相关产品推荐
相关产品推荐

