tf.train.init_from_checkpoint为何不初始化tf.Variable创建的变量?
tf.train.init_from_checkpoint只初始化tf.get_variable创建的变量,而非tf.Variable? 我发现TensorFlow中的tf.train.init_from_checkpoint方法仅会初始化通过tf.get_variable创建的变量,而不会处理tf.Variable创建的变量。以下是我的测试示例:
首先创建两个变量并保存到 checkpoint:
import tensorflow as tf tf.Variable(1.0, name='foo') tf.get_variable('bar',initializer=1.0) saver = tf.train.Saver() with tf.Session() as sess: tf.global_variables_initializer().run() saver.save(sess, './model', global_step=0)
如果使用tf.train.Saver加载 checkpoint,两个变量都能正常恢复为1.0:
import tensorflow as tf foo = tf.Variable(0.0, name='foo') bar = tf.get_variable('bar', initializer=0.0) saver = tf.train.Saver() with tf.Session() as sess: saver.restore(sess, './model-0') print(f'foo: {foo.eval()} bar: {bar.eval()}') # 输出: foo: 1.0 bar: 1.0
但使用tf.train.init_from_checkpoint时,只有tf.get_variable创建的bar被恢复为1.0,tf.Variable创建的foo仍保持初始值0.0:
import tensorflow as tf foo = tf.Variable(0.0, name='foo') bar = tf.get_variable('bar', initializer=0.0) tf.train.init_from_checkpoint('./model-0', {'/':'/'}) with tf.Session() as sess: tf.global_variables_initializer().run() print(f'foo: {foo.eval()} bar: {bar.eval()}') # 输出: foo: 0.0 bar: 1.0
请问这是预期行为吗?如果是,背后的原因是什么?
这确实是预期行为,背后和TensorFlow中两种变量创建方式的设计定位以及tf.train.init_from_checkpoint的功能目标直接相关:
首先,
tf.get_variable是TensorFlow变量共享机制的核心,它创建的变量会被注册到**变量作用域(variable scope)**的管理体系中,这类变量的命名、复用逻辑明确可控,通常用于构建可复用的模型组件(比如卷积层、LSTM单元),也是预训练模型迁移时最常需要复用的对象。而
tf.Variable是一种更直接的变量创建方式,每次调用都会生成一个新的变量实例(即使名称相同,也会自动添加后缀区分),它并不属于变量共享体系的一部分,更多用于创建一次性的、不需要复用的变量。
tf.train.init_from_checkpoint的设计初衷就是为了方便迁移预训练模型中的共享变量,它会在全局变量初始化阶段,针对那些通过tf.get_variable创建的、与checkpoint中名称匹配的变量进行赋值;而tf.Variable创建的变量不在它的处理范围内——因为这类变量通常不被视为“可复用的预训练组件”,所以会保持你在代码中设置的初始值。
简单来说,init_from_checkpoint是专门为变量共享场景设计的初始化工具,而tf.Saver则是通用的 checkpoint 恢复工具,会处理所有全局变量,这就是两者表现不同的核心原因。
内容的提问来源于stack exchange,提问作者user209974

