如何在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
相关产品推荐
相关产品推荐

