【问题标题】:How to get the global_step when restoring checkpoints in Tensorflow?在 Tensorflow 中恢复检查点时如何获取 global_step?
【发布时间】:2016-03-20 11:25:09
【问题描述】:

我正在像这样保存我的会话状态:

self._saver = tf.saver()
self._saver.save(self._session, '/network', global_step=self._time)

当我稍后恢复时,我想获取我从中恢复的检查点的 global_step 值。这是为了从中设置一些超参数。

执行此操作的 hacky 方法是运行并解析检查点目录中的文件名。但是必须有一种更好的内置方法才能做到这一点?

【问题讨论】:

    标签: tensorflow


    【解决方案1】:

    一般模式是有一个global_step 变量来跟踪步骤

    global_step = tf.Variable(0, name='global_step', trainable=False)
    train_op = optimizer.minimize(loss, global_step=global_step)
    

    然后你可以保存

    saver.save(sess, save_path, global_step=global_step)
    

    当您恢复时,global_step 的值也会恢复

    【讨论】:

    • 这不起作用,每次我恢复训练时,global_step 变量都会重置为 0
    • 这意味着您保存到检查点的 global_step 为 0,或者您在恢复后将其重新初始化为 0
    • 这将是一个很好的解决方案,但是如果saver.restore 可以返回 global_step,那就很简单了。我们可以只做 'global_step=saver.restore(...)' 你认为 tensorflow 团队可能对这个方向感兴趣吗?
    • 似乎在某些情况下它可能有用,但也似乎需要做很多工作——现在 TF 1.0 已经发布,对 API 的任何更改都必须经过 API 审查
    • @YaroslavBulatov 这不适用于此处的初始 v3 训练:github.com/tensorflow/models/tree/master/inception/inception 恢复模型后全局步长始终为 0。
    【解决方案2】:

    这有点小题大做,但其他答案对我根本不起作用

    ckpt = tf.train.get_checkpoint_state(checkpoint_dir) 
    
    #Extract from checkpoint filename
    step = int(os.path.basename(ckpt.model_checkpoint_path).split('-')[1])
    

    2017 年 9 月更新

    我不确定这是否由于更新而开始工作,但以下方法似乎可以有效地让 global_step 正确更新和加载:

    创建两个操作。一个用来保存 global_step,另一个用来增加它:

        global_step = tf.Variable(0, trainable=False, name='global_step')
        increment_global_step = tf.assign_add(global_step,1,
                                                name = 'increment_global_step')
    

    现在在您的训练循环中,每次运行训练操作时都运行增量操作。

    sess.run([train_op,increment_global_step],feed_dict=feed_dict)
    

    如果您想在任何时候将全局步长值检索为整数,只需在加载模型后使用以下命令:

    sess.run(global_step)
    

    这对于创建文件名或计算当前时期很有用,而无需第二个 tensorflow 变量来保存该值。例如,在加载时计算当前 epoch 类似于:

    loaded_epoch = sess.run(global_step)//(batch_size*num_train_records)
    

    【讨论】:

      【解决方案3】:

      我和 Lawrence Du 有同样的问题,我找不到通过恢复模型来获取 global_step 的方法。所以我将his hack 应用到我正在使用的the inception v3 training code in the Tensorflow/models github repo。下面的代码还包含与pretrained_model_checkpoint_path 相关的修复。

      如果您有更好的解决方案,或者知道我缺少什么,请发表评论!

      无论如何,这段代码对我有用:

      ...
      
      # When not restoring start at 0
      last_step = 0
      if FLAGS.pretrained_model_checkpoint_path:
          # A model consists of three files, use the base name of the model in
          # the checkpoint path. E.g. my-model-path/model.ckpt-291500
          #
          # Because we need to give the base name you can't assert (will always fail)
          # assert tf.gfile.Exists(FLAGS.pretrained_model_checkpoint_path)
      
          variables_to_restore = tf.get_collection(
              slim.variables.VARIABLES_TO_RESTORE)
          restorer = tf.train.Saver(variables_to_restore)
          restorer.restore(sess, FLAGS.pretrained_model_checkpoint_path)
          print('%s: Pre-trained model restored from %s' %
                (datetime.now(), FLAGS.pretrained_model_checkpoint_path))
      
          # HACK : global step is not restored for some unknown reason
          last_step = int(os.path.basename(FLAGS.pretrained_model_checkpoint_path).split('-')[1])
      
          # assign to global step
          sess.run(global_step.assign(last_step))
      
      ...
      
      for step in range(last_step + 1, FLAGS.max_steps):
      
        ...
      

      【讨论】:

      • 此方法不适用于 (download.tensorflow.org/models/inception_v3_2016_08_28.tar.gz) 提供的官方预训练 inception v3 模型检查点,因为检查点文件名仅包含 inception_v3.ckpt
      • @VipinPillai 在所有预训练模型中,全局步长重置为零。因此,该模型可用于初始化图以进行微调,而无需设置全局步长。
      【解决方案4】:

      您可以使用global_step 变量来跟踪步骤,但如果在您的代码中,您正在初始化或将此值分配给另一个step 变量,则可能不一致。

      例如,你定义你的global_step 使用:

      global_step = tf.Variable(0, name='global_step', trainable=False)
      

      分配给您的训练操作:

      train_op = optimizer.minimize(loss, global_step=global_step)
      

      保存在您的检查点中:

      saver.save(sess, checkpoint_path, global_step=global_step)
      

      并从您的检查点恢复:

      saver.restore(sess, checkpoint_path) 
      

      global_step 的值也已恢复,但如果您将其分配给另一个变量,例如 step,那么您必须执行以下操作:

      step = global_step.eval(session=sess)
      

      变量step,包含检查点中最后保存的global_step。

      最好也将图中的 global_step 定义为零变量(如前所述):

      global_step = tf.train.get_or_create_global_step()
      

      如果存在,这将获得您的最后一个 global_step,如果不存在,则创建一个。

      【讨论】:

      • 这是我在这件事上看到的最干净的解决方案! +1。
      【解决方案5】:

      TL;DR

      作为张量流变量(将在会话中评估)

      global_step = tf.train.get_or_create_global_step()
      # use global_step variable to calculate your hyperparameter 
      # this variable will be evaluated later in the session
      saver = tf.train.Saver()
      with tf.Session() as sess:
          # restore all variables from checkpoint
          saver.restore(sess, checkpoint_path)
          # than init table and local variables and start training/evaluation ...
      

      或者:作为 numpy 整数(没有任何会话):

      reader = tf.train.NewCheckpointReader(absolute_checkpoint_path)
      global_step = reader.get_tensor('global_step')
      


      长答案

      至少有两种方法可以从检查点检索全局。作为 tensorflow 变量或 numpy 整数。如果global_step 没有在Saver 的save 方法中作为参数提供,则无法解析文件名。对于预训练模型,请参阅答案末尾的备注。

      作为 TensorFlow 变量

      如果您需要global_step 变量来计算一些超参数,您可以使用tf.train.get_or_create_global_step()。这将返回一个张量流变量。因为该变量将在会话稍后进行评估,所以您只能使用 tensorflow 操作来计算您的超参数。所以例如:max(global_step, 100) 将不起作用。您必须使用等效于 tensorflow 的 tf.maximum(global_step, 100),可以在会话稍后进行评估。

      在会话中,您可以使用saver.restore(sess, checkpoint_path) 使用检查点初始化全局步骤变量

      global_step = tf.train.get_or_create_global_step()
      # use global_step variable to calculate your hyperparameter 
      # this variable will be evaluated later in the session
      hyper_parameter = tf.maximum(global_step, 100) 
      saver = tf.train.Saver()
      with tf.Session() as sess:
          # restore all variables from checkpoint
          saver.restore(sess, checkpoint_path)
          # than init table and local variables and start training/evaluation ...
      
          # for verification you can print the global step and your hyper parameter
          print(sess.run([global_step, hyper_parameter]))
      

      或者:作为 numpy 整数(无会话)

      如果您需要全局 step 变量作为标量而不启动会话,您也可以直接从检查点文件中读取此变量。你只需要一个NewCheckpointReader。由于旧 tensorflow 版本中有 bug,您应该将检查点文件的路径转换为绝对路径。使用阅读器,您可以将模型的所有张量作为 numpy 变量。 全局步骤变量的名称是一个常量字符串tf.GraphKeys.GLOBAL_STEP,定义为'global_step'。

      absolute_checkpoint_path = os.path.abspath(checkpoint_path)
      reader = tf.train.NewCheckpointReader(absolute_checkpoint_path)
      global_step = reader.get_tensor(tf.GraphKeys.GLOBAL_STEP)
      

      对预训练模型的说明:在大多数在线可用的预训练模型中,全局步长重置为零。因此,这些模型可用于初始化模型参数以进行微调,而不会覆盖全局步骤。

      【讨论】:

        【解决方案6】:

        现在的 0.10rc0 版本好像不一样了,没有 tf.saver() 了。现在是 tf.train.Saver()。此外,save 命令将信息添加到 global_step 的 save_path 文件名中,因此我们不能只在同一个 save_path 上调用 restore,因为那不是实际的保存文件。

        我现在看到的最简单的方法是使用 SessionManager 以及这样的保护程序:

        my_checkpoint_dir = "/tmp/checkpoint_dir"
        # make a saver to use with SessionManager for restoring
        saver = tf.train.Saver()
        # Build an initialization operation to run below.
        init = tf.initialize_all_variables()
        # use a SessionManager to help with automatic variable restoration
        sm = tf.train.SessionManager()
        # try to find the latest checkpoint in my_checkpoint_dir, then create a session with that restored
        # if no such checkpoint, then call the init_op after creating a new session
        sess = sm.prepare_session("", init_op=init, saver=saver, checkpoint_dir=my_checkpoint_dir))
        

        就是这样。现在你有一个从 my_checkpoint_dir 恢复的会话(在调用它之前确保该目录存在),或者如果那里没有检查点,那么它会创建一个新会话并调用 init_op 来初始化你的变量。

        当您想保存时,您只需保存到该目录中所需的任何名称并将 global_step 传入。这是一个示例,我将循环中的 step 变量保存为 global_step,因此如果您杀死程序并重新启动它以恢复检查点:

        checkpoint_path = os.path.join(my_checkpoint_dir, 'model.ckpt')
        saver.save(sess, checkpoint_path, global_step=step)
        

        这会在 my_checkpoint_dir 中创建文件,例如“model.ckpt-1000”,其中 1000 是传入的 global_step。如果它继续运行,那么您会得到更像“model.ckpt-2000”的文件。当程序重新启动时,上面的 SessionManager 会选择其中最新的一个。 checkpoint_path 可以是您想要的任何文件名,只要它位于 checkpoint_dir 中即可。 save() 将创建附加了 global_step 的文件(如上所示)。它还创建一个“检查点”索引文件,这是 SessionManager 找到最新保存检查点的方式。

        【讨论】:

          【解决方案7】:

          请注意我关于全局步骤保存和恢复的解决方案。

          保存:

          global_step = tf.Variable(0, trainable=False, name='global_step')
          saver.save(sess, model_path + model_name, global_step=_global_step)
          

          恢复:

          if os.path.exists(model_path):
              saver.restore(sess, tf.train.latest_checkpoint(model_path))
              print("Model restore finished, current globle step: %d" % global_step.eval())
          

          【讨论】:

            【解决方案8】:

            变量未按预期恢复的原因很可能是因为它是在您的 tf.Saver() 对象创建之后创建的。

            当您没有明确指定var_list 或为var_list 指定None 时,创建tf.Saver() 对象的位置很重要。许多程序员的预期行为是,当调用save() 方法时,图中的所有变量都会被保存,但事实并非如此,也许应该这样记录。图表中所有变量的快照会在对象创建时保存。

            除非您遇到任何性能问题,否则在您决定保存进度时立即创建保护程序对象是最安全的。否则,请确保在创建所有变量后创建保护程序对象。

            另外,传递给saver.save(sess, save_path, global_step=global_step) 的global_step 只是一个用于创建文件名的计数器,与是否将其恢复为global_step 变量无关。这是一个参数误称 IMO,因为如果您在每个时期结束时保存进度,最好将您的时期号传递给此参数。

            【讨论】:

              猜你喜欢
              • 1970-01-01
              • 2018-05-23
              • 1970-01-01
              • 1970-01-01
              • 1970-01-01
              • 1970-01-01
              • 2019-04-03
              • 2017-07-30
              • 2017-11-27
              相关资源
              最近更新 更多