【问题标题】:how to save only essential parameters in tensorflow?如何在张量流中只保存基本参数?
【发布时间】:2016-08-17 20:24:22
【问题描述】:

默认情况下,会话保护程序会保存所有创建的变量,这会导致检查点文件非常大。我只想保存模型参数和某些会话变量,例如优化器状态和全局步骤。除了在保护程序初始化期间将变量列入白名单之外,还有哪些最佳实践?

【问题讨论】:

  • 我正在尝试的一种方法是为我不想保存的其他变量创建一个集合,例如输出变量。然后将不在该集合中的所有变量列入白名单。

标签: tensorflow


【解决方案1】:

Saver 默认从all_variables() 获取变量列表,这是来自GraphKeys.VARIABLES 集合的所有变量。您可以使用 Variable(..., collections=[]) 从该集合中排除变量。或者你可以把它作为另一个集合,就像在代码库中对非检查点 limit_epochs 变量所做的那样

 with ops.name_scope(name, "limit_epochs", [tensor]) as name:
    zero64 = constant_op.constant(0, dtype=dtypes.int64)
    epochs = variables.Variable(
        zero64, name="epochs", trainable=False,
        collections=[ops.GraphKeys.LOCAL_VARIABLES])

【讨论】:

  • 如果collections=[] 变量永远不会初始化。 collections=[ops.GraphKeys.LOCAL_VARIABLES] 是正确的解决方案。然后需要使用tf.local_variables_initializer() 对其进行初始化才能使用。
【解决方案2】:

经过一些调查(使用不同的批量大小进行检查点并打印出all_variables),我发现我过于担心了。事实上,在 tensorflow 中,Op 的结果并没有被保存,例如y 中的 y = k * x + b。因此,与torch-nn 不同,您很少需要担心非参数会被保存。

【讨论】:

    【解决方案3】:

    您可以创建一个字典,其中包含您要保存的所有变量,键是它们的名称作为字符串。将此字典传递给 saver.save() 函数。 这就是api 的建议。

    【讨论】:

      猜你喜欢
      • 1970-01-01
      • 1970-01-01
      • 2020-09-16
      • 2021-03-14
      • 2018-07-10
      • 1970-01-01
      • 1970-01-01
      • 1970-01-01
      • 2021-08-30
      相关资源
      最近更新 更多