【发布时间】: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个文件名,那么问题就解决了。
-
非常感谢您的回复。但它会改变训练数据的分布吗,因为一些数据被复制了?如果是这样,将文件名列表重复
k次k = batch_size / GCD(num_instances, batch_size)会更好吗? GCD:最大公约数。 -
@Richard_wth,是的,谢谢你的建议!你说得对,我们可以考虑你建议的平衡数据分布的方法。
标签: tensorflow tensorflow-datasets