TensorFlow冻结.pb模型转ONNX失败,寻求可行转换方案
解决TensorFlow冻结.pb模型转ONNX的问题
首先,你之前用MMdnn的mmconvert报错,根源就是没指定TensorFlow冻结模型的输出节点——先给你把这个坑填上,再给你推荐几个更常用、靠谱的转换方法:
一、先搞定MMdnn的转换错误
要修复这个报错,你需要先确定模型的输出节点名称,然后在命令里补上--outNodeName参数:
快速找输出节点的方法:
- 用Netron可视化工具(强烈推荐):先安装
pip install netron,然后运行netron /frozen_models/model.pb,在浏览器打开对应的地址,就能直观看到模型的输入输出节点名称(比如常见的Softmax:0、predictions:0这类)。 - 或者用TensorFlow代码打印所有节点:
import tensorflow as tf with tf.compat.v1.Session() as sess: with tf.io.gfile.GFile('/frozen_models/model.pb', 'rb') as f: graph_def = tf.compat.v1.GraphDef() graph_def.ParseFromString(f.read()) tf.import_graph_def(graph_def, name='') # 遍历打印所有节点名称,从中找输出节点 for node in graph_def.node: print(node.name)
- 用Netron可视化工具(强烈推荐):先安装
修正后的mmconvert命令:
假设你的输出节点叫output,命令应该改成:mmconvert -sf tensorflow -iw /frozen_models/model.pb --inNodeName input --inputShape 512 --outNodeName output -df onnx -om tf_mobilenet.onnx
二、推荐其他更稳定的转换方法
1. 使用ONNX官方的tf2onnx工具(最常用)
tf2onnx是ONNX官方维护的转换工具,对TensorFlow模型的兼容性更好,步骤如下:
- 安装:
pip install tf2onnx - 转换命令(需要指定输入输出节点的完整名称,比如
input:0):python -m tf2onnx.convert --input /frozen_models/model.pb --inputs input:0 --outputs output_node_name:0 --output tf_mobilenet.onnx
2. 先转成SavedModel格式再转换
冻结的.pb模型是GraphDef格式,部分工具对SavedModel格式支持更友好,所以可以先转成SavedModel,再用tf2onnx转换:
- 用Python代码把冻结模型转成SavedModel:
import tensorflow as tf from tensorflow.python.saved_model import signature_constants, tag_constants with tf.compat.v1.Session() as sess: # 加载冻结模型 with tf.io.gfile.GFile('/frozen_models/model.pb', 'rb') as f: graph_def = tf.compat.v1.GraphDef() graph_def.ParseFromString(f.read()) sess.graph.as_default() tf.import_graph_def(graph_def, name='') # 替换成你的输入输出节点完整名称 input_tensor = sess.graph.get_tensor_by_name('input:0') output_tensor = sess.graph.get_tensor_by_name('output_node_name:0') # 保存为SavedModel builder = tf.compat.v1.saved_model.builder.SavedModelBuilder('./saved_model') signature = tf.compat.v1.saved_model.signature_def_utils.predict_signature_def( inputs={'input': input_tensor}, outputs={'output': output_tensor}) builder.add_meta_graph_and_variables( sess, [tag_constants.SERVING], signature_def_map={ signature_constants.DEFAULT_SERVING_SIGNATURE_DEF_KEY: signature }) builder.save() - 然后转换SavedModel到ONNX:
python -m tf2onnx.convert --saved-model ./saved_model --output tf_mobilenet.onnx
3. 曲线救国:转TFLite再转ONNX
如果上面的方法都遇到兼容性问题,还可以走这条间接路线:
- 先把冻结模型转成TensorFlow Lite格式:
import tensorflow as tf converter = tf.compat.v1.lite.TFLiteConverter.from_frozen_graph( graph_def_file='/frozen_models/model.pb', input_arrays=['input'], input_shapes={'input': [1, 512]}, # 必须加上batch维度 output_arrays=['output_node_name'] ) tflite_model = converter.convert() with open('model.tflite', 'wb') as f: f.write(tflite_model) - 再用
tflite2onnx工具转ONNX:
先安装pip install tflite2onnx,然后运行:tflite2onnx model.tflite tf_mobilenet.onnx
内容的提问来源于stack exchange,提问作者albusdemens
相关产品推荐
相关产品推荐

