【问题标题】:how to shuffle a Concatenated Tensorflow dataset如何打乱连接的 Tensorflow 数据集
【发布时间】:2019-01-16 19:02:42
【问题描述】:

我有多个结构相同的 tensorflow 数据集。 我想将它们组合成一个数据集。使用 tf.dataset.concatenate

但我发现,在对这个组合数据集进行改组时,数据集并未在整个数据集的范围内进行改组。但在每个分离的数据集中打乱。

有什么方法可以解决这个问题吗?

【问题讨论】:

    标签: tensorflow dataset


    【解决方案1】:

    从 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')
    

    【讨论】:

      【解决方案2】:

      您的随机播放缓冲区大小是多少?

      例如,如果您有 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)
      

      【讨论】:

        【解决方案3】:

        当你连接两个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,您可以像这样依赖 zipflat_mapconcatenate 的组合:

        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 解决问题?
        • 背景是:1.输入数据集是词序列。 2.我使用这个方案的原因是我应该根据输入句子的第一个单词来阅读不同的词汇文件。 3.现在我可以得到一个python列表,其中每一项都是一个tf.dataset。
        【解决方案4】:

        不是 100% 确定,但您可能想查看在数据集对象上调用不同操作的顺序。 shuffle() 的行为可能因顺序而异。另请参阅this 可能相关的问题。

        【讨论】:

          猜你喜欢
          • 1970-01-01
          • 1970-01-01
          • 2020-01-19
          • 2021-07-06
          • 2020-02-19
          • 2019-12-05
          • 2021-08-12
          • 1970-01-01
          • 1970-01-01
          相关资源
          最近更新 更多