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

Keras多分类模型中如何按epoch跟踪混淆矩阵指标?

问题解决与分析

1. 自定义TP指标的逻辑错误

你写的true_pos函数存在明显逻辑bug,导致训练时结果不符合预期:

def true_pos(y_true, y_pred):
    y_pred = K.argmax(y_pred, axis = 1)
    # 错误:第二个条件误写为预测值转float后等于真实标签,完全偏离TP的计算逻辑
    return tf.math.reduce_sum(tf.cast(tf.math.logical_and(tf.math.equal(y_pred, 0),
                              tf.math.equal(tf.cast(y_pred, tf.float32), y_true)), tf.int32))

要计算类别0的TP(真正例),正确逻辑是预测类别为0,且真实类别也为0,修正后的代码:

def true_pos(y_true, y_pred):
    y_pred = tf.argmax(y_pred, axis=1)
    # 真实标签是sparse格式的整数,直接与预测值比较
    return tf.math.reduce_sum(tf.cast(
        tf.math.logical_and(
            tf.math.equal(y_pred, 0),
            tf.math.equal(y_true, 0)  # 改为判断真实标签是否为0
        ), tf.int32))

2. 用tf.confusion_matrix跟踪epoch级指标的正确方式

直接用函数式自定义指标会出错,因为Keras默认会对每个batch的指标结果做平均,但混淆矩阵元素是计数类指标,需要累加而非平均。必须继承tf.keras.metrics.Metric类维护累加状态,才能得到每个epoch的正确数值:

自定义类别TP指标

class ClassTruePos(tf.keras.metrics.Metric):
    def __init__(self, class_id, name='tp', **kwargs):
        super().__init__(name=name, **kwargs)
        self.class_id = class_id
        # 初始化累加变量
        self.tp = self.add_weight(name='tp', initializer='zeros', dtype=tf.int32)

    def update_state(self, y_true, y_pred, sample_weight=None):
        y_pred = tf.argmax(y_pred, axis=1)
        # 筛选真实与预测均为目标类的样本
        matches = tf.math.logical_and(
            tf.math.equal(y_true, self.class_id),
            tf.math.equal(y_pred, self.class_id)
        )
        matches = tf.cast(matches, tf.int32)
        # 处理样本权重(可选)
        if sample_weight is not None:
            sample_weight = tf.cast(sample_weight, tf.int32)
            matches = tf.multiply(matches, sample_weight)
        # 累加当前batch的TP数
        self.tp.assign_add(tf.reduce_sum(matches))

    def result(self):
        return self.tp

    def reset_state(self):
        # 每个epoch开始前重置累加值
        self.tp.assign(0)

编译模型时添加指标

model.compile(
    optimizer='adam',
    loss='sparse_categorical_crossentropy',
    metrics=[
        tf.keras.metrics.CategoricalAccuracy(),
        ClassTruePos(0, name='tp_class0'),
        ClassTruePos(1, name='tp_class1'),
        ClassTruePos(2, name='tp_class2'),
        ClassTruePos(3, name='tp_class3')
    ]
)

训练时每个epoch会输出对应类别的累计TP值。如果需要跟踪FP、FN等其他混淆矩阵元素,用同样逻辑自定义Metric类即可。

3. 为什么Keras不内置多分类混淆矩阵指标?

主要两点原因:

  • 显示限制:Keras训练日志仅支持展示标量指标,混淆矩阵是二维数组,无法直接在日志里直观呈现;
  • 需求灵活性:不同用户对混淆矩阵的需求差异极大——有人需要单个类别的TP/FP,有人需要完整矩阵,有人需要归一化结果。内置通用实现无法覆盖所有场景,因此Keras仅提供准确率、召回率等通用标量指标,将复杂自定义需求交给用户。

另外,你用sklearn.metrics.confusion_matrix结合model.predict()获取最终混淆矩阵的方式完全合理,适合训练后的整体评估。


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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.14 11:41:02