You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.05.12 04:03:58