【问题标题】:Keras fit_generator issueKeras fit_generator 问题
【发布时间】:2019-04-25 16:17:41
【问题描述】:

我关注 this tutorial 为我的 Keras 模型创建了一个自定义生成器。这是一个 MWE,显示了我面临的问题:

import sys, keras
import numpy as np
import tensorflow as tf
import pandas as pd
from keras.models import Model
from keras.layers import Dense, Input
from keras.optimizers import Adam
from keras.losses import binary_crossentropy

class DataGenerator(keras.utils.Sequence):
    'Generates data for Keras'
    def __init__(self, list_IDs, batch_size, shuffle=False):
        'Initialization'
        self.batch_size = batch_size
        self.list_IDs = list_IDs
        self.shuffle = shuffle
        self.on_epoch_end()

    def __len__(self):
        'Denotes the number of batches per epoch'
        return int(np.floor(len(self.list_IDs) / self.batch_size))

    def __getitem__(self, index):
        'Generate one batch of data'
        # Generate indexes of the batch
        #print('self.batch_size: ', self.batch_size)
        print('index: ', index)
        sys.exit()

    def on_epoch_end(self):
        'Updates indexes after each epoch'
        self.indexes = np.arange(len(self.list_IDs))
        print('self.indexes: ', self.indexes)
        if self.shuffle == True:
            np.random.shuffle(self.indexes)

    def __data_generation(self, list_IDs_temp):
        'Generates data containing batch_size samples' # X : (n_samples, *dim, n_channels)

        X1 = np.empty((self.batch_size, 10), dtype=float)
        X2 = np.empty((self.batch_size, 12),  dtype=int)

        #Generate data
        for i, ID in enumerate(list_IDs_temp):
            print('i is: ', i, 'ID is: ', ID)

            #Preprocess this sample (omitted)
            X1[i,] = np.repeat(1, X1.shape[1])
            X2[i,] = np.repeat(2, X2.shape[1])

        Y = X1[:,:-1]
        return X1, X2, Y

if __name__=='__main__':
    train_ids_to_use = list(np.arange(1, 321)) #1, 2, ...,320 
    valid_ids_to_use = list(np.arange(321, 481)) #321, 322, ..., 480

    params = {'batch_size': 32}

    train_generator = DataGenerator(train_ids_to_use, **params)
    valid_generator = DataGenerator(valid_ids_to_use, **params)

    #Build a toy model
    input_1 = Input(shape=(3, 10))
    input_2 = Input(shape=(3, 12))
    y_input = Input(shape=(3, 10))

    concat_1 = keras.layers.concatenate([input_1, input_2])
    concat_2 = keras.layers.concatenate([concat_1, y_input])

    dense_1 = Dense(10, activation='relu')(concat_2)
    output_1 = Dense(10, activation='sigmoid')(dense_1)

    model = Model([input_1, input_2, y_input], output_1)
    print(model.summary())

    #Compile and fit_generator
    model.compile(optimizer=Adam(lr=0.001), loss=binary_crossentropy)
    model.fit_generator(generator=train_generator, validation_data = valid_generator, epochs=2, verbose=2)

我不想打乱我的输入数据。我认为这已经得到处理,但是在我的代码中,当我在__get_item__ 中打印出index 时,我得到了随机数。我想要连续的数字。请注意,我正在尝试在 __getitem__ 中使用 sys.exit 来终止进程,以查看发生了什么。

我的问题:

  1. 为什么index 不连续?我该如何解决这个问题?

  2. 当我在终端使用屏幕运行时,为什么它没有响应 Ctrl+C?

【问题讨论】:

  • 我认为您可以通过将shuffle=False 传递给fit_generator 方法来实现?
  • 您好,感谢您的回复。我在__init__ 中将其作为默认设置,然后我测试了索引值是否被on_epoch_end 中的if 语句打乱了。我发现if语句中的东西没有被执行,我认为这意味着shuffle确实是假的。
  • 您希望批量索引连续生成,对吗?这就是fit_generatorshuffle=False 参数。你试过了吗?
  • 是的,请看上面的评论。
  • 对不起,我不明白。在我的机器上,当我在fit_generator 调用中设置shuffle=False(不是在__init__ 方法中)时,我会得到连续的索引。

标签: python keras generator


【解决方案1】:

您可以使用fit_generator 方法的shuffle 参数连续生成批次。来自fit_generator()documentation

随机播放:布尔值。是否在每个 epoch 开始时打乱批次的顺序。仅用于 Sequence (keras.utils.Sequence) 的实例。当steps_per_epoch 不是None 时无效。

只需将shuffle=False 传递给fit_generator

model.fit_generator(generator=train_generator, shuffle=False, ...)

【讨论】:

  • 好吧,请让我检查一下我的理解:fit_generator 没有被覆盖,它属于 Keras,并且有自己的参数。我正在使用DataGenerator 创建自己的生成器(train_generator 和 valid_generator)来创建数据切片供 Keras 使用。但是为什么这意味着我需要在对fit_generator 的调用中指定shuffle=False,而不是在我自己的生成器中?
  • @StatsSorceress fit_generator 已经实现,它可以使用给定的索引值调用Sequence__getitem__ 方法。因此,fit_generator 为您提供了要生成的批次的索引。将 shuffle=False 传递给 fit_generator 强制它按顺序给出批次索引,即 0、1、2、3、...
猜你喜欢
  • 2017-11-17
  • 2020-08-01
  • 2017-06-07
  • 1970-01-01
  • 1970-01-01
  • 2017-12-30
  • 2019-10-26
  • 2017-09-13
  • 2018-10-07
相关资源
最近更新 更多