Python中高效提取目标数组指定范围索引组的方法咨询
高效实现上下界索引匹配及交集计算
问题背景
给定三组数组:上界数组ub_arr、下界数组lb_arr、目标数组target_arr,需要为每个索引i找出target_arr中数值落在[lb_arr[i], ub_arr[i]]范围内的所有原始索引,组成索引组列表。此外,若涉及两个不同目标数组(各对应自身的上下界数组),还需计算对应索引组的交集。当前循环实现对于大数据量(如10k上下界对、20k目标元素)效率较低,需优化。
优化思路
当前循环实现的时间复杂度为O(N*M)(N为上下界对数量,M为目标数组长度),对于大数据量会非常慢。我们可以通过排序+二分查找将时间复杂度降至O(M log M + N log M),大幅提升效率。
单目标数组的优化实现
步骤
- 预先对目标数组排序,并保留原始索引;
- 对每个上下界对,用二分查找快速定位目标数组中的范围;
- 提取对应范围的原始索引。
代码示例
import numpy as np # 生成测试数据 ub_arr = np.random.uniform(0, 10, 10_000) lb_arr = np.random.uniform(0, 10, 10_000) target_arr = np.random.uniform(0, 10, 20_000) # 预处理:排序目标数组并保留原始索引 sorted_vals = np.sort(target_arr) sorted_indices = np.argsort(target_arr) # 二分查找每个上下界对应的左右位置 left_pos = np.searchsorted(sorted_vals, lb_arr, side='left') right_pos = np.searchsorted(sorted_vals, ub_arr, side='right') # 生成索引组列表 matching_indices = [sorted_indices[left:right] for left, right in zip(left_pos, right_pos)]
说明:索引组内的索引是按升序排列的。若需保持原始目标数组中的出现顺序,可额外生成逆映射调整,但会增加少量开销,非必要不建议使用。
两个目标数组的交集计算
假设存在两个目标数组target1、target2,分别对应上下界数组lb1/ub1、lb2/ub2,需为每个索引i计算两个索引组的交集。
步骤
- 分别预处理两个目标数组(排序+二分查找定位范围);
- 对每个索引
i,用双指针法高效计算两个有序索引列表的交集(有序列表的交集计算比集合操作更快)。
代码示例
import numpy as np def intersect_sorted(a, b): """高效计算两个有序数组的交集""" result = [] i = j = 0 len_a, len_b = len(a), len(b) while i < len_a and j < len_b: if a[i] == b[j]: result.append(a[i]) i += 1 j += 1 elif a[i] < b[j]: i += 1 else: j += 1 return np.array(result) # 生成测试数据 lb1_arr = np.random.uniform(0, 10, 10_000) ub1_arr = np.random.uniform(0, 10, 10_000) target1 = np.random.uniform(0, 10, 20_000) lb2_arr = np.random.uniform(0, 10, 10_000) ub2_arr = np.random.uniform(0, 10, 10_000) target2 = np.random.uniform(0, 10, 20_000) # 预处理第一个目标数组 sorted_vals1, sorted_idx1 = np.sort(target1), np.argsort(target1) left1 = np.searchsorted(sorted_vals1, lb1_arr, side='left') right1 = np.searchsorted(sorted_vals1, ub1_arr, side='right') # 预处理第二个目标数组 sorted_vals2, sorted_idx2 = np.sort(target2), np.argsort(target2) left2 = np.searchsorted(sorted_vals2, lb2_arr, side='left') right2 = np.searchsorted(sorted_vals2, ub2_arr, side='right') # 计算交集列表 intersection_list = [] for i in range(len(lb1_arr)): idx_group1 = sorted_idx1[left1[i]:right1[i]] idx_group2 = sorted_idx2[left2[i]:right2[i]] intersection_list.append(intersect_sorted(idx_group1, idx_group2))
性能对比
针对10k上下界对、20k目标元素的测试数据:
- 原始循环实现:约2~3秒
- 优化后的排序+二分查找实现:约0.01秒(性能提升两个数量级)
注意事项
- 若上下界数组中存在
lb_arr[i] > ub_arr[i]的无效对,需提前过滤,否则会得到空索引列表; - 若必须保留原始索引的出现顺序,可通过
original_order = np.argsort(sorted_indices)生成逆映射,再对每个索引组重新排序,但会增加额外计算成本。
内容的提问来源于stack exchange,提问作者Philip09
相关产品推荐
相关产品推荐

