【问题标题】:TensorFlow: Restoring variables from from multiple checkpointsTensorFlow:从多个检查点恢复变量
【发布时间】:2016-06-14 12:29:45
【问题描述】:

我有以下情况:

  • 我有 2 个模型用 2 个单独的脚本编写:

  • 模型A由变量a1a2a3组成,写成A.py

  • 模型B由变量b1b2b3组成,用B.py编写

A.pyB.py中,我都有一个tf.train.Saver保存所有局部变量的检查点,我们分别调用检查点文件ckptAckptB

我现在想制作一个使用a1b1 的模型C。我可以通过使用 var_scope 在 A 和 C 中使用与 a1 完全相同的变量名(b1 也是如此)。

问题是我如何将a1b1ckptAckptB 加载到模型C 中?例如,以下是否可行?

saver.restore(session, ckptA_location)
saver.restore(session, ckptB_location)

如果您尝试两次恢复同一个会话,会引发错误吗?它会抱怨没有为额外的变量分配“插槽”(b2b3a2a3),还是会简单地恢复它可以恢复的变量,只有在有一些变量时才会抱怨C 中其他未初始化的变量?

我现在正在尝试编写一些代码来测试这一点,但我希望看到解决此问题的规范方法,因为在尝试重新使用一些预训练的权重时经常会遇到这种情况。

谢谢!

【问题讨论】:

    标签: tensorflow


    【解决方案1】:

    如果您尝试使用保护程序(默认情况下代表所有六个变量)从不包含保护程序所代表的所有变量的检查点恢复,您将获得tf.errors.NotFoundError。 (但是请注意,只要所有请求的变量都存在于相应的文件中,您就可以在同一会话中多次调用 Saver.restore() 来获取任何变量子集。)

    规范的方法是定义两个独立的tf.train.Saver 实例,覆盖完全包含在单个检查点中的每个变量子集。例如:

    saver_a = tf.train.Saver([a1])
    saver_b = tf.train.Saver([b1])
    
    saver_a.restore(session, ckptA_location)
    saver_b.restore(session, ckptB_location)
    

    根据您的代码的构建方式,如果您在本地范围内有指向称为a1b1tf.Variable 对象的指针,您可以在此处停止阅读。

    另一方面,如果变量a1b1 定义在不同的文件中,您可能需要采取一些创造性的措施来检索指向这些变量的指针。虽然并不理想,但人们通常会使用通用前缀,例如如下(假设变量名分别为"a1:0""b1:0"):

    saver_a = tf.train.Saver([v for v in tf.all_variables() if v.name == "a1:0"])
    saver_b = tf.train.Saver([v for v in tf.all_variables() if v.name == "b1:0"])
    

    最后一点:您不必费力地确保变量在 A 和 C 中具有相同的名称。您可以将 name-to-Variable 字典作为第一个参数传递给 @987654336 @构造函数,从而将检查点文件中的名称重新映射到代码中的Variable对象。如果A.pyB.py 具有类似命名的变量,或者如果在C.py 中您想在tf.name_scope() 中组织这些文件中的模型代码,这将有所帮助。

    【讨论】:

    • 第一个代码片段是指应该用 C.py 编写的东西,对吗?它写在图形定义的末尾,其中 a1 a2 a3 和 b1 b2 b3 在 C.py 中定义?我最初的问题是,如果在 C 中只定义了 a1 和 b1 而没有其他定义呢?
    • 道歉 - 更新了答案以涵盖此案例。如果你在 C.py 中重新定义变量,那么事情就容易多了!
    • 再跟进:当我们执行这条语句时,“saver_a.restore(session, ckptA_location)” saver_a 用单个变量 [a1] 进行实例化,没有别的,但 ckptA 包含所有 a1 的值, a2,a3。您是说这不是问题,因为保护程序只会在 ckpt 中搜索 a1 并将 a1 恢复为模型 C,而忽略 a2 和 a3?最后一个问题:C 中的 a1 将被标识为 ckptA 中的 a1,只要两个 a1 都以相同的名称实例化(在 C 和 A 中),对吗?
    • 没错(两部分)。对于后一部分,如果a1 在检查点中具有不同的名称,您可以指定显式名称到Variable 映射。默认情况下,Saver 构造函数使用Variable.name 属性进行查找。
    猜你喜欢
    • 1970-01-01
    • 2018-09-26
    • 2018-05-23
    • 2017-06-17
    • 1970-01-01
    • 2016-09-29
    • 1970-01-01
    • 2017-07-30
    • 1970-01-01
    相关资源
    最近更新 更多