Python中基于聚合计数表生成混淆矩阵的最优/最快方法
针对你这种已经聚合好「实际标签、预测标签、出现次数」的统计表格,最优最快的实现思路是避免把计数展开成单个样本列表——毕竟如果单条计数动辄上千上万,展开操作会浪费大量内存和计算时间。下面给你两种实用方案,优先推荐第一种:
方法1:用Pandas Pivot Table(高效首选)
这种方法直接利用Pandas的高效聚合能力,一步生成混淆矩阵,完全不需要展开数据,是大数据量下的最优解。
步骤示例:
首先把你的表格转换成Pandas DataFrame(如果还没转的话):
import pandas as pd # 模拟你的输入数据 data = { 'Actual Label': ['A', 'B', 'C', 'X'], 'Predicted Label': ['A', 'B', 'D', 'Y'], 'Count Occurrences': [200, 150, 15, 5] } df = pd.DataFrame(data)
然后生成混淆矩阵:
# 获取所有唯一标签(实际+预测的并集,防止漏掉类别) all_labels = pd.unique(df[['Actual Label', 'Predicted Label']].values.ravel('K')) all_labels.sort() # 可选,让标签按顺序排列,矩阵更易读 # 生成混淆矩阵的透视表,缺失的(实际,预测)对填充为0 confusion_matrix_df = df.pivot_table( index='Actual Label', columns='Predicted Label', values='Count Occurrences', fill_value=0, aggfunc='sum' # 重新索引,确保所有标签都在矩阵的行和列中 ).reindex(index=all_labels, columns=all_labels, fill_value=0) # 如果需要转换成NumPy数组(方便后续计算) confusion_matrix_np = confusion_matrix_df.values
为什么这是最快的?
- 直接基于聚合数据操作,没有冗余的重复元素生成
- Pandas的透视表底层用了高效的哈希表实现,处理百万级别的聚合行也毫无压力
方法2:用Scikit-learn(适合小数据/兼容Sklearn生态)
如果你已经习惯Sklearn的API,或者需要后续结合Sklearn的其他评估工具(比如分类报告),可以用这种方法,但仅适合小数据量——因为需要把计数展开成单个样本的标签列表。
步骤示例:
from sklearn.metrics import confusion_matrix # 把聚合数据展开成单个样本的标签列表 actual_labels = [] predicted_labels = [] for _, row in df.iterrows(): count = row['Count Occurrences'] actual_labels.extend([row['Actual Label']] * count) predicted_labels.extend([row['Predicted Label']] * count) # 生成混淆矩阵 cm = confusion_matrix(actual_labels, predicted_labels) # 获取标签顺序(和矩阵行/列对应) label_order = pd.unique(actual_labels + predicted_labels) label_order.sort()
注意事项:
- 如果你的计数总和很大(比如超过100万),这种方法会占用大量内存,运行速度远不如第一种
- 生成的矩阵顺序由
confusion_matrix自动排序,你可以用labels参数指定顺序,比如confusion_matrix(actual_labels, predicted_labels, labels=all_labels)
内容的提问来源于stack exchange,提问作者Raphael
相关产品推荐
相关产品推荐

