You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何在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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.09.25 14:36:05