如何识别NumPy二维数组中精确匹配指定元素数量的行对
问题
现有一个二维NumPy数组,其特性为每列包含互不重叠的独立整数范围(示例数组如下,比如数字8仅会出现在第3列):
import numpy as np b = np.array([[1, 5, 9, 11, 13], [1, 6, 8, 10, 14], [2, 4, 8, 12, 15], [2, 5, 7, 11, 13], [3, 4, 9, 10, 15], [3, 5, 7, 12, 14]])
需要找出所有行对,记录它们的行索引以及精确匹配的元素个数,理想输出格式如下:
out = [[0, 0, 5], [0, 1, 1], [0, 2, 0], [0, 3, 3], ...]
核心需求是高效计算行对之间精确匹配的元素数量。
解决方案
利用数组特性简化计算
因为数组每列的整数范围互不重叠,所以两个相等的元素必然处于同一列,不会出现不同列有相同数字的情况。因此,两行之间的匹配元素数量,等价于这两行在相同列位置上元素相等的个数。
代码实现
import numpy as np # 转换为numpy数组 b = np.array([[1, 5, 9, 11, 13], [1, 6, 8, 10, 14], [2, 4, 8, 12, 15], [2, 5, 7, 11, 13], [3, 4, 9, 10, 15], [3, 5, 7, 12, 14]]) # 计算所有行对的匹配数:广播比较每一行与所有行的对应列元素,再求和 match_counts = (b[:, np.newaxis] == b).sum(axis=2) # 生成所有行对的索引 rows, cols = np.indices(match_counts.shape) # 组合成目标输出格式 out = np.stack([rows.ravel(), cols.ravel(), match_counts.ravel()], axis=1) print(out)
代码说明
b[:, np.newaxis] == b:通过广播将数组变形为(n_rows, n_rows, n_cols)的三维数组,每个位置表示第i行和第j行在第k列的元素是否相等。.sum(axis=2):对列维度求和,得到(n_rows, n_rows)的矩阵,其中match_counts[i,j]就是第i行和第j行的匹配元素个数。np.indices(match_counts.shape):生成行对的索引矩阵,再通过ravel()扁平化,最后和匹配数组合成目标的二维数组。
内容的提问来源于stack exchange,提问作者user109387
相关产品推荐
相关产品推荐

