无.pb后缀模型转.tflite:如何获取输入输出张量名称?
解决无.pb后缀模型的张量名获取与TFLite转换问题
我之前也踩过类似模型格式不明确的坑,给你几个实用的方向一步步排查解决:
1. 先搞清楚你的模型到底是什么格式
别着急重命名文件,先确认模型的类型:
- 如果它是一个文件夹,里面包含
variables子文件夹和saved_model.pb文件,那这是标准的SavedModel格式,完全不需要重命名,直接用TensorFlow API加载即可。 - 如果它是单个无后缀文件,大概率是冻结的TensorFlow GraphDef二进制文件(就是常说的.pb文件,只是没加后缀),或者是Checkpoint的单个组件(但Checkpoint一般是多个文件配合使用)。
2. 用Python代码直接加载模型并查看张量/节点
比起TensorBoard,用代码加载模型后直接查看节点和张量信息更可靠,下面分两种情况给出代码:
情况A:模型是SavedModel格式(文件夹)
import tensorflow as tf # 替换成你的模型文件夹路径 model_path = "path/to/your/model/folder" # 加载SavedModel loaded_model = tf.saved_model.load(model_path) # 获取模型的可用签名(SavedModel的签名定义里自带输入输出信息) signature_keys = list(loaded_model.signatures.keys()) print(f"模型可用签名: {signature_keys}") # 取默认签名(一般是"serving_default") infer_func = loaded_model.signatures["serving_default"] # 打印输入张量详情 print("\n输入张量信息:") for input_name, tensor_info in infer_func.structured_input_signature[1].items(): print(f" 名称: {input_name}, 形状: {tensor_info.shape}, 数据类型: {tensor_info.dtype}") # 打印输出张量详情 print("\n输出张量信息:") for output_name, tensor_info in infer_func.structured_outputs.items(): print(f" 名称: {output_name}, 形状: {tensor_info.shape}, 数据类型: {tensor_info.dtype}")
情况B:模型是单个GraphDef二进制文件(无后缀)
import tensorflow as tf # 替换成你的模型文件路径(不管有没有后缀都能读) model_path = "path/to/your/model/file" # 读取GraphDef二进制内容 with tf.io.gfile.GFile(model_path, 'rb') as f: graph_def = tf.compat.v1.GraphDef() graph_def.ParseFromString(f.read()) # 导入GraphDef到默认图 tf.import_graph_def(graph_def, name='') # 打印所有节点名称 print("所有节点名称:") for node in graph_def.node: print(f" {node.name}") # 筛选可能的输入节点(通常是Placeholder类型) print("\n可能的输入节点(Placeholder类型):") for node in graph_def.node: if node.op == 'Placeholder': shape_info = node.attr['shape'].shape # 格式化输出形状 shape_str = [dim.size for dim in shape_info.dim] print(f" 名称: {node.name}, 形状: {shape_str}") # 筛选可能的输出节点(关键词比如output、predict、result等) print("\n可能的输出节点(含output关键词):") for node in graph_def.node: if 'output' in node.name.lower(): print(f" 名称: {node.name}")
3. 转换为Android可用的TFLite格式
拿到输入输出张量名后,就可以进行转换了,分两种情况操作:
针对SavedModel格式(最推荐,稳定性最高)
不需要手动指定张量名,直接用官方API转换:
import tensorflow as tf model_path = "path/to/your/model/folder" converter = tf.lite.TFLiteConverter.from_saved_model(model_path) # 如果模型有自定义算子,可能需要开启兼容模式 # converter.target_spec.supported_ops = [tf.lite.OpsSet.TFLITE_BUILTINS, tf.lite.OpsSet.SELECT_TF_OPS] tflite_model = converter.convert() # 保存为.tflite文件 with open("face_completion.tflite", "wb") as f: f.write(tflite_model)
针对GraphDef二进制文件
需要指定输入输出张量的名称(注意如果导入时用了name='prefix',要加上前缀,上面代码里是name='',直接用节点名即可):
import tensorflow as tf model_path = "path/to/your/model/file" converter = tf.lite.TFLiteConverter.from_frozen_graph( graph_def_file=model_path, input_arrays=["input_tensor_name"], # 替换成你找到的输入张量名 output_arrays=["output_tensor_name"], # 替换成你找到的输出张量名 # 建议指定输入形状,避免转换时出错 input_shapes={"input_tensor_name": [1, 256, 256, 3]} # 根据你的模型实际形状调整 ) # 同样,有自定义算子时开启兼容 # converter.target_spec.supported_ops = [tf.lite.OpsSet.TFLITE_BUILTINS, tf.lite.OpsSet.SELECT_TF_OPS] tflite_model = converter.convert() with open("general_completion.tflite", "wb") as f: f.write(tflite_model)
4. 为什么之前TensorBoard加载失败?
- 如果是SavedModel文件夹,正确的加载命令是
tensorboard --logdir=path/to/your/model/folder,而不是加载重命名后的单个文件; - 如果是单个GraphDef文件,TensorBoard需要先将其导入到TensorFlow图中,再保存为日志文件才能加载,这个流程反而不如直接用代码查看节点高效。
另外你之前运行Python代码没输出节点,大概率是代码没有兼容TensorFlow 2.x的API,上面的代码已经用tf.compat.v1处理了旧版GraphDef的兼容问题,应该能正常输出。
内容的提问来源于stack exchange,提问作者Death14Stroke
相关产品推荐
相关产品推荐

