TensorFlow模型转TFLite:input_arrays与output_arrays取值疑问
如何确定TFLite转换时的input_arrays和output_arrays取值
我来帮你搞定这个问题——你要找的input_arrays和output_arrays其实就是模型的输入张量名称和输出张量名称,针对TensorFlow Object Detection模型,这几个实用方法可以快速找到它们:
- 方法1:用
saved_model_cli工具(最省心)
如果你的模型是训练后导出的SavedModel格式,直接在终端运行这条命令:
saved_model_cli show --dir /path/to/your/saved_model --all
执行后会输出详细的模型签名信息,里面会明确标注Inputs和Outputs对应的张量名称。比如目标检测模型的输入通常是image_tensor,输出一般包含detection_boxes、detection_scores、detection_classes、num_detections这几个(毕竟要输出检测框、得分、类别这些核心信息)。
要是你手里只有冻结的.pb文件,也可以先把它转成SavedModel格式再用这个工具:
import tensorflow as tf # 加载冻结图 graph_def = tf.compat.v1.GraphDef() with tf.io.gfile.GFile('/tmp/frozen_graph.pb', 'rb') as f: graph_def.ParseFromString(f.read()) # 导出为SavedModel with tf.compat.v1.Session() as sess: tf.import_graph_def(graph_def, name='') builder = tf.compat.v1.saved_model.Builder('/tmp/converted_saved_model') builder.add_meta_graph_and_variables(sess, [tf.saved_model.SERVING]) builder.save()
- 方法2:直接打印冻结图的所有节点
不想转格式的话,用几行代码加载冻结图,把所有节点名称打印出来,再筛选输入和输出节点:
import tensorflow as tf graph_def = tf.compat.v1.GraphDef() with tf.io.gfile.GFile('/tmp/frozen_graph.pb', 'rb') as f: graph_def.ParseFromString(f.read()) # 遍历所有节点并打印名称 for node in graph_def.node: print(node.name)
输入节点一般是Placeholder类型的,名称里常带有input、image这类关键词;输出节点就是模型最终输出的张量,目标检测模型的输出节点就是前面提到的那几个检测相关的名称。
- 方法3:用TensorBoard可视化图结构
想直观看到模型的结构?用TensorBoard就能清晰定位输入输出节点:
- 先把
.pb文件转换成TensorBoard能识别的日志文件:
import tensorflow as tf from tensorflow.python.summary import summary graph_def = tf.compat.v1.GraphDef() with tf.io.gfile.GFile('/tmp/frozen_graph.pb', 'rb') as f: graph_def.ParseFromString(f.read()) with tf.compat.v1.Session() as sess: tf.import_graph_def(graph_def, name='') summary_writer = summary.FileWriter('/tmp/tb_logs', sess.graph) summary_writer.close()
- 启动TensorBoard:
tensorboard --logdir=/tmp/tb_logs
- 打开浏览器里显示的地址,进入「Graphs」页面,就能看到完整的模型图,点击节点就能查看它的名称,输入输出节点一眼就能找到。
另外提醒下:你示例里的代码是针对MobileNet分类模型的,和你的目标检测模型结构不一样,别直接套用示例里的input_arrays和output_arrays参数哦!
内容的提问来源于stack exchange,提问作者Dyboo
相关产品推荐
相关产品推荐

