如何用NumPy向量化方法替代循环实现一维近邻查找?
用NumPy矢量化方法替代循环查找一维数据的N个最近邻
肯定存在完全符合NumPy风格的矢量化方法,可以彻底替代原代码里的for循环,核心思路是利用NumPy的广播机制一次性计算所有val元素与arr的绝对差,再通过行级的排序/分区操作直接获取前N个最近邻的索引。
方法1:基于argsort的直接矢量化实现
这种方法逻辑和原代码最接近,只是把循环替换成了广播和行级argsort,代码更简洁:
import numpy as np def find_nnearest_vectorized(arr, val, N): # 广播计算所有val元素与arr的绝对差,得到shape为(len(val), len(arr))的矩阵 diff_matrix = np.abs(arr - val[:, np.newaxis]) # 对矩阵的每一行(对应一个val元素)排序,取前N个索引 nearest_idxs = diff_matrix.argsort(axis=1)[:, :N] return nearest_idxs # 测试验证 A = np.arange(10, 20) test = find_nnearest_vectorized(A, A, 3) print(test)
运行这段代码会得到和原代码完全一致的输出,且没有显式for循环。
方法2:基于argpartition的高效实现
如果arr长度很大,且N远小于arr的长度,用argpartition会比argsort更高效(时间复杂度更低)——它不需要对整行完全排序,只需找到前N小的元素位置,再对这些位置做局部排序即可:
import numpy as np def find_nnearest_vectorized_fast(arr, val, N): diff_matrix = np.abs(arr - val[:, np.newaxis]) # 找到每行中前N小元素的索引(此时索引未按距离排序) partitioned_idxs = np.argpartition(diff_matrix, N, axis=1)[:, :N] # 提取每行前N小元素的差值,对差值排序得到局部顺序 diff_subset = diff_matrix[np.arange(len(val))[:, None], partitioned_idxs] sorted_sub_idxs = diff_subset.argsort(axis=1) # 根据局部顺序重新排列索引,得到和原代码一致的有序结果 nearest_idxs = partitioned_idxs[np.arange(len(val))[:, None], sorted_sub_idxs] return nearest_idxs # 测试验证 A = np.arange(10, 20) test = find_nnearest_vectorized_fast(A, A, 3) print(test)
结果验证
两种方法的输出都和原代码完全相同,以示例输入为例,输出为:
[[0 1 2] [1 0 2] [2 1 3] [3 2 4] [4 3 5] [5 4 6] [6 5 7] [7 6 8] [8 7 9] [9 8 7]]
内容的提问来源于stack exchange,提问作者Sterling Butters
相关产品推荐
相关产品推荐

