使用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
相关产品推荐
相关产品推荐

