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

Keras自定义单类别Precision、Recall指标每轮训练多步后变为NaN

问题原因
  • 核心触发点是除零运算:你实现的precision指标分母为TP+FP(即当前batch中模型预测为正类的样本总数),如果某个batch内模型对目标类别(最后一维)的预测结果全部小于0.5,那么TP和FP都会为0,分母为0直接得到NaN;同理recall的分母是TP+FN(当前batch中真实为正类的样本总数),如果某个batch内目标类别的真实标签全为负,也会触发除零得到NaN。
  • 训练初期模型预测随机性强,每个batch基本都会有正类预测/真实正样本,因此指标正常;训练后期模型收敛,预测置信度变高,很容易出现某个batch内无目标类正预测/无真实正样本的情况,就会陆续出现NaN,且一旦单个batch出现NaN,Keras默认对batch级指标做平均的逻辑会导致后续所有平均指标都为NaN,符合你观察到的现象。
  • 验证集指标正常是因为验证集的batch数据分布更稳定,没有出现上述极端的batch分布。
解决方法

方案1:最简单的快速修复

给除法运算添加极小平滑项,避免除零,直接修改自定义指标的除法逻辑即可,以precision为例:

def precision(y_true, y_pred):
    '''
    Calculates precision metric over gun label
    Precision = TP/(TP+FP)
    '''
    # 仅关注最后一个标签
    y_true = y_true[:,-1]
    y_pred = y_pred[:,-1]
    y_pred = tf.where(y_pred>.5, 1, 0)

    y_pred = tf.cast(y_pred, tf.float32)
    y_true = tf.cast(y_true, tf.float32)

    true_positives = K.sum(y_true * y_pred)
    false_positive = tf.math.reduce_sum(tf.where(tf.logical_and(tf.not_equal(y_true,y_pred), y_pred==1), 1, 0))
    false_positive = tf.cast(false_positive, tf.float32)
    # 添加K.epsilon()避免除零
    precision = true_positives / (true_positives + false_positive + K.epsilon())
    return precision

recall的修改逻辑一致,在分母处添加K.epsilon()即可。

方案2:使用Keras内置指标(更推荐)

Keras自带的Precision、Recall指标已经内置了除零保护和全局状态累计功能,且支持指定计算单个类别的指标,直接调用即可无需自己实现,代码更稳定:

model.compile(
    loss='binary_crossentropy', 
    optimizer=optimizer, 
    metrics=[
        'accuracy',
        tf.keras.metrics.Precision(class_id=-1, name='precision'),
        tf.keras.metrics.Recall(class_id=-1, name='recall')
    ]
)

内置指标是跨batch累计全局的TP、FP、FN来计算整个epoch的指标,结果比单batch平均的自定义指标更准确,也不会出现单batch极端情况导致的NaN问题。

方案3:自定义全局指标类

如果需要自定义特殊逻辑,可以继承tf.keras.metrics.Metric实现全局累计的指标类,避免单batch的波动问题,示例如下:

class ClassPrecision(tf.keras.metrics.Metric):
    def __init__(self, class_id=-1, name='precision', **kwargs):
        super().__init__(name=name, **kwargs)
        self.class_id = class_id
        self.true_positives = self.add_weight(name='tp', initializer='zeros')
        self.false_positives = self.add_weight(name='fp', initializer='zeros')

    def update_state(self, y_true, y_pred, sample_weight=None):
        y_true = y_true[:, self.class_id]
        y_pred = y_pred[:, self.class_id]
        y_pred = tf.cast(y_pred > 0.5, tf.float32)
        y_true = tf.cast(y_true, tf.float32)
        
        tp = tf.reduce_sum(y_true * y_pred)
        fp = tf.reduce_sum(tf.where(tf.logical_and(y_true != y_pred, y_pred == 1), 1.0, 0.0))
        self.true_positives.assign_add(tp)
        self.false_positives.assign_add(fp)

    def result(self):
        return self.true_positives / (self.true_positives + self.false_positives + K.epsilon())

    def reset_state(self):
        self.true_positives.assign(0)
        self.false_positives.assign(0)

使用时直接传入metrics参数即可。

内容的提问来源于stack exchange,提问作者theastronomist

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.29 08:09:03