【问题标题】:Keras get model outputs after each batchKeras 在每批后获取模型输出
【发布时间】:2019-03-08 03:16:49
【问题描述】:

我正在使用生成器为分层循环模型生成顺序训练数据,该模型需要前一批的输出来生成下一批的输入。这与 Keras 参数 stateful=True 的情况类似,它为下一批保存隐藏状态,但它更复杂,所以我不能按原样使用。

到目前为止,我尝试在损失函数中添加 hack:

def custom_loss(y_true, y_pred):
    global output_ref
    output_ref[0] = y_pred[0].eval(session=K.get_session())
    output_ref[1] = y_pred[1].eval(session=K.get_session())

但这没有编译,我希望有更好的方法。 Keras 回调会有帮助吗?

【问题讨论】:

    标签: callback keras


    【解决方案1】:

    向here学习:

    model.compile(optimizer='adam')
    # hack after compile
    output_layers = [ 'gru' ]
    s_name = 's'
    model.metrics_names += [s_name]
    model.metrics_tensors += [layer.output for layer in model.layers if layer.name in output_layers]
    
    class my_callback(Callback):
        def on_batch_end(self, batch, logs=None):
            s_pred = logs[s_name]
            print('s_pred:', s_pred)
            return
    
    model.fit(..., callbacks=[my_callback()])
    

    【讨论】:

    • s_name是什么意思,我试试你的方法,遇到AttributeError: Model object has no attribute 'metric_names'的错误。
    【解决方案2】:

    我在 Keras 的 Tensorflow 版本中使用它,但它应该在没有 Tensorflow 的 Keras 中工作

    import tensorflow as tf
    
    class ModelOutput:
        ''' Class wrapper for a metric that stores the output passed to it '''
        def __init__(self, name):
            self.name = name
            self.y_true = None
            self.y_pred = None
    
        def save_output(self, y_true, y_pred):
            self.y_true = y_true
            self.y_pred = y_pred
            return tf.constant(True)
    
    class ModelOutputCallback(tf.keras.callbacks.Callback):
      def __init__(self, model_outputs):
        tf.keras.callbacks.Callback.__init__(self)
        self.model_outputs = model_outputs
    
      def on_train_batch_end(self, batch, logs=None):
        #use self.model_outputs to get the outputs here
    
    model_outputs = [
                    ModelOutput('rbox_score_map'),
                    ModelOutput('rbox_shapes'),
                    ModelOutput('rbox_angles')
                ]
    
    # Note the extra [] around m.save_output, this example is for a model with 
    # 3 outputs, metrics must be a list of lists if you type it out
    model.compile( ..., metrics=[[m.save_output] for m in self.model_outputs])
    
    model.fit(..., callbacks=[ModelOutputCallback(model_outputs)])
    
    

    【讨论】:

      猜你喜欢
      • 2017-05-25
      • 1970-01-01
      • 1970-01-01
      • 1970-01-01
      • 2019-01-27
      • 2017-11-21
      • 1970-01-01
      • 2018-03-05
      • 1970-01-01
      相关资源
      最近更新 更多