【问题标题】:Get Keras model input from inside a custom callback从自定义回调中获取 Keras 模型输入
【发布时间】:2019-03-17 22:28:11
【问题描述】:

我有一个非常简单的问题。我有一个为分类定义的 Keras 模型(TF 后端)。我想在训练期间转储输入到我的模型中的训练图像以进行调试。我正在尝试创建一个自定义回调,为此编写 Tensorboard 图像摘要。

但是如何在回调中获取真实的训练数据呢?

目前我正在尝试这个:

class TensorboardKeras(Callback):                                                                                                                                                                                                                                     
    def __init__(self, model, log_dir, write_graph=True):                                                                                                                                                                                                             
        self.model = model                                                                                                                                                                                                                                            
        self.log_dir = log_dir                                                                                                                                                                                                                                        
        self.session = K.get_session()                                                                                                                                                                                                                                

        tf.summary.image('input_image', self.model.input)                                                                                                                                                                                                             
        self.merged = tf.summary.merge_all()                                                                                                                                                                                                                          

        if write_graph:                                                                                                                                                                                                                                               
            self.writer = tf.summary.FileWriter(self.log_dir, K.get_session().graph)                                                                                                                                                                                  
        else:                                                                                                                                                                                                                                                         
            self.writer = tf.summary.FileWriter(self.log_dir)

    def on_batch_end(self, batch, logs=None):
        summary = self.session.run(self.merged, feed_dict={})                                                                                                                                                                                                         
        self.writer.add_summary(summary, batch)                                                                                                                                                                                                                       
        self.writer.flush()

但我收到错误消息: InvalidArgumentError(有关回溯,请参见上文):您必须为占位符张量“input_1”提供一个值,其 dtype 为 float 和 shape [?,224,224,3]

必须有一种方法可以查看作为输入的模型,对吧?

或者我应该尝试其他方式来调试它?

【问题讨论】:

  • 虽然您绝对应该能够做到这一点,但为什么不在输入点检查数据呢?例如。您的数据生成器。
  • 我正在使用直接输入 model.fit 方法的 tf.data.TFRecords 数据集。要直接检查数据,我必须编写一个逐批检索数据的包装代码。此外,它不会是训练的一部分,而是调试数据的辅助代码。或者,我可以使用回调并保持代码更简单。
  • 恐怕这是 Keras 中唯一的解决方案......从来没有在回调中看到任何与数据相关的东西,只有统计数据。

标签: python tensorflow keras


【解决方案1】:

您不需要为此进行回调。您需要做的就是实现一个生成图像及其标签作为元组的函数。 flow_from_directory 函数有一个名为 save_to_dir 的参数,它可以满足您的所有需求,如果没有,您可以这样做:

def trainGenerator(batch_size,train_path, image_size)
    #preprocessing see https://keras.io/preprocessing/image/ for details
    image_datagen = ImageDataGenerator(horizontal_flip=True)
    #create image generator see https://keras.io/preprocessing/image/#flow_from_directory for details
    train_generator = image_datagen.flow_from_directory(
        train_path,
        class_mode = "categorical",
        target_size = image_size,
        batch_size = batch_size,
        save_prefix  = "augmented_train",
        seed = seed)

    for (batch_imgs, batch_labels) in train_generator: 
        #do other stuff such as dumping images or further augmenting images
    yield (batch_imgs,batch_labels)


t_generator = trainGenerator(32, "./train_data", (224,224,3))
model.fit_generator(t_generator,steps_per_epoch=10,epochs=1)

【讨论】:

    猜你喜欢
    • 1970-01-01
    • 1970-01-01
    • 1970-01-01
    • 2020-04-12
    • 2021-03-20
    • 2011-10-24
    • 2020-09-10
    • 2020-04-30
    • 2017-11-21
    相关资源
    最近更新 更多