【问题标题】:How to apply DataGenerator to train and validation data?如何应用 DataGenerator 来训练和验证数据?
【发布时间】:2020-06-18 12:50:21
【问题描述】:

使用这里的代码https://keras.io/api/utils/python_utils/#sequence-class,我编写了一个自定义的DataGenerator

 # Here, `x_set` is list of path to the images
 # and `y_set` are the associated classes.

class DataGenerator(Sequence):

    def __init__(self, x_set, y_set, batch_size):
        self.x, self.y = x_set, y_set
        self.batch_size = batch_size

    def __len__(self):
        return math.ceil(len(self.x) / self.batch_size)

    def __getitem__(self, idx):
        batch_x = self.x[idx * self.batch_size:(idx + 1) *
        self.batch_size]
        batch_y = self.y[idx * self.batch_size:(idx + 1) *
        self.batch_size]

        return np.array([
            resize(imread(file_name), (224, 224))
               for file_name in batch_x]), np.array(batch_y)

现在,我想知道如何将数据生成器应用于我的训练数据和验证数据? 我有X_trainX_val,它们是包含我的图像文件的图像路径的列表以及y_trainy_val,它们是一个热门编码标签。

然后我可以使用此代码吗?

training_generator = DataGenerator(X_train, y_train)
validation_generator = DataGenerator(X_val, y_val)

然后拟合模型?

model.fit_generator(generator=training_generator,
                    validation_data=validation_generator)

【问题讨论】:

  • 您忘记将“batch_size”参数传递给DataGenerator 类初始化方法。如果你愿意,你可以在方法声明中设置一个默认值。

标签: python data-generation


【解决方案1】:

你写的基本上是正确的。不要忘记将batch_size 参数传递给您的DataGenerator

另一方面,epochs 参数(正如您在评论中提到的)应该传递给model.fit_generator(更好的是,使用model.fit 代替,因为fit_generator 方法是deprecated)。如果您不传递它,epochs 的默认值将是 1。

另外请查看this tutorial,了解如何使用Sequence 类(您可以跳到使用DataGenerator 的底部)。在本教程中,除了batch_size 之外的几个参数被传递给DataGenerator,因为它们被定义为__init__ 方法的输入。只要不定义它们就不必传递它们。

【讨论】:

猜你喜欢
  • 2021-01-05
  • 2022-07-31
  • 2019-12-26
  • 2019-05-01
  • 2020-09-14
  • 1970-01-01
  • 2020-05-11
  • 2021-07-21
  • 2018-03-09
相关资源
最近更新 更多