优化代码:为数组A中每个元素寻找B中n₀个最近子集点
高效生成最近邻布尔矩阵的优化方案
给定一维数组A(长度m)、B(长度n),以及满足0 < n₀ < n的数值n₀,需要高效完成以下任务:为A中每个元素找到B里距离最近的n₀个点,生成一个m×n的布尔数组M,需满足:
- M的第i列恰好包含n₀个1和n-n₀个0;
- M[:,i]中值为1的位置,对应B中与A[i]距离最近的n₀个点的索引。
已尝试的实现代码
import numpy as np m = 1000 n = 1200 n_0 = 300 A = np.random.rand(m) B = np.random.rand(n) A_dist, B_dist = np.meshgrid(A, B, sparse=True, copy=False) dist = (A_dist - B_dist)**2 qq = np.quantile(dist, n_0/n, axis = 0) M = (dist <= qq)
原方法的可优化点
原方法通过计算分位数筛选最近的n₀个点,但未利用距离矩阵的核心特性:dist矩阵的每一列对应A中一个元素到所有B元素的距离,所有列的计算都基于同一组B元素。直接对每列计算分位数的方式未充分挖掘计算潜力,当m和n较大时效率受限。
优化思路与实现
核心优化点在于利用B的有序性减少重复计算:
- 先对B数组排序,距离的比较本质是数值大小的相对关系,排序后可通过二分查找快速定位A中元素在B中的位置,进而高效锁定最近的n₀个点;
- 排序后的B中,每个A[i]的最近n₀个点必然集中在其附近的连续区间,通过二分找到插入位置后,直接确定左右扩展范围,避免计算所有B元素的距离。
优化后的代码:
import numpy as np m = 1000 n = 1200 n_0 = 300 A = np.random.rand(m) B = np.random.rand(n) B_sorted = np.sort(B) B_argsort = np.argsort(B) # 记录B排序前的索引,用于还原位置 M = np.zeros((m, n), dtype=bool) for i, a in enumerate(A): # 找到a在排序后B中的插入位置 pos = np.searchsorted(B_sorted, a) # 计算左右需要取的元素数量,优先取更近的一侧 left = pos right = n - pos take_left = min(left, n_0 // 2) take_right = min(right, n_0 - take_left) # 调整左右数量,确保总共取n_0个 if take_left + take_right < n_0: if left > right: take_left += n_0 - (take_left + take_right) else: take_right += n_0 - (take_left + take_right) # 确定选中的索引范围 start = pos - take_left end = pos + take_right # 处理边界情况(当pos靠近数组两端时) if start < 0: end += abs(start) start = 0 if end > n: start -= end - n end = n # 获取排序后B中选中元素的原索引 selected_indices = B_argsort[start:end] # 标记M矩阵 M[i, selected_indices] = True
优化效果说明
- 时间复杂度从原方法的O(mn)降低到O(n log n + m log n),m、n较大时性能提升显著;
- 避免生成完整的距离矩阵,节省大量内存空间;
- 利用有序数组特性,通过二分查找和区间扩展直接定位最近点,无需计算所有距离值,进一步提升效率。
内容的提问来源于stack exchange,提问作者Trevor J Richards
相关产品推荐
相关产品推荐

