如何在Keras中实现可用于model.fit的自定义有状态指标
Keras自定义有状态指标在model.fit中不更新的解决方法
问题原因
官方示例本身没有问题,但它的演示场景默认传入的y_pred是已经完成二值化的0/1标签,和model.fit的实际使用场景不匹配:
你在代码中使用sigmoid作为输出层激活,模型输出的是0~1区间的连续概率值,原代码直接用tf.cast(y_pred, tf.bool)转换时,所有非0的概率值都会被判定为True,相当于所有样本都被预测为正例。此时计算得到的真阳性数值等于数据集内真实正例的总数,所以会出现逐epoch数值不变、和内置Recall结果冲突的问题。
内置的Recall指标默认内置了阈值处理逻辑,所以不需要额外调整就可以直接适配概率输出。
修复方案
修改BinaryTruePositives类的update_state方法,增加阈值判断逻辑即可:
def update_state(self, y_true, y_pred, sample_weight=None): y_true = tf.cast(y_true, tf.bool) # 新增阈值判断步骤,将sigmoid概率输出转为二值预测,默认阈值为0.5 y_pred = tf.greater_equal(y_pred, 0.5) y_pred = tf.cast(y_pred, tf.bool) values = tf.logical_and(tf.equal(y_true, True), tf.equal(y_pred, True)) values = tf.cast(values, self.dtype) if sample_weight is not None: sample_weight = tf.cast(sample_weight, self.dtype) values = tf.multiply(values, sample_weight) self.true_positives.assign_add(tf.reduce_sum(values))
如果需要调整分类阈值,直接修改tf.greater_equal的第二个参数即可,比如设置阈值为0.3就改为tf.greater_equal(y_pred, 0.3)。
修改完成后重新运行代码,自定义的真阳性指标就会随着训练进程正常更新,数值也会和内置Recall指标匹配(Recall = 真阳性数 / 真实正例总数)。
内容的提问来源于stack exchange,提问作者BioBroo
相关产品推荐
相关产品推荐

