【问题标题】:Multiple independent iterators for a single dataset单个数据集的多个独立迭代器
【发布时间】:2018-09-22 19:57:34
【问题描述】:

假设我有一个训练数据集“data_train”,我想创建两个独立的迭代器,它们都迭代 data_train。我将使用第一个迭代器来训练我的网络“iter_train”,其中 iter_train.get_next() 的输出将是我训练的批次。第二个迭代器将用于在我训练时评估整个训练数据集,“iter_eval”,以监控训练进度。

目前,如果我只有一个迭代器“iter_single”,并且我想在一个时期的中途评估训练损失,我将不得不重置迭代器,使用 iter_single 评估整个数据集,然后从头开始训练带有 iter_single 的数据集。因此,我不会完成我之前的 epoch 并忽略一半的数据集,除非我浪费时间迭代数据而不对其进行操作。

我已经尝试为一个数据集创建两个迭代器,但是,通过重置一个迭代器会重置另一个迭代器,这使得拥有两个迭代器毫无意义。

【问题讨论】:

  • 请提供一个最小的代码示例来澄清您的确切问题

标签: python tensorflow tensorflow-datasets


【解决方案1】:

只是如果您的数据集大小不是很大(我所说的巨大是指在训练和评估期间可以将训练和验证数据都保存在内存中),您可以使用以下代码:

首先,读取和解析您的数据,并将它们传递给 Tensorflow 数据集对象:

def get_image_dataset(dir_path, batch_size, split=0.7):

    # Parse data and return them in array format (Numpy)
    train_data, val_data = parse_data(dir_path, split)

    # Create the dataset for our train data
    train_data = tf.data.Dataset.from_tensor_slices(train_data)
    train_data = train_data.batch(batch_size)

    # Create the dataset for our test data
    val_data = tf.data.Dataset.from_tensor_slices(val_data)
    val_data = val_data.batch(batch_size)

    return train_data, val_data

第二,为你的训练和验证数据定义一个迭代器和初始化器:

def get_data():
    with tf.name_scope('data'):

        train_data, test_data =  get_image_dataset(self.batch_size)
        iterator = tf.data.Iterator.from_structure(output_types=train_data.output_types, output_shapes=train_data.output_shapes)

        # Define one iterator for your data
        img, self.label = iterator.get_next()

        # Example of application on MNIST dataset
        img = tf.reshape(img, [-1, CNN_INPUT_HEIGHT, CNN_INPUT_WIDTH, CNN_INPUT_CHANNELS])

        # Define two initializers for either train or test (validation) data
        self.train_init = iterator.make_initializer(train_data)
        self.test_init = iterator.make_initializer(test_data)

第三第四,在训练/测试您的网络时,使用train/test初始化您的Tensorflow图像这样的数据集:

训练

def train_network_one_epoch(...):

    # Initialize training
    sess.run(self.train_init)

    # Run training graph nodes

    return something

测试

def evaluate_network(...):

    # Initialize testing
    sess.run(self.test_init)

    # Run evaluation graph nodes

    return something

您可以查看this 示例,该示例清楚地演示了此过程。

【讨论】:

  • 感谢 hexpheus,但就我而言,我希望训练和验证数据集相同。这是因为我想在训练时检查训练集的损失和准确性。因此 val_data = train_data。如果您进行这个简单的替换,那么运行 sess.run(self.test_init) 与运行 sess.run(self.train_init) 相同,至少从我自己进行的一项测试来看是这样。
猜你喜欢
  • 1970-01-01
  • 1970-01-01
  • 2019-03-25
  • 2015-10-19
  • 2022-06-24
  • 2021-01-25
  • 1970-01-01
  • 2020-03-22
  • 2022-01-02
相关资源
最近更新 更多