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

tf.contrib.framework.init_from_checkpoint函数无法正常工作求助

问题原因与解决方案

这个问题主要出在两个关键的使用误区上,都是针对tf.contrib.framework.init_from_checkpoint的特性理解不到位导致的:

1. 变量创建方式不匹配

init_from_checkpoint是专门为tf.get_variable设计的API,它会修改get_variable创建的变量的初始化逻辑。但你用了tf.Variable来定义新模型的变量——tf.Variable的初始化器在创建时就被固定为你指定的[0,0,0],init_from_checkpoint无法覆盖这个默认初始化器,自然加载不了checkpoint里的值。

2. 全局初始化覆盖了checkpoint加载逻辑

就算你用对了变量创建方式,直接调用tf.global_variables_initializer()也会出问题:这个操作会强制执行所有变量的默认初始化器(也就是你设置的[0,0,0]),完全忽略init_from_checkpoint设置的从checkpoint加载的逻辑。


修复后的代码示例

把tf.Variable换成tf.get_variable,并使用tf.contrib.framework.initialize_all_variables()来初始化变量(它会尊重init_from_checkpoint的设置):

import tensorflow as tf
model_name = "./my_model.ckp"
### MY MODEL IS COMPOSED BY 2 VARIABLES
with tf.variable_scope("A"):
    A = tf.get_variable("A1", initializer=[1, 2, 3])
with tf.variable_scope("B"):
    B = tf.get_variable("B1", initializer=[4, 5, 6])
# INITIALIZING AND SAVING THE MODEL
with tf.Session() as sess:
    tf.global_variables_initializer().run(session=sess)
    print(sess.run([A, B]))
    saver = tf.train.Saver()
    saver.save(sess, model_name)
#### CLEANING UP
tf.reset_default_graph()
### CREATING OTHER "MODEL"
with tf.variable_scope("C"):
    # 使用get_variable替代Variable,指定shape和类型即可,不需要默认初始值
    A = tf.get_variable("A1", shape=[3], dtype=tf.int32)
with tf.variable_scope("B"):
    B = tf.get_variable("B1", shape=[3], dtype=tf.int32)
# MAPPING THE VARIABLES FROM MY CHECKPOINT TO MY NEW SET OF VARIABLES
tf.contrib.framework.init_from_checkpoint(
    model_name, {"A/": "C/", "B/": "B/"})
with tf.Session() as sess:
    # 使用initialize_all_variables,它会执行init_from_checkpoint设置的加载逻辑
    tf.contrib.framework.initialize_all_variables().run(session=sess)
    print(sess.run([A, B]))

运行这段代码,第二次输出就会符合预期:[array([1, 2, 3], dtype=int32), array([4, 5, 6], dtype=int32)]


补充说明

如果你坚持要使用tf.Variable(不推荐),可以通过以下方式实现:

  1. 先调用tf.global_variables_initializer()初始化变量到默认值
  2. 手动获取init_from_checkpoint添加的赋值操作并运行

但这种方式需要额外收集操作,代码会更繁琐,所以优先推荐使用tf.get_variable的方案。

内容的提问来源于stack exchange,提问作者Tiago Freitas Pereira

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.15 07:10:34