【发布时间】:2019-04-26 09:00:36
【问题描述】:
我的模型在每个输入批次中使用按时间顺序排列的序列。因此,我在改组输入数据之前创建批次。这带来了一个问题,即批次总是在整个数据集中包含相同的数据样本(从相同的索引开始 - 移动 batch_size),我通过缓存初始数据集并从跳过的数据集中采样解决了这个问题,但这会占用内存相当快(虽然我的数据集只有 150MB):
dataset = tf.data.Dataset.from_tensor_slices(data)
dataset = dataset.window(size=window_size, shift=window_shift, stride=window_stride, drop_remainder=True).flat_map(lambda x: x.batch(window_size))
dataset = dataset.map(process_fn, num_parallel_calls=8)
dataset = dataset.cache()
datasets = []
for i in range(0, batch_size):
d = dataset.skip(i)
d = d.batch(batch_size, drop_remainder=True)
datasets.append(d)
dataset = tf.data.experimental.sample_from_datasets(datasets)
dataset = dataset.shuffle(buffer_size=30000, reshuffle_each_iteration=False)
dataset = dataset.repeat()
还有其他方法可以实现这种行为吗?我想涵盖批次内第一个序列开始的所有可能索引。
【问题讨论】:
-
你找到更好的方法来减少内存使用了吗?
-
遗憾的是,我还没有重构这部分代码。我不确定现在是否有更好的方法。
-
我看到你的设置
buffere_size为 3,000。你试过像3这样的小数字吗?还有什么dataset尺寸和形状? -
您应该提供更多类似
process_fn和batch_size的代码,以便我们重现问题。 -
batch_size是 128
标签: python tensorflow tensorflow-datasets tensorflow-estimator