Keras不使用回调实现自定义指标 如何获取model.evaluate的y_pred值
Keras获取model.evaluate生成的y_pred的实现方案
默认情况下model.evaluate不会对外暴露前向计算生成的y_pred,但可以通过自定义测试步的方式无额外开销拿到对应结果,不需要额外调用model.predict产生冗余计算。
方案1:重写test_step存储y_pred(适用于TensorFlow 2.x)
重写Keras模型的test_step方法,将前向计算得到的y_pred暂存在模型实例属性中,evaluate执行完成后直接读取即可,原有evaluate的计算逻辑、性能几乎不受影响。
示例代码如下:
import tensorflow as tf from tensorflow import keras class CustomModel(keras.Model): def test_step(self, data): x, y = data # 和evaluate默认逻辑一致,前向计算得到y_pred y_pred = self(x, training=False) # 暂存当前批次的y_pred if not hasattr(self, "eval_preds"): self.eval_preds = [] self.eval_preds.append(y_pred) # 保持原有损失、指标计算逻辑不变 self.compiled_loss(y, y_pred, regularization_losses=self.losses) self.compiled_metrics.update_state(y, y_pred) return {m.name: m.result() for m in self.metrics} # 复用原有模型的输入输出结构即可,不需要修改模型定义 # 示例:假设原有模型是Functional API定义的 inputs = keras.Input(shape=(28,)) dense1 = keras.layers.Dense(32, activation="relu")(inputs) outputs = keras.layers.Dense(10, activation="softmax")(dense1) # 用CustomModel封装原有输入输出 model = CustomModel(inputs, outputs) model.compile(optimizer="adam", loss="sparse_categorical_crossentropy", metrics=["accuracy"]) # 调用evaluate后直接读取y_pred model.evaluate(x_test, y_test, verbose=0) # 拼接所有批次的预测结果 y_pred_eval = tf.concat(model.eval_preds, axis=0) # 用完清空暂存值,避免下一次evaluate数据累积 model.eval_preds.clear()
方案2:结合回调在epoch结束自动获取y_pred
如果需要每个epoch训练结束后自动在验证集上执行evaluate并计算自定义指标,可以在回调中直接读取模型暂存的y_pred,不需要额外调用predict:
class CustomMetricCallback(keras.callbacks.Callback): def on_epoch_end(self, epoch, logs=None): logs = logs or {} # 执行验证集evaluate self.model.evaluate(x_val, y_val, verbose=0) # 直接读取evaluate生成的y_pred y_pred_val = tf.concat(self.model.eval_preds, axis=0) # 计算自定义指标 custom_score = your_custom_metric_func(y_val, y_pred_val) logs["val_custom_score"] = custom_score # 清空暂存 self.model.eval_preds.clear() # 训练时传入回调即可 model.fit( x_train, y_train, epochs=10, callbacks=[CustomMetricCallback()] )
注意事项
- 该方案完全复用了
evaluate过程中的前向计算结果,没有多余的性能开销,和原生evaluate的执行速度几乎一致 - 如果你使用的是Sequential模型,同样可以通过继承
keras.Sequential重写test_step的方式实现相同效果
内容的提问来源于stack exchange,提问作者Luka
相关产品推荐
相关产品推荐

