如何将.pb格式的SavedModel模型转换为Keras的.h5格式模型?
问题解答
.pb格式SavedModel转.h5格式的实现方案
tf.saved_model.load加载得到的是SavedModel实例,不是原生Keras模型对象,无法直接保存为.h5格式,需根据模型来源选择对应方案:
- 如果该.pb模型是Keras训练后通过
model.save()导出的SavedModel格式:
直接添加compile=False参数加载,再保存为.h5即可,示例代码:# 加载SavedModel目录 model = tf.keras.models.load_model(input_model_path, compile=False) # 保存为h5格式 model.save("output_model.h5") - 如果该.pb模型是非Keras训练导出的SavedModel(比如TensorFlow 1.x导出、第三方工具导出的模型):
先获取推理签名,再手动封装为Keras模型后保存,示例代码:
注意:封装时需要将示例中的输入shape、输出节点名称替换为你模型的实际参数,可通过# 加载SavedModel loaded_model = tf.saved_model.load(input_model_path) # 获取默认推理签名 infer = loaded_model.signatures["serving_default"] # 按照实际输入输出调整shape和名称,此处以输入为(None, 224, 224, 3)的图像、输出单张量为例 input_layer = tf.keras.Input(shape=(224,224,3), name="input") output_layer = infer(input_layer)["output_name"] # 替换为实际输出节点的名称 # 封装为Keras模型 keras_model = tf.keras.Model(inputs=input_layer, outputs=output_layer) # 保存为h5格式 keras_model.save("output_model.h5")print(infer.structured_outputs)查看输出节点信息
tf.keras.models.load_model加载.pb报IndexError的解决方法
该报错的核心原因是tf.keras.models.load_model默认加载Keras原生格式的模型,当目标文件缺少Keras元数据、路径配置错误时就会触发索引异常,可按以下步骤排查:
- 检查传入路径是否正确:
tf.keras.models.load_model加载SavedModel时需要传入整个SavedModel目录的路径,不要直接指向单个.pb文件,目录下需同时存在saved_model.pb文件和variables目录 - 加载时添加
compile=False参数:跳过模型编译阶段的元数据校验,可解决大部分Keras导出的SavedModel加载报错问题,用法为model = tf.keras.models.load_model(input_model_path, compile=False) - 确认模型来源:如果是TensorFlow 1.x或者非Keras框架导出的SavedModel,不要直接用
tf.keras.models.load_model加载,改用上述封装为Keras模型的方案即可
内容的提问来源于stack exchange,提问作者lkuebler
相关产品推荐
相关产品推荐

