【问题标题】:Training GAN in keras with .fit_generator()使用 .fit_generator() 在 keras 中训练 GAN
【发布时间】:2020-03-06 05:12:34
【问题描述】:

我一直在使用以下训练循环训练类似于 Pix2Pix 的条件 GAN 架构:

for epoch in range(start_epoch, end_epoch):
    for batch_i, (input_batch, target_batch) in enumerate(dataLoader.load_batch(batch_size)):
                fake_batch= self.generator.predict(input_batch)

                d_loss_real = self.discriminator.train_on_batch(target_batch, valid)
                d_loss_fake = self.discriminator.train_on_batch(fake_batch, invalid)
                d_loss = np.add(d_loss_fake, d_loss_real) * 0.5

                g_loss = self.combined.train_on_batch([target_batch, input_batch], [valid, target_batch])

现在这很好用,但效率不高,因为数据加载器很快就会成为时间瓶颈。 我研究了 keras 提供的 .fit_generator() 函数,它允许生成器在工作线程中运行并且运行得更快。

self.combined.fit_generator(generator=trainLoader,
                                    validation_data=evalLoader
                                    callbacks=[checkpointCallback, historyCallback],
                                    workers=1,
                                    use_multiprocessing=True)

我花了一些时间才发现这是不正确的,我不再单独训练我的生成器和鉴别器,并且鉴别器根本没有被训练,因为它在组合模型中设置为 trainable = False,本质上破坏任何形式的对抗性损失,我还不如用 MSE 自己训练我的生成器。

现在我的问题是,是否有一些解决方法,例如在自定义回调中训练我的鉴别器,每批 .fit_generator() 方法都会触发该回调?可以实现创建自定义回调,例如:

class MyCustomCallback(tf.keras.callbacks.Callback):
  def on_train_batch_end(self, batch, logs=None):
    discriminator.train_on_batch()

另一种可能性是将原始训练循环并行化,但恐怕我现在没有时间这样做。

