如何用Python代码将TensorFlow PB模型转换为TFLite格式?
当然没问题!很多人都偏好用代码来控制模型转换的整个流程,毕竟能更灵活地添加自定义配置。下面是一份完整的Python代码示例,帮你把PB格式的模型转换成TFLite,适配TensorFlow Mobile框架:
将PB模型转换为TFLite的Python代码示例
前置准备
首先确保你安装了合适版本的TensorFlow(推荐2.x版本,API更稳定),如果还没装,可以用命令安装:
pip install tensorflow>=2.0
基础转换代码
这是最核心的转换逻辑,你只需要替换对应的路径和张量名称即可:
import tensorflow as tf # 1. 定义路径 pb_model_path = "/your/path/to/frozen_model.pb" # 你的PB模型路径 tflite_save_path = "/your/path/to/save/model.tflite" # TFLite模型保存路径 # 2. 初始化转换器 converter = tf.lite.TFLiteConverter.from_frozen_graph( graph_def_file=pb_model_path, input_arrays=["INPUT_TENSOR_NAME"], # 替换为你的模型输入张量名称 output_arrays=["OUTPUT_TENSOR_NAME"], # 替换为你的模型输出张量名称 # 可选:如果模型输入形状固定,可以指定,比如 input_shapes={"INPUT_TENSOR_NAME": [1, 224, 224, 3]} ) # 3. 执行转换 tflite_model = converter.convert() # 4. 保存转换后的模型 with open(tflite_save_path, "wb") as f: f.write(tflite_model) print(f"TFLite模型已成功保存到:{tflite_save_path}")
如何获取输入输出张量名称?
如果你不知道模型的输入输出张量名,可以用下面的小脚本查看PB模型里的所有节点名称:
import tensorflow as tf with tf.io.gfile.GFile(pb_model_path, "rb") as f: graph_def = tf.compat.v1.GraphDef() graph_def.ParseFromString(f.read()) # 打印所有节点名称 print("PB模型中的所有节点名称:") for node in graph_def.node: print(node.name)
通常输入节点会带有input、x之类的关键词,输出节点会带有output、predictions之类的关键词。
可选:添加模型优化(量化)
为了让模型更适合移动端部署,你可以开启TensorFlow Lite的优化选项,比如整数量化,能大幅减小模型体积并提升运行速度:
# 在初始化转换器后添加以下配置 converter.optimizations = [tf.lite.Optimize.DEFAULT] # 可选:提供代表性数据集做量化校准(提升量化精度) def representative_data_generator(): # 这里需要返回符合模型输入形状的真实数据,下面是示例用随机数据代替 for _ in range(100): # 替换成你的真实输入,比如预处理后的图片 yield [tf.random.normal([1, 224, 224, 3])] converter.representative_dataset = representative_data_generator # 配置全整数量化 converter.target_spec.supported_ops = [tf.lite.OpsSet.TFLITE_BUILTINS_INT8] converter.inference_input_type = tf.int8 converter.inference_output_type = tf.int8
把这段代码加在基础转换代码的第2步和第3步之间即可。
内容的提问来源于stack exchange,提问作者Nael Marwan
相关产品推荐
相关产品推荐

