You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

使用Keras训练模型转.pb后,转换为.tflite时出错求助

解决Keras .h5转.pb再转.tflite的错误问题

我帮你梳理下这个问题的解决方案,你遇到的.pb转.tflite失败的问题,大多是因为中间.pb转换不完整、节点名称不对或者TensorFlow版本兼容问题,下面分步骤给你解决思路:

1. 先确保.pb文件转换正确(适配TF2.x环境)

你原来的脚本是偏向TF1.x的写法,在TF2.x环境下容易出现兼容问题,而且代码看起来没写完(nb_classes = 1后面没有后续逻辑)。这里给你两种正确的.pb转换方式:

方式一:用SavedModel格式导出(推荐,TF2.x原生支持)

这种方式更稳定,后续转.tflite也更方便:

import tensorflow as tf
from keras.models import load_model

# 加载你的.h5模型
model = load_model("/media/hsmnzaydn/8AD030E8D030DBDF/Projects/Machine Learning/Basic Keras/CancerDetected/modelim.h5")

# 导出为SavedModel格式(TF2.x推荐的模型格式)
saved_model_dir = "./saved_model"
model.save(saved_model_dir, save_format='tf')

# 可选:如果一定要生成.pb文件,也可以基于SavedModel转换,不过更推荐直接用SavedModel转.tflite

方式二:补全TF1.x风格的冻结图转换(适合兼容旧代码)

如果你坚持用graph_util的方式,需要补全完整的冻结逻辑:

from tensorflow.python.framework import graph_util
from tensorflow.python.framework import graph_io
from keras.models import load_model
from keras import backend as K
import os

# 关闭学习模式,确保模型处于推理状态
K.set_learning_phase(0)
# 加载模型
model = load_model("/media/hsmnzaydn/8AD030E8D030DBDF/Projects/Machine Learning/Basic Keras/CancerDetected/modelim.h5")
nb_classes = 1  # 你的模型类别数

# 获取输入、输出节点的名称(后续转.tflite需要用到)
input_node_name = model.input.name.split(':')[0]
output_node_name = model.output.name.split(':')[0]

# 冻结图,将变量转为常量
sess = K.get_session()
graph_def = sess.graph.as_graph_def()
frozen_graph_def = graph_util.convert_variables_to_constants(
    sess, graph_def, [output_node_name]
)

# 保存.pb文件
output_dir = "./pb_model"
os.makedirs(output_dir, exist_ok=True)
graph_io.write_graph(frozen_graph_def, output_dir, "model.pb", as_text=False)
print(f"PB文件已保存:{output_dir}/model.pb")
print(f"输入节点名称:{input_node_name},输出节点名称:{output_node_name}")

2. 从.pb文件转.tflite的正确步骤

转.tflite时最容易踩的坑是输入输出节点名称错误,一定要用上面打印的节点名来指定:

import tensorflow as tf

# 加载冻结的.pb文件
converter = tf.lite.TFLiteConverter.from_frozen_graph(
    graph_def_file="./pb_model/model.pb",
    input_arrays=[input_node_name],  # 替换成上面打印的输入节点名,比如"input_1"
    output_arrays=[output_node_name],  # 替换成上面打印的输出节点名,比如"dense_2/Softmax"
    # 如果你的模型输入有固定形状,需要指定,比如:input_shapes={"input_1": [1, 224, 224, 3]}
)

# 转换并保存.tflite文件
tflite_model = converter.convert()
with open("model.tflite", "wb") as f:
    f.write(tflite_model)

3. 更简单的捷径:直接从.h5转.tflite

其实完全不需要中间转.pb这一步,直接用TF Lite转换器从.h5模型转换,能避免很多中间环节的错误:

import tensorflow as tf

# 直接加载.h5模型并转换
converter = tf.lite.TFLiteConverter.from_keras_model_file(
    "/media/hsmnzaydn/8AD030E8D030DBDF/Projects/Machine Learning/Basic Keras/CancerDetected/modelim.h5"
)

# 可选:如果需要模型量化(减小体积、加速推理),可以添加以下配置
# converter.optimizations = [tf.lite.Optimize.DEFAULT]

# 生成并保存.tflite文件
tflite_model = converter.convert()
with open("model.tflite", "wb") as f:
    f.write(tflite_model)

常见错误排查点

  • 节点名称不匹配:如果报错提示找不到节点,用netron工具打开.pb文件查看准确的输入输出节点名,或者在导出.pb时打印出来核对。
  • TF版本不一致:确保训练模型时的Keras/TensorFlow版本,和转换时的版本一致,TF2.x和TF1.x的转换逻辑差异很大。
  • 自定义层问题:如果模型包含自定义层,需要在转换时添加支持:
    converter.target_spec.supported_ops = [tf.lite.OpsSet.TFLITE_BUILTINS, tf.lite.OpsSet.SELECT_TF_OPS]
    
  • 输入形状不明确:如果模型输入没有固定形状,转.tflite时需要通过input_shapes参数指定输入的维度。

内容的提问来源于stack exchange,提问作者Serkan Özaydin

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.05.21 04:33:28