【问题标题】:TensorFlow: "Cannot capture a stateful node by value" in tf.contrib.data APITensorFlow:tf.contrib.data API 中的“无法按值捕获有状态节点”
【发布时间】:2017-11-06 12:43:51
【问题描述】:

对于迁移学习,人们通常使用网络作为特征提取器来创建特征数据集,在该数据集上训练另一个分类器(例如 SVM)。

我想使用 Dataset API (tf.contrib.data) 和 dataset.map() 来实现这一点:

# feature_extractor will create a CNN on top of the given tensor
def features(feature_extractor, ...):
    dataset = inputs(...)  # This creates a dataset of (image, label) pairs

    def map_example(image, label):
        features = feature_extractor(image, trainable=False)
        #  Leaving out initialization from a checkpoint here... 
        return features, label

    dataset = dataset.map(map_example)

    return dataset

为数据集创建迭代器时执行此操作失败。

ValueError: Cannot capture a stateful node by value.

这是真的,网络的内核和偏差是变量,因此是有状态的。对于这个特定的示例,它们不必是。

有没有办法让 Ops,特别是 tf.Variable 对象无状态?

由于我使用的是tf.layers,因此我不能简单地将它们创建为常量,并且设置trainable=False 也不会创建常量,只是不会将变量添加到GraphKeys.TRAINABLE_VARIABLES 集合中。

【问题讨论】:

    标签: tensorflow tensorflow-datasets


    【解决方案1】:

    不幸的是,tf.Variable 本质上是有状态的。但是,仅当您使用 Dataset.make_one_shot_iterator() 创建迭代器时才会出现此错误。* 为避免此问题,您可以改用 Dataset.make_initializable_iterator(),但需要注意的是,您还必须在返回的迭代器上运行 iterator.initializer之后 为输入管道中使用的 tf.Variable 对象运行初始化程序。


    * 造成这种限制的原因是 Dataset.make_one_shot_iterator() 的实现细节以及它用于封装数据集定义的工作在进行中的 TensorFlow 函数 (Defun) 支持。由于使用查找表和变量等有状态资源比我们最初想象的更受欢迎,因此我们正在研究放宽这一限制的方法。

    【讨论】:

    • 对不起,有状态/无状态节点的概念是什么?提前感谢@mrry
    猜你喜欢
    • 2020-04-08
    • 2016-10-10
    • 1970-01-01
    • 2021-10-03
    • 1970-01-01
    • 2021-09-08
    • 2022-01-19
    • 2018-06-06
    • 1970-01-01
    相关资源
    最近更新 更多