TensorFlow在自定义类中初始化变量失败的原因是什么?
为什么自定义类中TensorFlow变量初始化会失败?
先把你提供的简化代码补全(推测build_model里的完整逻辑):
import tensorflow as tf class ExModel(object): def __init__(self, graph=None): if graph is None: self.graph = tf.get_default_graph() else: self.graph = graph self.sess = tf.Session(graph=self.graph) # 先创建初始化操作,再构建模型 self.init_op = tf.group(tf.global_variables_initializer(), tf.local_variables_initializer()) self.build_model() print(self.sess.run(tf.report_uninitialized_variables())) # 输出[b'global_step'] print(self.sess.run(self._global_step)) # 抛出未初始化错误 def build_model(self): with self.graph.as_default(): self._global_step = tf.Variable(0, trainable=False, name='global_step')
问题根源
你踩了TensorFlow变量初始化的一个典型坑:初始化操作的创建时机早于变量定义。
TensorFlow的tf.global_variables_initializer()会收集调用它时已经存在的所有全局变量,生成对应的初始化操作。你在__init__里先创建了self.init_op,之后才调用build_model()定义_global_step变量——这就导致init_op里根本没有包含_global_step的初始化逻辑,自然运行时这个变量就处于未初始化状态。
解决方法
调整代码顺序,先构建所有变量,再创建初始化操作,修改后的__init__方法如下:
def __init__(self, graph=None): if graph is None: self.graph = tf.get_default_graph() else: self.graph = graph self.sess = tf.Session(graph=self.graph) # 先构建模型,创建所有变量 self.build_model() # 再创建包含所有变量的初始化操作 self.init_op = tf.group(tf.global_variables_initializer(), tf.local_variables_initializer()) # 执行初始化 self.sess.run(self.init_op) print(self.sess.run(tf.report_uninitialized_variables())) # 输出空数组 print(self.sess.run(self._global_step)) # 正常输出0
这样修改后,tf.global_variables_initializer()就能收集到build_model()里定义的_global_step,执行self.sess.run(self.init_op)后所有变量都会被正确初始化。
内容的提问来源于stack exchange,提问作者imhuay
相关产品推荐
相关产品推荐

