TensorFlow:恢复的变量数值异常,仅恢复部分变量时遇问题
部分TensorFlow变量恢复异常的排查与解决
我来帮你梳理下这个问题,结合TensorFlow的变量机制,大概率是这几个环节出了问题:
1. 别让全局初始化覆盖了恢复的变量
你提到在会话前初始化了权重,但如果训练时你先执行了tf.global_variables_initializer(),这个操作会把所有变量重新初始化,包括你要恢复的w1,直接覆盖掉从checkpoint加载的值,看起来就像是随机的异常情况。
解决办法是拆分初始化操作,只初始化那些不需要恢复的变量:
# 定义需要恢复的目标变量 weights = { '1': tf.Variable(tf.random_normal([n_input, n_hidden_1], mean=0, stddev=tf.sqrt(2*1.67/(n_input+n_hidden_1))), name='w1') } # 定义其他需要初始化的变量(比如偏置、未预训练的层权重) other_trainable_vars = [tf.Variable(tf.zeros([n_hidden_1]), name='b1'), ...] # 分别创建初始化操作 target_weights_init = tf.variables_initializer(weights.values()) others_init = tf.variables_initializer(other_trainable_vars) with tf.Session() as sess: # 第一步:先恢复目标变量,这一步要在初始化其他变量之前 weights_saver = tf.train.Saver(var_list=weights) weights_saver.restore(sess, "./your_checkpoint_dir/model.ckpt") # 第二步:只初始化不需要恢复的变量 sess.run(others_init) # 接下来再启动训练流程
2. 确认变量名完全匹配
如果之前保存模型时,这个变量的名字和你现在定义的w1不一致(比如之前定义时加了命名空间、或者name参数写的不一样),Saver会找不到对应的变量,直接跳过恢复操作,导致变量还是初始的随机值。
你可以先检查旧checkpoint里的变量名:
from tensorflow.python.tools.inspect_checkpoint import print_tensors_in_checkpoint_file # 打印checkpoint里所有变量的名字和形状 print_tensors_in_checkpoint_file("./old_checkpoint_path/model.ckpt", tensor_name='', all_tensors=False, all_tensor_names=True)
确保你现在定义的weights['1']的name属性和checkpoint里的完全一致,再用这个name构建var_list。
3. 检查Saver的创建顺序
正确的流程应该是:先定义完所有需要恢复的变量,再创建对应的tf.train.Saver。如果在创建Saver之后又修改了变量结构,或者新增了变量,可能会导致Saver没有正确关联到目标变量,恢复操作失效。
4. 立刻验证恢复结果
在恢复操作执行后,马上打印变量的值,确认是否成功加载了checkpoint里的内容,而不是初始的随机值:
with tf.Session() as sess: weights_saver.restore(sess, "./your_checkpoint_dir/model.ckpt") # 打印恢复后的变量值,对比checkpoint里的预期值 print("恢复后的w1部分值:", sess.run(weights['1'])[:5, :5]) # 再执行其他初始化和训练操作
这样能快速定位是恢复环节没生效,还是训练过程中其他逻辑意外修改了变量。
内容的提问来源于stack exchange,提问作者L.Ech
相关产品推荐
相关产品推荐

