TensorFlow1转TensorFlow2时detection_graph.as_default()报AttributeError如何解决
错误原因
你混淆了tf.Graph和tf.GraphDef两个类的功能定位:
- TensorFlow 1.x中的
tf.Graph()是计算图实例类,才具备as_default()上下文管理器方法 - 你修改后的代码错误地将第一行的实例化对象换成了
tf.compat.v1.GraphDef(),该类是序列化冻结图二进制内容的存储容器,本身不存在as_default属性,因此触发AttributeError。
修复代码(兼容模式)
如果沿用原TensorFlow 1.x的图逻辑写法,只需修正第一行的实例化类即可:
detection_graph = tf.compat.v1.Graph() with detection_graph.as_default(): od_graph_def = tf.compat.v1.GraphDef() with tf.io.gfile.GFile(PATH_TO_FROZEN_GRAPH, 'rb') as fid: serialized_graph = fid.read() od_graph_def.ParseFromString(serialized_graph) tf.compat.v1.import_graph_def(od_graph_def, name='')
如果后续需要通过会话执行模型推理,建议在代码开头添加禁用TensorFlow 2默认eager执行的语句:
import tensorflow as tf tf.compat.v1.disable_eager_execution()
可选:TensorFlow 2原生适配方案
如果不想使用兼容模式的图上下文,可以直接将冻结图转换为TensorFlow 2原生可调用函数:
def load_frozen_graph(model_path): # 加载冻结图二进制内容 with tf.io.gfile.GFile(model_path, 'rb') as f: graph_def = tf.compat.v1.GraphDef() graph_def.ParseFromString(f.read()) # 包装为TF2可调用函数 wrapped_model = tf.compat.v1.wrap_function( lambda: tf.compat.v1.import_graph_def(graph_def, name=''), [] ) # 替换为你的模型实际输入、输出张量名称 input_tensor = wrapped_model.graph.get_tensor_by_name('input:0') output_tensor = wrapped_model.graph.get_tensor_by_name('detection_out:0') return wrapped_model.prune(feeds=input_tensor, fetches=output_tensor) # 调用示例 model = load_frozen_graph(PATH_TO_FROZEN_GRAPH) prediction = model(your_input_data)
内容的提问来源于stack exchange,提问作者TourEiffel
相关产品推荐
相关产品推荐

