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

如何在Keras训练二元分类器时获取每类误分类数?

获取Keras二元分类器验证集的每类误分类数量

当然可行!其实只需要拿到验证集的真实标签和模型的预测结果,再通过混淆矩阵就能轻松统计每类的误分类数量。下面分两种场景给你具体实现方案:

一、训练完成后一次性统计

假设你已经用train_test_split得到了训练集X_train, y_train和验证集X_val, y_val,并且完成了模型训练。

1. 获取验证集的预测结果

根据你模型最后一层的激活函数,选择对应的预测处理方式:

  • 如果最后一层是sigmoid(输出单个类别的概率值):
    import numpy as np
    
    # 获取预测概率
    y_pred_probs = model.predict(X_val, verbose=0)
    # 用0.5作为阈值转换为类别标签(可根据需求调整阈值)
    y_pred = (y_pred_probs > 0.5).astype(int)
    
  • 如果最后一层是softmax(输出两个类别的概率分布):
    import numpy as np
    
    y_pred_probs = model.predict(X_val, verbose=0)
    # 取概率最大的类别作为预测结果
    y_pred = np.argmax(y_pred_probs, axis=1)
    

2. 计算混淆矩阵并提取误分类数

用sklearn的confusion_matrix生成混淆矩阵,里面的元素直接对应各类的正确/错误分类数量:

from sklearn.metrics import confusion_matrix

# 生成混淆矩阵(注意y_val如果是one-hot编码,要先转成一维标签:y_val = np.argmax(y_val, axis=1))
cm = confusion_matrix(y_val, y_pred)

二元分类的混淆矩阵结构是:

[[真阴性(TN), 假阳性(FP)],
 [假阴性(FN), 真阳性(TP)]]
  • FP:原本是类别0,被误分为类别1的数量
  • FN:原本是类别1,被误分为类别0的数量

直接提取并打印:

misclass_class0 = cm[0][1]  # 类别0的误分类数
misclass_class1 = cm[1][0]  # 类别1的误分类数

print(f"验证集类别0的误分类数量:{misclass_class0}")
print(f"验证集类别1的误分类数量:{misclass_class1}")

二、训练过程中每轮统计误分类数

如果想在每个epoch结束后自动输出验证集的误分类情况,可以自定义一个Keras回调函数:

from tensorflow.keras.callbacks import Callback
from sklearn.metrics import confusion_matrix
import numpy as np

class MisclassificationTracker(Callback):
    def __init__(self, val_data):
        super().__init__()
        self.X_val, self.y_val = val_data

    def on_epoch_end(self, epoch, logs=None):
        # 获取预测结果(根据你的模型输出调整这里的处理逻辑)
        y_pred_probs = self.model.predict(self.X_val, verbose=0)
        y_pred = (y_pred_probs > 0.5).astype(int)
        
        # 处理one-hot编码的真实标签(如果你的y_val是one-hot的话)
        if len(self.y_val.shape) == 2:
            y_val_flat = np.argmax(self.y_val, axis=1)
        else:
            y_val_flat = self.y_val
        
        # 计算混淆矩阵
        cm = confusion_matrix(y_val_flat, y_pred)
        misclass0 = cm[0][1]
        misclass1 = cm[1][0]
        
        # 打印结果
        print(f"\nEpoch {epoch+1} 验证集统计:")
        print(f"类别0误分类数:{misclass0} | 类别1误分类数:{misclass1}\n")

# 训练时传入这个回调
tracker_callback = MisclassificationTracker(val_data=(X_val, y_val))
model.fit(
    X_train, y_train,
    validation_data=(X_val, y_val),
    callbacks=[tracker_callback],
    epochs=10,
    batch_size=32
)

注意事项

  • 如果你的真实标签y_val是one-hot编码格式(比如[[0,1],[1,0],...]),一定要先转换成一维的类别索引(比如[1,0,...])再计算混淆矩阵,否则会出错。
  • 二元分类的阈值(比如0.5)可以根据你的业务需求调整,比如想要更高的召回率,可以适当降低阈值。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.13 07:32:18