TensorFlow预训练模型冻结后仅含输出占位符,无法导出完整图与权重
解决TensorFlow模型冻结后.pb文件仅含输出占位符的问题
我之前也碰到过一模一样的情况,咱们一步步拆解排查:
1. 先确认输出节点名称是否准确
你设置的output_node_names = "pos"大概率是问题根源。很多时候模型实际的输出节点名并非直观的"pos"——可能是pos/BiasAdd、pos/Sigmoid这类带运算后缀的名称,或是定义模型时用tf.identity包装后的节点名。
验证方法:
在恢复会话后,打印所有节点名称,找到真正的输出节点:
for node in graph.as_graph_def().node: print(node.name)
把找到的正确节点名替换output_node_names,多输出的话用逗号分隔即可。
2. 检查模型保存与恢复的完整性
如果原模型用tf.train.Saver()保存时,没有包含所有可训练变量(比如用tf.get_variable创建的变量没加入Saver的var_list),就会导致恢复后没有权重数据。
验证方法:
恢复会话后,打印所有可训练变量的内容:
for var in tf.trainable_variables(): print(var.name, var.eval(sess))
如果变量为空或者值异常,说明原模型保存环节有问题,得回去检查保存代码,确保所有权重都被Saver追踪。
3. 避免默认图的干扰
有时候导入.meta文件时,会和当前默认图的残留节点冲突,试试显式创建新图来隔离:
with tf.Graph().as_default() as graph: sess = tf.Session(graph=graph) saver = tf.train.import_meta_graph(path+model_name+'.meta', clear_devices=True) saver.restore(sess, path+model_name) # 后续的graph_util转换代码都放在这个代码块里
这样能保证我们操作的是导入的目标模型图,而非其他无关节点。
4. 清理训练相关节点再冻结
训练时的一些节点(比如Dropout、BatchNorm的更新操作)会干扰图的冻结,试试先清理这些节点:
from tensorflow.python.framework import graph_util # 先获取清理后的图定义 input_graph_def = graph.as_graph_def() input_graph_def = tf.graph_util.remove_training_nodes(input_graph_def) # 再转换为常量图 output_graph_def = graph_util.convert_variables_to_constants( sess, input_graph_def, output_node_names.split(",") )
5. 验证冻结后的模型
生成.pb文件后,别只用xxd看(二进制文件的十六进制视图很难直观判断),用TensorFlow工具验证:
# 查看模型的所有节点和输入输出 saved_model_cli show --dir ./model-gaze/ --all
或者用Python加载.pb文件检查:
with tf.gfile.GFile(output_graph, 'rb') as f: graph_def = tf.GraphDef() graph_def.ParseFromString(f.read()) tf.import_graph_def(graph_def, name='') for node in graph_def.node: print(node.name)
如果这里能看到完整的计算图节点,说明冻结已经成功了。
内容的提问来源于stack exchange,提问作者kalu
相关产品推荐
相关产品推荐

