TensorFlow中sess.run执行全局变量初始化报错:fetch为None的解决方法
修复TensorFlow中
sess.run(tf.compat.v1.global_variables_initializer())的TypeError错误 问题场景
执行代码中的sess.run(tf.compat.v1.global_variables_initializer())语句时,触发错误:
TypeError: Argument
fetch= None has invalid type "NoneType". Cannot be None
相关代码
sess = tf.compat.v1.Session(config=config) with open(r'C:\Users\User\Downloads\New folder (8)\Docify-master\api\data\ctpn.pb', 'rb') as f: graph_def = tf.compat.v1.GraphDef() # print("sdsds",graph_def) # text_format.Merge(f.read(), graph_def) graph_def.ParseFromString(f.read()) # print("sss",graph_def.ParseFromString(f.read())) sess.graph.as_default() tf.import_graph_def(graph_def, name='') # print("hhdsd",tf.compat.v1.global_variables_initializer()) sess.run(tf.compat.v1.global_variables_initializer()) input_img = sess.graph.get_tensor_by_name('Placeholder:0') output_cls_prob = sess.graph.get_tensor_by_name('Reshape_2:0') output_box_pred = sess.graph.get_tensor_by_name('rpn_bbox_pred/Reshape_1:0') textdetector = TextDetector()
修复方法
核心原因
从预训练的pb文件导入计算图时,图中的变量已经被固化为常量,此时tf.compat.v1.global_variables_initializer()会返回None,传给sess.run()就会触发类型错误。
具体解决方案
直接删除初始化语句
删掉sess.run(tf.compat.v1.global_variables_initializer())这一行即可,因为预训练模型不需要重新初始化变量。修改后的代码:sess = tf.compat.v1.Session(config=config) with open(r'C:\Users\User\Downloads\New folder (8)\Docify-master\api\data\ctpn.pb', 'rb') as f: graph_def = tf.compat.v1.GraphDef() graph_def.ParseFromString(f.read()) sess.graph.as_default() tf.import_graph_def(graph_def, name='') # 移除初始化语句 input_img = sess.graph.get_tensor_by_name('Placeholder:0') output_cls_prob = sess.graph.get_tensor_by_name('Reshape_2:0') output_box_pred = sess.graph.get_tensor_by_name('rpn_bbox_pred/Reshape_1:0') textdetector = TextDetector()安全初始化(兼容有可初始化变量的场景)
如果不确定图中是否存在需要初始化的变量,可以先检查全局变量列表,再选择性初始化:sess = tf.compat.v1.Session(config=config) with open(r'C:\Users\User\Downloads\New folder (8)\Docify-master\api\data\ctpn.pb', 'rb') as f: graph_def = tf.compat.v1.GraphDef() graph_def.ParseFromString(f.read()) sess.graph.as_default() tf.import_graph_def(graph_def, name='') # 检查并初始化存在的全局变量 global_vars = tf.compat.v1.global_variables() if global_vars: sess.run(tf.compat.v1.variables_initializer(global_vars)) input_img = sess.graph.get_tensor_by_name('Placeholder:0') output_cls_prob = sess.graph.get_tensor_by_name('Reshape_2:0') output_box_pred = sess.graph.get_tensor_by_name('rpn_bbox_pred/Reshape_1:0') textdetector = TextDetector()
内容的提问来源于stack exchange,提问作者vishnu Vs
相关产品推荐
相关产品推荐

