【问题标题】:Produce balanced mini batch with Dataset API使用 Dataset API 生成平衡的小批量
【发布时间】:2018-04-06 21:41:51
【问题描述】:

我有一个关于新数据集 API (tensorflow 1.4rc1) 的问题。 我对标签01 有一个不平衡的数据集。我的目标是在预处理期间创建平衡的小批量。

假设我有两个过滤数据集:

ds_pos = dataset.filter(lambda l, x, y, z: tf.reshape(tf.equal(l, 1), []))
ds_neg = dataset.filter(lambda l, x, y, z: tf.reshape(tf.equal(l, 0), [])).repeat()

有没有办法将这两个数据集组合起来,使生成的数据集看起来像ds = [0, 1, 0, 1, 0, 1]

类似这样的:

dataset = tf.data.Dataset.zip((ds_pos, ds_neg))
dataset = dataset.apply(...)
# dataset looks like [0, 1, 0, 1, 0, 1, ...]
dataset = dataset.batch(20)

我目前的做法是:

def _concat(x, y):
   return tf.cond(tf.random_uniform(()) > 0.5, lambda: x, lambda: y)
dataset = tf.data.Dataset.zip((ds_pos, ds_neg))
dataset = dataset.map(_concat)

但我觉得有一种更优雅的方式。

提前致谢!

【问题讨论】:

标签: tensorflow tensorflow-datasets


【解决方案1】:

你在正确的轨道上。下面的例子使用Dataset.flat_map()将每对正例和负例在结果中变成两个连续的例子:

dataset = tf.data.Dataset.zip((ds_pos, ds_neg))

# Each input element will be converted into a two-element `Dataset` using
# `Dataset.from_tensors()` and `Dataset.concatenate()`, then `Dataset.flat_map()`
# will flatten the resulting `Dataset`s into a single `Dataset`.
dataset = dataset.flat_map(
    lambda ex_pos, ex_neg: tf.data.Dataset.from_tensors(ex_pos).concatenate(
        tf.data.Dataset.from_tensors(ex_neg)))

dataset = dataset.batch(20)

【讨论】:

  • 非常感谢!现在我可以删除每个数据点的采样步骤 :)
  • 我想对多类分类问题使用相同的方法。然而,它的速度非常缓慢。有没有更有效的方法在非二元分类任务中产生平衡的小批量?
  • @aseipel 一种选择是将您的示例分成每个类一个文件,然后使用Dataset.range(NUM_CLASSES).interleave(dataset_for_class, cycle_length=NUM_CLASSES)(其中dataset_for_class(c) 是一个为类c 加载该文件的函数)。跨度>
  • @aseipel -- 你知道在多类情况下哪个部分慢吗?通常重新初始化迭代器非常慢,所以我认为zipflat_map 解决方案比为每个类都有一个迭代器要好。
  • 关于此解决方案要记住的一点是zip 将创建一个大小等于正在压缩的最小数据集的数据集。我认为这就是 OP 在ds_neg 上使用repeat 的原因。我认为这是为了确保使用所有多数类的数据ds_pos
猜你喜欢
  • 1970-01-01
  • 2018-11-09
  • 1970-01-01
  • 1970-01-01
  • 2019-02-14
  • 2019-06-13
  • 1970-01-01
  • 2020-01-24
  • 1970-01-01
相关资源
最近更新 更多