numpy中查找数组内与测试数组值匹配的所有元素的索引位置
numpy高效匹配行索引实现
场景1:匹配元素固定在a的指定列(示例中为最后一列)
该场景下可通过排序+二分查找实现,时间复杂度为O(n log n + m log n)(n为a的行数,m为b的长度),性能远高于嵌套循环的O(n*m)复杂度。
实现代码如下:
import numpy as np a = np.array([[1, 2, 3, 456], [2, 3, 4, 789], [3, 4, 5, 101112], [4, 5, 6, 131415]]) b = np.array([101112, 456]) # 提取待匹配列,此处为最后一列,可按需调整索引 match_col = a[:, -1] # 构建排序索引 sorter = np.argsort(match_col) sorted_col = match_col[sorter] # 二分查找b中元素的位置 pos = np.searchsorted(sorted_col, b) # 映射回原始行索引 result = sorter[pos] print(result) # 输出 [2 0]
如果需要处理b中存在a无匹配的异常场景,可增加校验逻辑:
# 校验是否匹配成功 match_valid = sorted_col[pos] == b # 仅保留匹配成功的结果 valid_result = sorter[pos[match_valid]]
场景2:匹配元素可出现在a的任意列
该场景下可使用广播向量化比较实现,无嵌套循环:
# 广播比较得到每个b元素在a中各位置的匹配结果 match_mask = b[:, None, None] == a # 按行判断是否存在匹配,提取对应行索引 result = np.where(match_mask.any(axis=-1))[1] print(result) # 输出 [2 0]
内容的提问来源于stack exchange,提问作者D_C
相关产品推荐
相关产品推荐

