【问题标题】:How can I read endlessly from a Tensorflow tf.data.Dataset?如何从 Tensorflow tf.data.Dataset 无休止地读取数据?
【发布时间】:2018-03-19 15:33:05
【问题描述】:

我正在将旧数据层(使用队列)切换到“新”和推荐的数据集 API。我是第一次使用它,所以我提供了代码示例,以防我遇到根本错误。

我从一个生成器创建我的数据集(它将读取一个文件,并提供 n 个样本)。这是一个小数据集和 n_iterations >> n_samples,所以我只是想一遍又一遍地阅读这个数据集,最好是随机播放。

sample_set = tf.data.Dataset.from_generator( data_generator(filename),  
    (tf.uint8, tf.uint8), (tf.TensorShape([256,256,4]), tf.TensorShape([256,256,1]))
)

使用数据生成器:

class data_generator:
    def __init__(self, filename):
        self.filename= filename

    def __call__(self):
        with filename.open() as f:
           for idx in f: yield img[idx], label[idx]

为了实际使用数据,我知道我需要定义一个Iterator

sample = sample_set.make_one_shot_iterator().get_next()

然后我们就可以读取数据了

while True:
    try: my_sample = sess.run(sample)
    except tf.errors.OutOfRangeError: break   # this happens after dset is read once

但所有可用的迭代器似乎都是“有限的”,因为它们只读取数据集一次。

有没有一种简单的方法可以让数据集的读取无止境?

【问题讨论】:

    标签: tensorflow tensorflow-datasets


    【解决方案1】:

    数据集有repeatshuffle 方法。

    BUF_SIZE = 100 # choose it depending on your data
    sample_set = tf.data.Dataset.from_generator( data_generator(filename),  
        (tf.uint8, tf.uint8), (tf.TensorShape([256,256,4]), 
        tf.TensorShape([256,256,1]))
    ).repeat().shuffle(BUF_SIZE)
    

    【讨论】:

    • 完美,正是我的想法!我会选择这个,因为随机播放也非常有用。
    【解决方案2】:

    如果您不将明确的count 传递给它,Dataset.repeat() 转换将无休止地重复数据集:

    sample_set = tf.data.Dataset.from_generator(
        data_generator(filename), (tf.uint8, tf.uint8),
        (tf.TensorShape([256,256,4]), tf.TensorShape([256,256,1])))
    
    # Repeats `sample_set` endlessly.
    sample_set = sample_set.repeat()
    
    sample = sample_set.make_one_shot_iterator().get_next()
    

    【讨论】:

    • dataset.repeat() 功能是否以my answer 的方式实现(即有点像这样:try: my_sample = sess.run(sample) except tf.errors.OutOfRangeError: sess.run(sample_set_init_op) # re-initialize on same dataset)?因为我的日志仍然显示“Out of range: ...-Errors,但它继续运行。
    • 每次原始data_generator() 到达一个重复的末尾时,它可能会打印这个。 (我相信这个额外的日志记录在更新的 TensorFlow 版本中被删除了。)
    【解决方案3】:

    reinitializable Iterator 可以在同一个数据集上重新初始化,所以这段代码会一遍又一遍地读取同一个数据集:

    sample = tf.data.Iterator.from_structure(sample_set.output_types,
                                             sample_set.output_shapes).get_next()
    
    sample_it.make_initializer(sample_set)     # create initialize op
    
    with tf.Session(config=config) as sess:
        sess.run(sample_set_init_op)           # initialize in the beginning
    
        while True:
            try: 
                 my_sample = sess.run(sample)
            except tf.errors.OutOfRangeError:
                 sess.run(sample_set_init_op)  # re-initialize on same dataset
    

    【讨论】:

    • 这可能对其他人有用,但是除了感觉很脏以产生和捕获这样的错误之外,这需要“重新启动”图表。在这个玩具示例中,图是微不足道的,但通常数据输入是大型训练图的一部分。对于我的带有检查点管理和许多其他东西的主循环,包括output = sess.run(),我使用了其他人提供的包装器,所以基本上我无法访问sess.run(),并且希望模型能够处理“无限数据读取”本身(不需要我重新运行错误模型)。
    猜你喜欢
    • 1970-01-01
    • 2018-06-26
    • 2016-09-26
    • 2011-10-27
    • 2020-01-24
    • 1970-01-01
    • 2019-05-07
    • 1970-01-01
    • 1970-01-01
    相关资源
    最近更新 更多