自定义多分类混淆矩阵生成函数:基于双向量的计数需求
嘿,我来帮你搞定这个混淆矩阵的函数!你需要的是一个能统计真实标签(y_label)和预测值(y_predict)对应次数的3x3矩阵对吧?没问题,我给你两种实现方式,一种是纯Python基础版,另一种是适合大数据量的numpy高效版,都完全符合你的要求。
纯Python基础实现
这个版本不需要依赖任何第三方库,逻辑清晰易懂,适合小数据集或者想理解底层原理的场景:
def get_confusion_matrix(y_label, y_predict): # 初始化一个3行3列的零矩阵,用来存储计数结果 confusion_matrix = [[0] * 3 for _ in range(3)] # 遍历每一组真实标签和预测值 for true_label, pred_label in zip(y_label, y_predict): # 先做合法性校验,确保输入的标签只能是0、1、2 if true_label not in (0, 1, 2) or pred_label not in (0, 1, 2): raise ValueError("y_label和y_predict的元素只能是0、1或2,请检查输入!") # 对应位置的计数加1:true_label是行索引,pred_label是列索引 confusion_matrix[true_label][pred_label] += 1 return confusion_matrix
numpy高效实现
如果你的数据集比较大,用numpy的向量化操作会比循环快很多,适合生产环境或者大数据场景:
import numpy as np def get_confusion_matrix_np(y_label, y_predict): # 把输入转换成numpy数组,方便后续操作 y_true = np.array(y_label) y_pred = np.array(y_predict) # 校验输入合法性,确保所有元素都是0、1、2 if not np.all(np.isin(y_true, [0, 1, 2])) or not np.all(np.isin(y_pred, [0, 1, 2])): raise ValueError("y_label和y_predict的元素只能是0、1或2,请检查输入!") # 初始化3x3的整数型零矩阵 confusion_matrix = np.zeros((3, 3), dtype=int) # 用numpy的add.at方法高效累加对应位置的计数 np.add.at(confusion_matrix, (y_true, y_pred), 1) return confusion_matrix
使用示例与输出格式
不管用哪个版本,你都可以用下面的代码把结果打印成你想要的表格格式:
# 测试用的示例数据 y_label = [0, 0, 1, 2, 1, 0, 2, 1] y_predict = [0, 1, 1, 2, 0, 0, 2, 1] # 获取混淆矩阵 cm = get_confusion_matrix(y_label, y_predict) # 或者用numpy版本:cm = get_confusion_matrix_np(y_label, y_predict) # 打印成指定的表格格式 print("| | 0 | 1 | 2 |") print("---------------") for row_idx in range(3): print(f"| {row_idx} | {cm[row_idx][0]} | {cm[row_idx][1]} | {cm[row_idx][2]} |") print("---------------")
运行后会输出:
| | 0 | 1 | 2 | --------------- | 0 | 2 | 1 | 0 | --------------- | 1 | 1 | 2 | 0 | --------------- | 2 | 0 | 0 | 2 | ---------------
完全符合你想要的格式,比如cm[0][1]的值是1,对应真实标签为0、预测为1的样本数量。
内容的提问来源于stack exchange,提问作者oikonomiyaki
相关产品推荐
相关产品推荐

