【发布时间】:2018-09-12 16:13:14
【问题描述】:
我有一个从一个 tfrecord 文件创建的数据集。该数据集包含 5 个不同的类。
现在我想从每个批次中创建具有固定数量元素(例如 8 个)的批次。所以它应该创建包含每个类的 8 个元素的 40 个元素的批次。
tf.data 可以做到吗?
【问题讨论】:
标签: python tensorflow tensorflow-datasets
我有一个从一个 tfrecord 文件创建的数据集。该数据集包含 5 个不同的类。
现在我想从每个批次中创建具有固定数量元素(例如 8 个)的批次。所以它应该创建包含每个类的 8 个元素的 40 个元素的批次。
tf.data 可以做到吗?
【问题讨论】:
标签: python tensorflow tensorflow-datasets
最简单的做法是(也许不是很方便):
a) 准备 5 个不同的TFRecords,每个都只包含一个特定类的元素。
b) 创建5 不同的tf.data.TFRecordDataset 实例,从而创建5 不同的迭代器。
c) 然后在主代码中:
iterators = [....] # Store your iterators in a list
data = list(map(lambda x : x.get_next(), iterators))
data_to_use = tf.concat(....) # Concat your data in one single batch of `40` elements.
a) 仅使用一个 TFRecord。但是创建 5 它的不同实例
b) 在每个实例中,使用tf.data API 的tf.data.filter(predicate) 方法来过滤属于一个特定类的记录。为此,您必须编写一个函数,该函数可以检查每条记录的类。
c) 然后按照上一个解决方案中的步骤c)。
【讨论】: