【问题标题】:resetting a Tensorflow graph after OutOfRangeError when using Dataset使用数据集时在 OutOfRangeError 之后重置 Tensorflow 图
【发布时间】:2017-08-30 09:00:10
【问题描述】:

我正在尝试使用 the from_generator interface for the Dataset API 将多个“轮”输入注入到图表中。

在我的first attempt 上,我使用repeat() function 使生成器连续运行3 次。但是,batch_join call 的批量大小不是每轮迭代次数的偶数倍(批量大小为 3 的 10 次迭代),来自不同“轮”/“时期”的数据" 最终在同一批次中(取决于处理张量的顺序;图中存在一些并行性)。

在我的 second attempt 上,我尝试在每个 epoch 完成后重新运行迭代器。但是,只要tf.errors.OutOfRangeError is thrown,所有后续对sess.run() on the output of the batch call 的调用都会再次抛出OutOfRangeError,即使在rerunning the iterator's initializer 之后也是如此。

我想将多轮输入连续注入到图形中,而不是像第一个示例那样让它们重叠(例如,在批处理选项上使用 allow_smaller_final_batch)。我在自定义 Tensorflow 分支中实例化的一些内核重启起来非常昂贵,例如mmaping 一个 O(10gb) 的文件,所以我想以某种方式充分利用这两个世界。

【问题讨论】:

  • 请添加一个小的可运行示例来显示您的输入管道,以便我们重现您的问题
  • 这应该可以在 Tensorflow 存储库的主分支上运行,这对于数据集的 from_iterator 函数是必需的。如果此示例不适用于该版本,我可以修复它。

标签: tensorflow


【解决方案1】:

我认为问题源于使用 tf.contrib.data.Dataset(支持重新初始化)和 tf.train.batch_join()(使用 TensorFlow 队列和队列运行器,因此不支持重新初始化)。

我并不完全清楚您的代码在做什么,但我认为您可以将整个管道实现为Dataset。替换以下代码片段:

my_iterator = MyIterator(iterations=iterations)
dataset = ds.Dataset.from_generator(my_iterator, 
output_types=my_iterator.output_types, 
output_shapes=my_iterator.output_shapes)
#dataset = dataset.repeat(count=repetitions)
iterator = dataset.make_initializable_iterator()
next_elem = iterator.get_next()

#change constant to 1 or 2 or something to see that the batching is more predictable
ripple_adds = [(tf.stack((next_elem[0], next_elem[1] + constant)),) 
for constant in ripple_add_coefficients]
batch = tf.train.batch_join(ripple_adds, batch_size=batch_size, 
enqueue_many=False, name="sink_queue")

...类似于以下内容:

my_iterator = MyIterator(iterations=iterations)
dataset = tf.contrib.data.from_generator(my_iterator,
                                         output_types=my_iterator.output_types,
                                         output_shapes=my_iterator.output_shapes)

def ripple_add_map_func(x, y):
  return (tf.contrib.data.Dataset.range(num_ripples)
          .map(lambda r: tf.stack([x, y + r])))

dataset = dataset.flat_map(ripple_add_map_func).batch(batch_size)

iterator = dataset.make_initializable_iterator()
batch = iterator.get_next()

【讨论】:

    猜你喜欢
    • 1970-01-01
    • 2022-01-02
    • 1970-01-01
    • 1970-01-01
    • 1970-01-01
    • 2019-02-13
    • 1970-01-01
    • 2019-05-24
    • 2021-06-22
    相关资源
    最近更新 更多