如何将ONNX模型导入TensorFlow2.x?解决自定义改进版LeNet模型转换后无法完整加载的问题
解决方案:将ONNX模型导入TensorFlow 2.x并生成完整SavedModel结构
针对你遇到的问题——从ONNX模型转换后得到的单独.pb文件无法被tf.keras.models.load_model正常加载(因为缺少variables和assets文件夹),这里有两种可行的解决思路:
方法一:直接用onnx-tf导出完整的TensorFlow SavedModel格式
你之前使用onnx_tf.backend.prepare只是得到了模型的TensorFlow表示,但没有导出完整的SavedModel结构。可以通过以下步骤直接生成符合TF2.x要求的SavedModel(自动包含variables和assets文件夹):
import onnx from onnx_tf.backend import prepare # 加载ONNX模型 onnx_model = onnx.load("1645088924.84102.onnx") tf_rep = prepare(onnx_model) # 导出为完整的SavedModel格式 tf_rep.export_graph("./converted_tf_savedmodel/")
执行这段代码后,./converted_tf_savedmodel/目录下就会生成TF2.x加载所需的全部文件结构。之后你可以直接用常规方式加载模型:
import tensorflow as tf loaded_model = tf.keras.models.load_model("./converted_tf_savedmodel/") # 验证模型结构是否正确 loaded_model.summary()
方法二:加载单独的.pb文件并转换为Keras模型(备选方案)
如果你已经有了单独的.pb文件,也可以通过TensorFlow低阶API加载它,再封装为Keras模型使用。步骤如下:
import tensorflow as tf def load_frozen_pb_as_keras_model(pb_path): # 加载冻结模型的GraphDef with tf.io.gfile.GFile(pb_path, "rb") as f: graph_def = tf.compat.v1.GraphDef() graph_def.ParseFromString(f.read()) # 导入图并获取输入输出张量 with tf.Graph().as_default() as graph: tf.import_graph_def(graph_def, name="") # 注意:需要替换为你模型实际的输入输出节点名称 # 可以通过tf.compat.v1.get_default_graph().get_operations()查看所有节点名 input_tensor = graph.get_tensor_by_name("input_layer:0") output_tensor = graph.get_tensor_by_name("activation/Softmax:0") # 封装为Keras模型 keras_model = tf.keras.Model(inputs=input_tensor, outputs=output_tensor) return keras_model # 加载模型 frozen_model = load_frozen_pb_as_keras_model("your_model.pb") frozen_model.summary()
这种方法需要手动匹配模型的输入输出节点名称,生成的模型可能缺少部分Keras原生特性(如训练相关配置),因此更推荐使用方法一。
额外注意事项
- 确保
onnx、onnx-tf与TensorFlow版本兼容,建议通过pip install --upgrade onnx-tf升级到最新稳定版 - 可以用
onnx_model.opset_import查看ONNX模型的opset版本,确保其与你的TensorFlow版本适配
内容的提问来源于stack exchange,提问作者subspring
相关产品推荐
相关产品推荐

