ProtoBuf模型输入张量查找及推理占位符缺失报错咨询
嘿,这两个问题我刚好熟,给你一步步拆解解决!
解决InvalidArgumentError错误
错误原因分析
这个报错直白说就是:你的推理代码里没有给名为'w2'的float类型占位符传入数据。冻结后的.pb模型里,所有保留下来的Placeholder节点(不管是主输入还是训练时的辅助参数),在推理阶段都必须被喂入对应数据类型和形状的值,少一个都会触发这个报错。
快速修复步骤
- 先确认'w2'的身份:它可能是训练时的某个辅助输入(比如权重参数、dropout保留率、注意力掩码这类),冻结模型时被保留了下来。
- 在你的推理代码里,把这个张量加入
feed_dict,确保传入对应的数据:
import tensorflow as tf import numpy as np # 加载冻结模型 with tf.gfile.GFile('你的模型路径.pb', 'rb') as f: graph_def = tf.GraphDef() graph_def.ParseFromString(f.read()) with tf.Session() as sess: sess.graph.as_default() tf.import_graph_def(graph_def, name='') # 找到报错的'w2'张量,以及你的主输入、输出张量 main_input = sess.graph.get_tensor_by_name('你的主输入节点名:0') w2_input = sess.graph.get_tensor_by_name('w2:0') output_tensor = sess.graph.get_tensor_by_name('你的输出节点名:0') # 准备数据:注意dtype必须是float,形状要和w2的shape匹配 # 如果不知道shape,可以先打印w2_input.shape.as_list()查看 main_data = 你的主输入数据 # 比如预处理后的图像数组 w2_data = np.random.rand(*w2_input.shape.as_list()).astype(np.float32) # 如果w2是训练时的固定参数,也可以喂对应常量(比如dropout的keep_prob喂1.0) # 执行推理 result = sess.run(output_tensor, feed_dict={ main_input: main_data, w2_input: w2_data })
如何获取.pb模型的所有输入张量
这里有两种实用方法,按需选择:
方法一:用代码遍历筛选输入节点
直接写代码遍历模型图,找出所有Placeholder类型的节点(也就是输入),还能获取它们的名称、数据类型和形状:
import tensorflow as tf with tf.gfile.GFile('你的模型路径.pb', 'rb') as f: graph_def = tf.GraphDef() graph_def.ParseFromString(f.read()) with tf.Session() as sess: sess.graph.as_default() tf.import_graph_def(graph_def, name='') # 筛选所有Placeholder节点 input_info_list = [] for node in sess.graph.as_graph_def().node: if node.op == 'Placeholder': tensor = sess.graph.get_tensor_by_name(f"{node.name}:0") input_info_list.append({ '节点名称': node.name, '数据类型': tensor.dtype.name, '形状': tensor.shape.as_list() }) # 打印所有输入信息 for idx, info in enumerate(input_info_list): print(f"=== 输入{idx+1} ===") print(f"名称: {info['节点名称']}") print(f"数据类型: {info['数据类型']}") print(f"形状: {info['形状']}\n")
方法二:用TensorBoard可视化查看
如果想看模型的整体结构,用TensorBoard更直观:
- 先把
.pb模型导出为TensorBoard能识别的日志文件:
import tensorflow as tf with tf.gfile.GFile('你的模型路径.pb', 'rb') as f: graph_def = tf.GraphDef() graph_def.ParseFromString(f.read()) with tf.Session() as sess: sess.graph.as_default() tf.import_graph_def(graph_def, name='') # 保存日志到指定目录 writer = tf.summary.FileWriter('./model_logs', sess.graph) writer.close()
- 打开终端,运行命令:
tensorboard --logdir=./model_logs - 按照终端提示的地址(一般是http://localhost:6006)打开浏览器,在Graphs页面就能看到模型的所有节点,找
Placeholder类型的就是输入节点,还能查看它们的连接关系。
额外注意点
- 有些模型的输入可能不止一个(比如BERT的输入就有token_ids、attention_mask、token_type_ids三个),一定要把所有
Placeholder都喂数据 - 如果某个
Placeholder是训练专用的(比如dropout的keep_prob),推理时喂固定值就行(比如1.0) - 输入数据的dtype必须和Placeholder完全匹配,比如报错里的float,就不能喂int或者double类型的数据
内容的提问来源于stack exchange,提问作者Ignacio Peletier
相关产品推荐
相关产品推荐

