【问题标题】:How to pad to fixed BATCH_SIZE in tf.data.Dataset?如何在 tf.data.Dataset 中填充到固定的 BATCH_SIZE?
【发布时间】:2018-01-18 16:11:35
【问题描述】:

我有一个包含 11 个样本的数据集。而当我选择BATCH_SIZE为2时,下面的代码会报错:

dataset = tf.contrib.data.TFRecordDataset(filenames) 
dataset = dataset.map(parser)
if shuffle:
    dataset = dataset.shuffle(buffer_size=128)
dataset = dataset.batch(batch_size)
dataset = dataset.repeat(count=1)

问题出在dataset = dataset.batch(batch_size),当Dataset循环到最后一批时,剩余的样本数只有1个,那么有没有办法从之前访问的样本中随机抽取一个,生成最后一批?

【问题讨论】:

  • 通过填充filenames解决。
  • 请问如何填充文件名?我试过tf.data.TFRecordDataset(filename).repeat() 还是有这个问题。
  • 嗨@Richard_wth,我只是复制文件名,让文件名的数量是批量大小的倍数。比如11个文件名到12个文件名,那么问题就解决了。
  • 非常感谢您的回复。但它会改变训练数据的分布吗,因为一些数据被复制了?如果是这样,将文件名列表重复kk = batch_size / GCD(num_instances, batch_size) 会更好吗? GCD:最大公约数。
  • @Richard_wth,是的,谢谢你的建议!你说得对,我们可以考虑你建议的平衡数据分布的方法。

标签: tensorflow tensorflow-datasets


【解决方案1】:

@mining 通过填充文件名提出了一个解决方案。

另一种解决方案是使用tf.contrib.data.batch_and_drop_remainder。这将以固定的批量大小对数据进行批量处理,并丢弃最后一个较小的批量。

在您的示例中,有 11 个输入和 2 个批量大小,这将产生 5 个批次,每批 2 个元素。

这是文档中的示例:

dataset = tf.data.Dataset.range(11)
batched = dataset.apply(tf.contrib.data.batch_and_drop_remainder(2))

【讨论】:

【解决方案2】:

您可以在对batch 的调用中设置drop_remainder=True

dataset = dataset.batch(batch_size, drop_remainder=True)

来自documentation

drop_remainder:(可选)一个 tf.bool 标量 tf.Tensor,表示 最后一批是否应该在其数量较少的情况下被丢弃 比 batch_size 元素;默认行为是不删除 小批量。

【讨论】:

    猜你喜欢
    • 1970-01-01
    • 1970-01-01
    • 1970-01-01
    • 1970-01-01
    • 2012-04-02
    • 2023-04-07
    • 1970-01-01
    • 1970-01-01
    • 2013-12-17
    相关资源
    最近更新 更多