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(不推荐),可以通过以下方式实现:
- 先调用
tf.global_variables_initializer()初始化变量到默认值 - 手动获取
init_from_checkpoint添加的赋值操作并运行
但这种方式需要额外收集操作,代码会更繁琐,所以优先推荐使用tf.get_variable的方案。
内容的提问来源于stack exchange,提问作者Tiago Freitas Pereira

