【发布时间】:2016-08-17 20:24:22
【问题描述】:
默认情况下,会话保护程序会保存所有创建的变量,这会导致检查点文件非常大。我只想保存模型参数和某些会话变量,例如优化器状态和全局步骤。除了在保护程序初始化期间将变量列入白名单之外,还有哪些最佳实践?
【问题讨论】:
-
我正在尝试的一种方法是为我不想保存的其他变量创建一个集合,例如输出变量。然后将不在该集合中的所有变量列入白名单。
标签: tensorflow
默认情况下,会话保护程序会保存所有创建的变量,这会导致检查点文件非常大。我只想保存模型参数和某些会话变量,例如优化器状态和全局步骤。除了在保护程序初始化期间将变量列入白名单之外,还有哪些最佳实践?
【问题讨论】:
标签: tensorflow
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() 对其进行初始化才能使用。
经过一些调查(使用不同的批量大小进行检查点并打印出all_variables),我发现我过于担心了。事实上,在 tensorflow 中,Op 的结果并没有被保存,例如y 中的 y = k * x + b。因此,与torch-nn 不同,您很少需要担心非参数会被保存。
【讨论】:
您可以创建一个字典,其中包含您要保存的所有变量,键是它们的名称作为字符串。将此字典传递给 saver.save() 函数。 这就是api 的建议。
【讨论】: