【问题标题】:Split tensorflow dataset in dataset per class在每个类的数据集中拆分 tensorflow 数据集
【发布时间】:2018-09-12 16:13:14
【问题描述】:

我有一个从一个 tfrecord 文件创建的数据集。该数据集包含 5 个不同的类。

现在我想从每个批次中创建具有固定数量元素(例如 8 个)的批次。所以它应该创建包含每个类的 8 个元素的 40 个元素的批次。

tf.data 可以做到吗?

【问题讨论】:

    标签: python tensorflow tensorflow-datasets


    【解决方案1】:

    最简单的做法是(也许不是很方便):

    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)

    【讨论】:

      猜你喜欢
      • 1970-01-01
      • 2018-12-10
      • 2020-06-24
      • 2018-04-19
      • 2020-12-29
      • 1970-01-01
      • 1970-01-01
      • 2022-01-01
      • 1970-01-01
      相关资源
      最近更新 更多