TensorFlow 2.1/Keras调用freeze_session报错:output_node/Identity不在图中
解决TensorFlow 2.1.0中冻结Keras .h5模型时
output_node/Identity is not in graph的错误 我之前碰到过一模一样的问题,这本质是TensorFlow 2.x与1.x的兼容性差异,再加上你的freeze_session函数里的一个小逻辑问题共同导致的。下面给你几个可行的解决方案:
原因分析
- 你的
freeze_session函数里有一行output_names += [v.op.name for v in tf.compat.v1.global_variables()],这会把所有全局变量的节点名强行加入输出列表,干扰了模型真实输出节点的识别。 - 在TF2.x的即时执行模式下,直接用
K.get_session()获取的会话,可能没有正确加载模型的完整图结构,导致output_node/Identity节点无法被检索到。
方案一:修正freeze_session函数(快速修复)
先修改你的freeze_session函数,删掉那行多余的输出节点添加逻辑:
def freeze_session(session, keep_var_names=None, output_names=None, clear_devices=True): """ Freezes the state of a session into a pruned computation graph. Creates a new computation graph where variable nodes are replaced by constants taking their current value in the session. The new graph will be pruned so subgraphs that are not necessary to compute the requested outputs are removed. @param session The TensorFlow session to be frozen. @param keep_var_names A list of variable names that should not be frozen, or None to freeze all the variables in the graph. @param output_names Names of the relevant graph outputs. @param clear_devices Remove the device directives from the graph for better portability. @return The frozen graph definition. """ graph = session.graph with graph.as_default(): freeze_var_names = list(set(v.op.name for v in tf.compat.v1.global_variables()).difference(keep_var_names or [])) output_names = output_names or [] # 删掉这行!不要把全局变量加入输出节点列表 # output_names += [v.op.name for v in tf.compat.v1.global_variables()] input_graph_def = graph.as_graph_def() if clear_devices: for node in input_graph_def.node: node.device = "" frozen_graph = tf.compat.v1.graph_util.convert_variables_to_constants( session, input_graph_def, output_names, freeze_var_names) return frozen_graph
然后调用时,尝试使用模型输出层的名字(不带/Identity后缀):
frozen_graph = freeze_session(K.get_session(), output_names=["output_node"]) write_graph(frozen_graph, './', 'graph.pbtxt', as_text=True) write_graph(frozen_graph, './', 'graph.pb', as_text=False)
方案二:使用TF2.x原生方式冻结模型(推荐)
TF2.x已经不再推荐使用会话(Session)这种TF1.x风格的写法,更适合用SavedModel格式来处理模型的保存与冻结。步骤如下:
- 加载你的.h5模型,保存为SavedModel格式:
import tensorflow as tf from tensorflow import keras as kr model = kr.models.load_model("model.h5") # 保存为SavedModel格式,这会生成包含完整图结构的文件夹 tf.saved_model.save(model, "./my_saved_model")
- 加载SavedModel并冻结图:
# 加载SavedModel saved_model = tf.saved_model.load("./my_saved_model") # 获取默认的推理签名 infer_signature = saved_model.signatures["serving_default"] # 提取输入输出节点名(去掉张量的索引后缀,比如":0") input_names = [tensor.name.split(':')[0] for tensor in infer_signature.inputs] output_names = [tensor.name.split(':')[0] for tensor in infer_signature.outputs] # 在兼容TF1.x的会话中冻结图 graph = tf.compat.v1.get_default_graph() with tf.compat.v1.Session(graph=graph) as sess: # 初始化所有变量 tf.compat.v1.global_variables_initializer().run() # 转换变量为常量,完成图冻结 frozen_graph = tf.compat.v1.graph_util.convert_variables_to_constants( sess, graph.as_graph_def(), output_names=output_names ) # 保存冻结后的图文件 tf.io.write_graph(frozen_graph, './', 'graph.pbtxt', as_text=True) tf.io.write_graph(frozen_graph, './', 'graph.pb', as_text=False)
这种方法更贴合TF2.x的设计理念,能避免大部分兼容性问题,也是长期来看最稳妥的方案。
方案三:手动重新定义模型输出节点(备选)
如果上面的方法都无效,你可以在加载模型后,手动重新定义输出节点,绕过Identity的问题:
model = kr.models.load_model("model.h5") # 直接取输出层的原始输出,而非模型默认的带Identity的输出 new_output = model.get_layer("output_node").output # 创建新的模型实例 new_model = kr.models.Model(inputs=model.input, outputs=new_output) # 此时打印输出节点名,应该是"output_node" print(new_output.op.name) # 用这个新模型执行冻结操作 frozen_graph = freeze_session(K.get_session(), output_names=[new_output.op.name])
内容的提问来源于stack exchange,提问作者marco890
相关产品推荐
相关产品推荐

