如何将TensorFlow的.pb模型文件转换为Keras的.h5格式?
嘿,这个问题我刚好踩过几次坑,其实分两种常见情况来处理就好,取决于你的.pb模型是哪种类型:
情况1:你的.pb是SavedModel格式(附带variables文件夹)
这种是最省心的情况,TensorFlow 2.x的Keras原生支持直接加载SavedModel,然后转存为.h5格式就行:
from tensorflow import keras # 注意:这里的路径是包含saved_model.pb的文件夹,不是单个.pb文件哦 loaded_model = keras.models.load_model("/path/to/your/saved_model_folder") # 直接保存为.h5格式 loaded_model.save("converted_model.h5")
如果遇到版本兼容提示,可以显式指定保存格式:
loaded_model.save("converted_model.h5", save_format="h5")
情况2:你的.pb是冻结的单文件模型(仅单个.pb文件,无variables文件夹)
这种情况需要先解析冻结图,手动指定输入输出节点来重建Keras模型,步骤稍微复杂一点:
步骤1:加载冻结图并获取输入输出张量
首先你得知道模型的输入和输出节点名称——可以通过tf.compat.v1.get_default_graph().get_operations()打印所有节点,从中找到对应的输入输出(比如常见的输入节点名可能是input_1:0,输出可能是dense_1/Softmax:0)。
import tensorflow as tf from tensorflow import keras def load_frozen_graph(pb_file_path): # 加载冻结的计算图 graph = tf.Graph() with graph.as_default(): graph_def = tf.compat.v1.GraphDef() with tf.io.gfile.GFile(pb_file_path, 'rb') as f: graph_def.ParseFromString(f.read()) # 导入图到当前默认图 tf.import_graph_def(graph_def, name='') # 替换成你自己的输入输出节点名称! input_tensor = graph.get_tensor_by_name("your_input_node_name:0") output_tensor = graph.get_tensor_by_name("your_output_node_name:0") return input_tensor, output_tensor
步骤2:构建Keras模型并保存
拿到输入输出张量后,用Keras的Model类包装,然后直接保存:
# 加载冻结图的输入输出 input_tensor, output_tensor = load_frozen_graph("/path/to/your/frozen_model.pb") # 构建Keras模型 keras_model = keras.Model(inputs=input_tensor, outputs=output_tensor) # 保存为.h5格式 keras_model.save("converted_frozen_model.h5")
注意事项
- 如果你的模型包含自定义层或自定义操作,加载/保存前需要先注册这些自定义对象,比如:
from your_custom_layers import CustomLayer keras.utils.get_custom_objects().update({"CustomLayer": CustomLayer}) - 要是不确定节点名称,可以运行以下代码打印所有节点:
for op in graph.get_operations(): print(op.name)
内容的提问来源于stack exchange,提问作者Vampavi
相关产品推荐
相关产品推荐

