TensorFlow加载checkpoint报错Key global_step not found问题问询
问题定位
报错核心原因:官方发布的预训练权重仅保存了模型推理相关参数,未存储global_step训练步数变量,但训练代码中构造的tf.train.Saver将global_step加入了待恢复变量列表,TensorFlow默认加载逻辑要求Saver指定的所有变量必须在checkpoint中存在匹配键,因此抛出找不到Key的错误。
另外代码中存在一处逻辑错误:在会话上下文内调用tf.reset_default_graph()会直接清空当前已构建的计算图,即使权重加载成功,后续也会因找不到已定义的计算节点报错。
排查步骤
- 执行以下代码打印checkpoint内存储的所有变量键,可直接验证权重中确实不存在
global_step项:
from tensorflow.python.tools import inspect_checkpoint as chkp # 替换为本地checkpoint的实际路径 chkp.print_tensors_in_checkpoint_file( "/home/stereo_magnification/stereo-magnification-master/model/siggraph_model_20180701/model.latest", tensor_name='', all_tensors=False )
- 核对Saver绑定的变量列表:当前代码中Saver加载范围为所有模型变量加
global_step,超出了预训练权重的存储范围。 - 定位冗余错误代码:删除
saver.restore()前的tf.reset_default_graph()调用,该行在已启动的会话上下文中无任何正向作用,只会破坏计算图结构。
解决方法
根据使用场景二选一即可:
场景1:加载官方预训练权重做微调
这是预训练权重的标准使用场景,不需要继承原训练步数,加载时跳过权重中不存在的global_step变量,使用代码中初始化的0值作为训练起始步数即可:
- 删除冗余的
tf.reset_default_graph()行 - 将原代码中直接调用
saver.restore(sess, checkpoint)的逻辑替换为以下代码,仅加载权重中存在的匹配变量:
# 读取checkpoint内存储的变量列表 reader = tf.train.NewCheckpointReader(checkpoint) ckpt_var_names = set(reader.get_variable_to_shape_map().keys()) # 筛选出当前图中与checkpoint匹配的模型参数 loadable_vars = [ var for var in tf.model_variables() if var.name.split(':')[0] in ckpt_var_names ] # 构造专属加载器恢复权重 load_saver = tf.train.Saver(loadable_vars) load_saver.restore(sess, checkpoint) tf.logging.info("Loaded %d pretrained parameters, global_step initialized to 0", len(loadable_vars))
修改后原有训练逻辑不需要调整,global_step会从0开始正常计数,不影响微调训练。
场景2:从自己之前中断的训练任务断点续训
这种场景不要指向官方发布的预训练权重路径,将checkpoint_dir指向你自己训练过程中输出的checkpoint目录即可,自己训练保存的checkpoint会自动存储global_step变量,原代码逻辑可正常运行。官方预训练权重仅包含模型参数,本身就不是为断点续训场景准备的。
内容的提问来源于stack exchange,提问作者Carl Mier
相关产品推荐
相关产品推荐

