【发布时间】: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