Keras中如何在model.predict预测阶段使用自定义回调函数
问题描述
我有一个如下所示的keras模型:
针对该模型,我编写了如下callback function(回调函数):
import tensorflow as tf import numpy as np class WriteLayerValCallback(tf.keras.callbacks.Callback): def __init__(self): self.data = np.random.rand(1,10) def on_epoch_end(self, epoch, logs=None): #dns_layer = self.model.layers[6] dns_layer = self.model.get_layer('activation') outputs = dns_layer(self.data) tf.print(f'\n input: {self.data}') tf.print(f'\n output: {outputs}')
当前我使用如下代码执行模型预测:
yhat = model.predict(X)
我希望在执行Keras prediction(Keras预测)流程时调用上述自定义回调函数,请问具体该如何实现?
实现方法
核心问题有两点:
- 你当前写的回调只实现了
on_epoch_end方法,这个钩子仅在fit训练阶段的轮次结束时触发,预测流程不会调用该方法 - Keras的
predict方法原生支持传入callbacks参数,只要给回调实现预测阶段对应的生命周期钩子,不需要修改核心预测逻辑就能触发
具体操作步骤
- 第一步:修改自定义回调类,把需要在预测阶段执行的逻辑,绑定到预测流程对应的钩子方法上。预测阶段可用的常用钩子如下:
on_predict_begin: 整个预测流程启动时触发on_predict_batch_begin/on_predict_batch_end: 每个预测batch开始/结束时触发on_predict_end: 整个预测流程全部完成时触发
- 第二步:调用
model.predict时,把实例化后的回调对象传入callbacks参数即可
可直接运行的修改示例
调整后的回调代码
如果需要在预测全部结束后打印目标激活层的输出,把逻辑迁移到on_predict_end方法中即可:
import tensorflow as tf import numpy as np class WriteLayerValCallback(tf.keras.callbacks.Callback): def __init__(self): self.data = np.random.rand(1,10) # 预测流程结束后自动执行该方法 def on_predict_end(self, logs=None): dns_layer = self.model.get_layer('activation') outputs = dns_layer(self.data) tf.print(f'\n input: {self.data}') tf.print(f'\n output: {outputs}')
预测时传入回调
# 实例化自定义回调 custom_cb = WriteLayerValCallback() # 传入callbacks参数即可正常触发 yhat = model.predict(X, callbacks=[custom_cb])
补充:如果需要在预测过程中逐batch获取层输出、统计中间值,把对应逻辑写到
on_predict_batch_end方法中即可,该钩子可以拿到当前batch索引、输入输出等上下文信息,用法和训练阶段的batch钩子一致。
内容的提问来源于stack exchange,提问作者S M Abrar Jahin
相关产品推荐
相关产品推荐