【问题讨论】:

    标签: python keras deep-learning generative-adversarial-network


    【解决方案1】:

    更新:为此内置了队列:

    您可以在此答案中查看使用它们的快速方法:https://stackoverflow.com/a/59214794/2097240


    旧答案:

    我正是为此目的创建了这个并行化迭代器。我在训练中使用它;

    这是你使用它的方式:

    for epoch, batchIndex, originalBatchIndex, xAndY in ParallelIterator(
                                           generator, 
                                           epochs, 
                                           shuffle_bool, 
                                           use_on_epoch_end_from_generator_bool,
                                           workers = 8, 
                                           queue_size=10):
        #loop content
        x_train_batch, y_train_batch = xAndY
        model.train_on_batch(x_train_batch, y_train_batch)
    
    
    

    generator 应该是你的dataloader,但它必须是keras.utils.Sequence,而不仅仅是一个产量生成器。

    但是如果你需要的话,适应起来并不是很复杂。 (我只是不知道它是否会正确并行化,但我不知道是否可以正确并行化 yield 循环)
    在下面的迭代器定义中,您应该替换:

    • len(keras_sequence)steps_per_epoch
    • keras_sequence[i]next(keras_sequence)
    • use_on_epoch_end = False

    这是迭代器的定义:

    
    import multiprocessing.dummy as mp
    
    #A generator that wraps a Keras Sequence and simulates a `fit_generator` behavior for custom training loops
    #It will also work with any iterator that has `__len__` and `__getitem__`.    
    def ParallelIterator(keras_sequence, epochs, shuffle, use_on_epoch_end, workers = 4, queue_size = 10):
    
        sourceQueue = mp.Queue()                     #queue for getting batch indices
        batchQueue = mp.Queue(maxsize = queue_size)  #queue for getting actual batches 
        indices = np.arange(len(keras_sequence))     #array of indices to be shuffled
    
        use_on_epoch_end = 'on_epoch_end' in dir(keras_sequence) if use_on_epoch_end == True else False
        batchesLeft = 0
    
    #     printQueue = mp.Queue()                      #queue for printing messages
    #     import threading
    #     screenLock = threading.Semaphore(value=1)
    #     totalWorkers= 0
    
    #     def printer():
    #         nonlocal printQueue, printing
    #         while printing:
    #             while not printQueue.empty():
    #                 text = printQueue.get(block=True)
    #                 screenLock.acquire()
    #                 print(text)
    #                 screenLock.release()
    
        #fills the batch indices queue (called when sourceQueue is empty -> a few batches before an epoch ends)
        def fillSource():
            nonlocal batchesLeft
    
    #         printQueue.put("Iterator: fill source - source qsize = " + str(sourceQueue.qsize()))
            if shuffle == True:
                np.random.shuffle(indices)
    
            #puts the indices in the indices queue
            batchesLeft += len(indices)
    #         printQueue.put("Iterator: batches left:" + str(batchesLeft))
            for i in indices:
                sourceQueue.put(i)
    
        #function that will load batches from the Keras Sequence
        def worker():
            nonlocal sourceQueue, batchQueue, keras_sequence, batchesLeft
    #         nonlocal printQueue, totalWorkers
    #         totalWorkers += 1
    #         thisWorker = totalWorkers
    
            while True:
    #             printQueue.put('Worker: ' + str(thisWorker) + ' will try to get item')
                index = sourceQueue.get(block = True) #get index from the queue
    #             printQueue.put('Worker: ' + str(thisWorker) + ' got item ' +  str(index) + " - source q size = " + str(sourceQueue.qsize()))
    
                if index is None:
                    break
    
                item = keras_sequence[index] #get batch from the sequence
                batchesLeft -= 1
    #             printQueue.put('Worker: ' + str(thisWorker) + ' batches left ' + str(batchesLeft))
    
                batchQueue.put((index,item), block=True) #puts batch in the batch queue
    #             printQueue.put('Worker: ' + str(thisWorker) + ' added item ' + str(index) + ' - queue: ' + str(batchQueue.qsize()))
    
    #         printQueue.put("hitting end of worker" + str(thisWorker))
    
    #       #printing pool that will print messages from the print queue
    #     printing = True
    #     printPool = mp.Pool(1, printer)
    
        #creates the thread pool that will work automatically as we get from the batch queue
        pool = mp.Pool(workers, worker)    
        fillSource()   #at this point, data starts being taken and stored in the batchQueue
    
        #generation loop
        for epoch in range(epochs):
    
            #if not waiting for epoch end synchronization, always keeps 1 epoch filled ahead
            if (use_on_epoch_end == False):
                if epoch + 1 < epochs: #only fill if not last epoch
                    fillSource()
    
            for batch in range(len(keras_sequence)):
    
                #if waiting for epoch end synchronization, wait for workers to have no batches left to get, then call epoch end and fill
                if use_on_epoch_end == True:
                    if batchesLeft == 0:
                        keras_sequence.on_epoch_end()
                        if epoch + 1 < epochs:  #only fill if not last epoch
                            fillSource()
                        else:
                            batchesLeft = -1   #in the last epoch, prevents from calling epoch end again and again
    
                #yields batches for the outside loop that is using this generator
                originalIndex, batchItems = batchQueue.get(block = True)
                yield epoch, batch, originalIndex, batchItems
    
    
    #         print("iterator epoch end")
    #     printQueue.put("closing threads")
    
        #terminating the pool - add None to the queue so any blocked worker gets released
        for i in range(workers):
            sourceQueue.put(None)
        pool.terminate()
        pool.close()
        pool.join()
    #     printQueue.put("terminated")
    
    #     printing = False
    #     printPool.terminate()
    #     printPool.close()
    #     printPool.join()
    
    
        del pool,sourceQueue,batchQueue
    #     del printPool, printQueue
    

    【讨论】:

    • 我还没有时间尝试你的答案,但不妨在赏金到期之前奖励它。无论如何感谢您的回复!
    • 谢谢 :D -- 它确实有效,我现在在我的代码中使用它:))
    • 我终于有时间来看看你的迭代器,它超级聪明。谢谢!您的回答绝对值得 50 次代表:D
    • @AhmadMoussa ,我发现了两个用更少的代码完成这项工作的队列 :),更新了我的答案。
    • 您好,感谢您的更新!你说的太好了,如果可以的话,我会再次投票给你!
    【解决方案2】:

    虽然您的问题已经有了解决方案,但如果您可以在组合模型中的自定义回调中训练您的鉴别器,我想回答您的原始问题。

    简单的答案是是的

    编译模型(判别器和组合模型)时要小心,并按照此处所述的步骤操作: https://github.com/keras-team/keras/issues/8585#issuecomment-385729276

    调用您的组合模型拟合或拟合生成器:

    combined_model.fit_generator(train_loader, epochs, callbacks=[gan_callback])
    

    gan_callback 是一个自定义回调类,覆盖您调用的 on_batch_end(如您所说)

    def on_batch_end(self, batch_idx, logs=None):
        logs_disc = model_disc.train_on_batch(x, y)
    

    要在回调中获取鉴别器模型,可以在构造时将其作为参数提供,也可以通过继承的 self.model (model.layers) 变量获取。

    当你想将损失和指标输出到张量板时,我认为这个解决方案很优雅。

    在 gan_callback 的 on_batch_end 函数中,您可以直接获得两个日志(包含损失和指标的值):

    • 来自鉴别器的logs_disc
    • 来自生成器的日志,是 on_batch_end() 的参数

    根据您的配置,这可能会产生一个可以忽略的警告:

    UserWarning: Method on_batch_end() is slow compared to the batch update (0.151899).    Check your callbacks.
    

    【讨论】:

      猜你喜欢
      • 2018-12-06
      • 1970-01-01
      • 2020-03-20
      • 2017-09-18
      • 2017-06-20
      • 1970-01-01
      • 1970-01-01
      • 1970-01-01
      • 1970-01-01
      相关资源
      最近更新 更多