如何用numpy直接查找两个等长numpy数组的所有匹配集合
Numpy 实现匹配集合查找方案
核心逻辑
你要找的匹配集合本质是row值与col值构成的二分图的连通分量:只要两个col值共享同一个row值,或者两个row值共享同一个col值,就属于同一组。我们用并查集(Union-Find)数据结构实现快速连通分量查找,全程基于numpy接口,效率远高于嵌套循环方案。
完整实现代码
import numpy as np # 测试数据 rows = np.array([3, 5, 6, 8, 8, 9, 9, 9, 10, 10, 10, 11, 11, 12, 13, 14]) cols = np.array([11, 7, 11, 4, 7, 2, 4, 7, 2, 4, 7, 4, 7, 7, 11, 11]) # 步骤1:统一编码row和col节点,避免id冲突 max_row = rows.max() cols_encoded = cols + max_row + 1 # 步骤2:初始化并查集 all_nodes = np.unique(np.concatenate([rows, cols_encoded])) parent = np.arange(all_nodes.max() + 1) # 查找根节点(带路径压缩) def find(u): while parent[u] != u: parent[u] = parent[parent[u]] u = parent[u] return u # 合并两个节点 def union(u, v): u_root = find(u) v_root = find(v) if u_root != v_root: parent[v_root] = u_root # 步骤3:合并所有有对应关系的row和col for r, c in zip(rows, cols_encoded): union(r, c) # 步骤4:分组获取结果 root_labels = np.vectorize(find)(all_nodes) unique_roots, group_ids = np.unique(root_labels, return_inverse=True) result = [] for g_id in range(len(unique_roots)): group_nodes = all_nodes[group_ids == g_id] # 拆分row和编码后的col group_rows = group_nodes[group_nodes <= max_row] group_cols = group_nodes[group_nodes > max_row] - (max_row + 1) result.append((group_rows, group_cols)) # 打印输出 for r_set, c_set in result: print(f"匹配row集合:{r_set},匹配col集合:{c_set}")
输出结果
匹配row集合:[ 3 6 13 14],匹配col集合:[11] 匹配row集合:[ 5 8 9 10 11 12],匹配col集合:[2 4 7]
方案优势
- 时间复杂度接近线性,远高于原有嵌套循环方案,处理百万级数据也不会卡顿
- 结果自动去重,不需要额外后处理
- 仅依赖numpy原生接口,无额外第三方依赖
内容的提问来源于stack exchange,提问作者Deepak Garud
相关产品推荐
相关产品推荐

