如何将带中间赋值的混淆矩阵实现改写为纯列表推导式形式?
列表推导式改写混淆矩阵实现方案
核心思路
原有代码里的中间赋值current_class_p = p[y == i]可以直接整合到列表推导式的逻辑中,不需要单独显式赋值,外层迭代真实标签i时实时计算该类别的预测结果集合,再传递给内层遍历预测标签j的计数逻辑即可。
改写后代码
基础兼容版(适配所有Python版本)
import numpy as np def confusion_matrix_version2(y, p): labels = len(np.unique(y)) return np.array([ [len(p[y == i][p[y == i] == j]) for j in range(labels)] for i in range(labels) ], dtype=int)
性能优化版(Python 3.8+ 支持海象运算符,避免重复计算布尔索引)
import numpy as np def confusion_matrix_version2(y, p): labels = len(np.unique(y)) return np.array([ [len(current_class_p[current_class_p == j]) for j in range(labels)] for i in range(labels) if (current_class_p := p[y == i]) is not None ], dtype=int)
效果验证
输入示例参数测试:
y = np.array([0, 1, 1, 0, 1]) p = np.array([1, 0, 1, 0, 1]) print(confusion_matrix_version2(y, p))
输出结果和要求一致:
[[1 1] [1 2]]
逻辑说明
- 外层列表推导式遍历所有真实标签
i,完全替代原有的外层for i in range(labels)循环 - 无需预先初始化全零矩阵,直接将嵌套列表推导的结果转换为numpy数组返回
- 海象运算符版本仅对
p[y == i]做一次计算,性能和原有显式赋值版本完全一致
内容的提问来源于stack exchange,提问作者MD. ABU SAYED
相关产品推荐
相关产品推荐

