tf.train.Saver从何处收集待保存变量?为何需在图末尾定义?
为什么
tf.train.Saver必须放在计算图变量定义之后才能保存全部变量? 其实核心原因很简单:tf.train.Saver是在你创建它的那一刻,从当前图的GLOBAL_VARIABLES集合里一次性读取所有变量的,之后再新增的变量它根本“看不到”。
具体原理拆解
- 当你创建
tf.Variable时,TensorFlow默认会把这个变量自动加入到tf.GraphKeys.GLOBAL_VARIABLES集合中(除非你手动设置collections=None)。 tf.train.Saver的构造函数在被调用的瞬间,就会去读取这个集合里的所有变量,生成对应的保存/恢复逻辑。一旦初始化完成,这个Saver实例就不会再去监听集合的变化了。
结合你的测试代码分析
看你写的这段代码:
def how_saver_work(): g = tf.Graph() with g.as_default(): a = tf.Variable(1, name='a') b = tf.Variable(2, name='b') print(tf.get_collection(tf.GraphKeys.GLOBAL_VARIABLES)) # 这里输出[a, b] saver = tf.train.Saver() # 此时Saver只捕获了a和b c = tf.Variable(3, name='c') print(tf.get_collection(tf.GraphKeys.GLOBAL_VARIABLES)) # 这里虽然能看到[a, b, c],但Saver已经定型了
这时候你用这个saver去保存模型,只会保存a和b两个变量,c不会被包含进去——因为Saver创建的时候,c还没被加入到GLOBAL_VARIABLES集合里。
解决办法
- 如果要保存图中所有变量,最稳妥的方式就是把
tf.train.Saver()的调用放在所有变量定义完成之后,确保它能一次性收集到全部变量。 - 如果你只需要保存特定变量,也可以在创建Saver时显式指定变量列表,比如:
这种情况下,不管变量的定义顺序如何,Saver只会保存你指定的那些变量。saver = tf.train.Saver(var_list=[a, c])
内容的提问来源于stack exchange,提问作者imhuay
相关产品推荐
相关产品推荐

