【问题标题】:Tensorflow: how to make sure all samples in each batch are with the same label?Tensorflow:如何确保每批中的所有样本都具有相同的标签?
【发布时间】:2018-07-21 11:47:23
【问题描述】:

我想知道是否有一些方法可以对要在 Tensorflow 中生成的批次应用约束。例如,假设我们正在一个巨大的数据集上训练一个 CNN 来进行图像分类。是否可以强制 Tensorflow 生成所有样本都属于同一类的批次?比如,一组图像都标记为“Apple”,另一组样本都标记为“Orange”。

我问这个问题的原因是我想做一些实验,看看不同级别的洗牌如何影响最终训练的模型。为 CNN 训练进行样本级洗牌是一种常见的做法,每个人都在这样做。我只是想自己去查一下,从而获得更生动的第一手资料。

谢谢!

【问题讨论】:

    标签: tensorflow batch-processing


    【解决方案1】:

    Dataset.filter()可以用:

    labels = np.random.randint(0, 10, (10000))
    data = np.random.uniform(size=(10000, 5))
    
    ds = tf.data.Dataset.from_tensor_slices((data, labels))
    ds = ds.filter(lambda data, labels: tf.equal(labels, 1)) #comment this line out for unfiltered case
    ds = ds.batch(5)
    iterator = ds.make_one_shot_iterator()
    vals = iterator.get_next()
    
    with tf.Session() as sess:
        for _ in range(5):
            py_data, py_labels = sess.run(vals)
            print(py_labels)
    

    ds.filter():

     > [1 1 1 1 1]
       [1 1 1 1 1]
       [1 1 1 1 1]
       [1 1 1 1 1]
       [1 1 1 1 1]
    

    没有ds.filter():

      > [8 0 7 6 3]
        [2 4 7 6 1]
        [1 8 5 5 5]
        [7 1 7 4 0]
        [7 1 8 0 0]
    

    编辑。以下代码显示了如何使用可馈送迭代器即时执行批量标签选择。见“Creating an iterator

    labels = ['Apple'] * 100 + ['Orange'] * 100
    data = list(range(200))
    random.shuffle(labels)
    
    batch_size = 4
    
    ds_apple = tf.data.Dataset.from_tensor_slices((data, labels)).filter(
      lambda data, label: tf.equal(label, 'Apple')).batch(batch_size)
    ds_orange = tf.data.Dataset.from_tensor_slices((data, labels)).filter(
      lambda data, label: tf.equal(label, 'Orange')).batch(batch_size)
    
    handle = tf.placeholder(tf.string, [])
    iterator = tf.data.Iterator.from_string_handle(
      handle, ds_apple.output_types, ds_apple.output_shapes)
    batch = iterator.get_next()
    
    apple_iterator = ds_apple.make_one_shot_iterator()
    orange_iterator = ds_orange.make_one_shot_iterator()
    
    with tf.Session() as sess:
      apple_handle = sess.run(apple_iterator.string_handle())
      orange_handle = sess.run(orange_iterator.string_handle())
    
      # loop and switch back and forth between apples and oranges
      for _ in range(3):
        feed_dict = {handle: apple_handle}
        print(sess.run(batch, feed_dict=feed_dict))
        feed_dict = {handle: orange_handle}
        print(sess.run(batch, feed_dict=feed_dict))
    

    典型的输出如下。请注意,data 值在 Apple 和 Orange 批次中单调增加,表明迭代器没有重置。

    > (array([2, 3, 6, 7], dtype=int32), array([b'Apple', b'Apple', b'Apple', b'Apple'], dtype=object))
      (array([0, 1, 4, 5], dtype=int32), array([b'Orange', b'Orange', b'Orange', b'Orange'], dtype=object))
      (array([ 9, 13, 15, 19], dtype=int32), array([b'Apple', b'Apple', b'Apple', b'Apple'], dtype=object))
      (array([ 8, 10, 11, 12], dtype=int32), array([b'Orange', b'Orange', b'Orange', b'Orange'], dtype=object))
      (array([21, 22, 23, 25], dtype=int32), array([b'Apple', b'Apple', b'Apple', b'Apple'], dtype=object))
      (array([14, 16, 17, 18], dtype=int32), array([b'Orange', b'Orange', b'Orange', b'Orange'], dtype=object))
    

    【讨论】:

    • 在上面显示的示例中,过滤器应用于每个样本。我想知道过滤器可以应用于每批吗?我需要每批中的样品都带有相同的标签(全部带有 1、全部带有 2、全部带有 3 等),不一定只有 1。如果没有一个批次包含两个带有不同标签的样品,那没关系。总之:我想过滤包含具有不同标签的样本的批次。或者我想强制 Tensorflow 生成遵守此规则的批次。
    • 不确定我是否完全理解。上面的条件语句可以更改为tf.equal(labels, target),其中target 可以在其他地方设置为您想要的任何标签。如果过滤器应用于批次,那么批次大小似乎会在步骤之间发生变化。
    • 只是为了让自己完全理解:像这样[Apple,Apple,Apple,Apple]或像这样[Orange,Orange,Orange,Orange]这样的批次(4个样本)都可以。我只希望每批中所有样本的标签都相同,对标签没有额外的限制或要求。所以你建议我在训练期间动态设置“目标”以选择“Apple”或“Orange”?
    • 好的,我编辑了我的答案以展示如何使用可馈送迭代器动态过滤批次。我以前从未使用过,所以这对我来说是一次很好的体验。希望它能满足您的需求。
    • 谢谢!这正是我想要的。
    猜你喜欢
    • 1970-01-01
    • 2020-10-23
    • 1970-01-01
    • 1970-01-01
    • 2016-01-17
    • 1970-01-01
    • 2019-10-17
    • 2021-09-16
    • 2013-09-10
    相关资源
    最近更新 更多