如何高效获取与测试数组存在非空交集的numpy一维数组行索引?
高效获取与测试数组存在非空交集的行索引
嘿,这个场景我太熟悉了!当你手里有大量int32类型的numpy一维数组(或者列表的列表)时,用Python式的列表推导式+集合交集确实会在数据量上去之后变得很慢——毕竟Python循环和集合操作的开销不小。这里有几个比原生列表推导式高效得多的方案,根据你的数据形态选就行:
方案1:numpy向量化操作(适合规整的2D numpy数组)
如果你的“行集合”已经是一个规整的2D numpy数组,直接用numpy的向量化API就能搞定,全程都是C级别的循环,速度比Python循环快一个数量级以上:
import numpy as np # 假设你的行集合是规整的2D numpy数组 arr = np.array([[1, 2], [2, 3], [4]], dtype=np.int32) test_arr = np.array([1, 4], dtype=np.int32) # 核心逻辑:检查每行是否有元素在测试数组中 # np.in1d 会逐个判断arr中的元素是否在test_arr里,返回一维布尔数组 # reshape回原数组形状后,用any按行判断是否存在交集 mask = np.any(np.in1d(arr, test_arr).reshape(arr.shape), axis=1) # 提取符合条件的索引 indexes = np.where(mask)[0] print(indexes) # 输出: [0 2]
这个方法的优势是完全利用numpy的优化,不需要额外依赖,代码也简洁。如果你的原始数据是列表的列表,先转成numpy数组(np.array(iterable_of_iterables))就行——只要行长度一致,这个转换成本可以忽略。
方案2:numba加速循环(适合行长度不规则的集合)
如果你的行集合是长度不一的列表(比如有的行有2个元素,有的只有1个),numpy的规整数组就不太方便了,这时候用numba把Python循环编译成机器码,速度会比纯Python循环快很多:
from numba import njit import numpy as np # 先把测试数组转成集合,numba对集合的in操作优化很好 test_set = set(np.array([1, 4], dtype=np.int32)) # 用njit装饰器编译函数,把Python代码转成机器码 @njit def find_intersect_indices(iterable): indexes = [] for idx, row in enumerate(iterable): for num in row: if num in test_set: indexes.append(idx) break # 找到一个交集就停止,避免多余遍历 return indexes # 测试用例 iterable = [[1, 2], [2, 3], [4]] print(find_intersect_indices(iterable)) # 输出: [0, 2]
第一次运行这个函数会有编译的开销,但之后每次调用都是纯机器码执行,速度碾压纯Python循环。如果你的数据量很大,这个编译成本完全值得。
方案3:稀疏矩阵法(适合元素范围小的超大规模数据)
如果你的行数量特别多(比如几十万甚至上百万行),且元素的取值范围不大,用稀疏矩阵来做会更高效——它能避免不必要的内存占用,同时利用稀疏矩阵的快速运算:
import numpy as np from scipy.sparse import csr_matrix arr = np.array([[1, 2], [2, 3], [4]], dtype=np.int32) test_arr = np.array([1, 4], dtype=np.int32) # 构建稀疏矩阵:行对应原数组的行索引,列对应元素值,存在的元素标记为1 row_indices = np.repeat(np.arange(arr.shape[0]), arr.shape[1]) col_indices = arr.flatten() data = np.ones_like(row_indices) sparse_mat = csr_matrix((data, (row_indices, col_indices)), shape=(arr.shape[0], np.max(arr)+1)) # 对测试数组对应的列求和,求和结果>0的行就是有交集的行 mask = sparse_mat[:, test_arr].sum(axis=1).A.flatten() > 0 indexes = np.where(mask)[0] print(indexes) # 输出: [0 2]
这个方法的核心是把“元素是否存在”转换成稀疏矩阵的列,通过列求和快速判断行是否有交集,适合元素取值范围有限的超大规模数据集。
内容的提问来源于stack exchange,提问作者demon.mhm
相关产品推荐
相关产品推荐

