如何初始化TensorFlow函数内部定义的变量?
嘿,这个问题我之前踩过坑,核心是变量创建时机和初始化时机不匹配导致的:
你看,代码里先运行了tf.global_variables_initializer(),但这时候foo()里的变量x还没被定义,所以这个初始化操作里根本不包含x。等你之后调用foo()才创建x,再去运行计算的时候,自然会报错说变量没初始化。
下面给你几种可行的解决办法:
办法1:提前创建变量(最简单直接)
把函数内的变量创建逻辑提前,确保在执行初始化前,所有需要的变量都已经被定义好了:
import tensorflow as tf def foo(): x = tf.Variable(tf.ones([1])) y = tf.ones([1]) return x+y if __name__ == '__main__': # 先调用foo生成计算图和变量,再初始化 result_tensor = foo() with tf.Session() as sess: init = tf.global_variables_initializer() sess.run(init) print(sess.run(result_tensor))
办法2:用tf.get_variable配合作用域(更灵活)
如果需要在函数内动态创建变量,或者要复用变量的话,用tf.get_variable会更合适,同样要确保初始化前变量已被创建:
import tensorflow as tf def foo(): # 使用get_variable来定义变量,可通过参数控制是否复用 x = tf.get_variable('x', shape=[1], initializer=tf.ones_initializer()) y = tf.ones([1]) return x+y if __name__ == '__main__': with tf.Session() as sess: # 先调用foo创建变量 result_tensor = foo() init = tf.global_variables_initializer() sess.run(init) print(sess.run(result_tensor))
办法3:函数内单独初始化变量(适合特殊场景,不推荐常规用)
如果必须在函数调用时才创建变量,那可以在函数内单独初始化这个变量,但要注意每次调用函数都会重新初始化变量,可能不符合你的需求:
import tensorflow as tf def foo(sess): x = tf.Variable(tf.ones([1])) # 单独初始化当前函数内的变量 sess.run(x.initializer) y = tf.ones([1]) return x+y if __name__ == '__main__': with tf.Session() as sess: print(sess.run(foo(sess)))
额外提一句:如果能升级到TensorFlow 2.x会更省心
TF2.x默认是即时执行模式,不需要手动管理Session和初始化,代码会简洁很多:
import tensorflow as tf def foo(): x = tf.Variable(tf.ones([1])) y = tf.ones([1]) return x+y if __name__ == '__main__': print(foo().numpy())
内容的提问来源于stack exchange,提问作者mehrtash
相关产品推荐
相关产品推荐

