【发布时间】:2019-01-16 19:02:42
【问题描述】:
我有多个结构相同的 tensorflow 数据集。 我想将它们组合成一个数据集。使用 tf.dataset.concatenate
但我发现,在对这个组合数据集进行改组时,数据集并未在整个数据集的范围内进行改组。但在每个分离的数据集中打乱。
有什么方法可以解决这个问题吗?
【问题讨论】:
标签: tensorflow dataset
我有多个结构相同的 tensorflow 数据集。 我想将它们组合成一个数据集。使用 tf.dataset.concatenate
但我发现,在对这个组合数据集进行改组时,数据集并未在整个数据集的范围内进行改组。但在每个分离的数据集中打乱。
有什么方法可以解决这个问题吗?
【问题讨论】:
标签: tensorflow dataset
从 tensorflow 1.9 开始,您还可以使用sample_from_datasets 方法。
例如下面的代码
datasets = [tf.data.Dataset.from_tensors("foo").repeat(3).apply(tf.data.experimental.enumerate_dataset()).repeat(),
tf.data.Dataset.from_tensors("bar").repeat(3).apply(tf.data.experimental.enumerate_dataset()).repeat(),
tf.data.Dataset.from_tensors("baz").repeat(3).apply(tf.data.experimental.enumerate_dataset()).repeat()]
dataset = tf.data.experimental.sample_from_datasets(datasets) # from 1.12
# dataset = tf.contrib.data.sample_from_datasets(datasets) # between 1.9 and 1.12
iterator = dataset.make_one_shot_iterator();next_element = iterator.get_next()
with tf.Session() as sess:
for i in range(10):
print(sess.run(next_element))
将打印
(0, b'bar')
(0, b'foo')
(1, b'bar')
(0, b'baz')
(2, b'bar')
(1, b'foo')
(1, b'baz')
(2, b'foo')
(2, b'baz')
(0, b'foo')
【讨论】:
您的随机播放缓冲区大小是多少?
例如,如果您有 3 个数据集,每个数据集包含 1000 个项目,那么您需要应用 shuffle(3000) 来随机化所有项目的顺序。
这是一个例子:
这应该洗牌所有 3000 个项目:
dataset = dataset1.concatenate(dataset2).concatenate(dataset3)
dataset = dataset.shuffle(3000)
但是,这不会打乱整个数据集:
dataset1 = dataset1.shuffle(1000)
dataset2 = dataset2.shuffle(1000)
dataset3 = dataset3.shuffle(1000)
dataset = dataset1.concatenate(dataset2).concatenate(dataset3)
【讨论】:
当你连接两个Datasets 时,你会得到第一个的元素,然后是第二个的元素。如果您对结果进行混洗,如果您的混洗缓冲区小于Dataset 的大小,您将无法获得良好的混音。
您需要的是从数据集中交错样本。如果您使用 TF >= 1.9,最好的方法是使用专用的tf.contrib.data.choose_from_datasets 函数。直接来自文档的示例:
datasets = [tf.data.Dataset.from_tensors("foo").repeat(),
tf.data.Dataset.from_tensors("bar").repeat(),
tf.data.Dataset.from_tensors("baz").repeat()]
# Define a dataset containing `[0, 1, 2, 0, 1, 2, 0, 1, 2]`.
choice_dataset = tf.data.Dataset.range(3).repeat(3)
result = tf.contrib.data.choose_from_datasets(datasets, choice_dataset)
如果在批次中保留样本顺序和/或它们的比率很重要,最好对输入数据集进行洗牌。
如果您使用的是早期版本的 TF,您可以像这样依赖 zip、flat_map 和 concatenate 的组合:
a = tf.data.Dataset.range(3).repeat()
b = tf.data.Dataset.range(100, 105).repeat()
value = (tf.data.Dataset
.zip((a, b))
.flat_map(lambda x, y: tf.data.Dataset.concatenate(
tf.data.Dataset.from_tensors([x]),
tf.data.Dataset.from_tensors([y])))
.make_one_shot_iterator()
.get_next())
sess = tf.InteractiveSession()
for _ in range(10):
print(value.eval())
【讨论】:
tf 1.4 解决问题?
不是 100% 确定,但您可能想查看在数据集对象上调用不同操作的顺序。 shuffle() 的行为可能因顺序而异。另请参阅this 可能相关的问题。
【讨论】: