You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何用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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.09.30 09:54:03