如何快速实现100维空间中200万数据点的最近邻欧氏距离计算?
优化思路与代码实现
针对你处理百万级高维数据找最近邻的场景,原代码的主要瓶颈在于逐行迭代Pandas DataFrame和不必要的全量排序。下面从逻辑和代码层面给出具体优化方案:
核心优化方向
- 抛弃逐行迭代,用向量运算代替循环:Numpy的广播机制能把高维距离计算转化为批量操作,效率比Pandas行循环高几个数量级。
- 用
argmin替代排序找Top1:排序是O(nlogn)复杂度,而找最小索引的argmin是O(n),对于百万级数据来说节省大量时间。 - 从DataFrame转成Numpy数组计算:Pandas的行操作有额外开销,纯Numpy数组的数值计算速度更快。
优化后的代码
import pandas as pd import numpy as np # 模拟数据(保持你的测试数据结构) data = {'Name': ['Ly','Gr','Er','Ca','Cy','Sc','Cr','Cn','Le','Cs','An','Ta','Sa','Ly','Az','Sx','Ud','Lr','Si','Au','Co','Ck','Mj','wa'], 'dim0': [33,-9,18,-50,39,-23,-19,89,-74,81,8,23,-63,-62,-14,45,39,-46,74,19,7,97,-29,71,], 'dim1': [-7,75,77,-93,-89,4,-96,-64,41,-27,-87,23,-69,-77,-92,18,21,27,-76,-57,-44,20,15,-76,], 'dim2': [-31,54,-14,-93,72,-14,65,44,-88,19,48,-51,-25,36,-46,98,8,0,53,-47,-29,95,65,-3,], 'dim3': [-12,-86,10,93,-79,-55,-6,-79,-12,66,-81,-14,44,84,9,-19,-69,29,-50,-59,35,-28,90,-73,], } df = pd.DataFrame(data) # 拆分目标集和对比集 target_indices = df.tail(5).index df_target = df.loc[target_indices].copy() df_compare = df.drop(target_indices) # 提取数值维度转成Numpy数组(关键:用数组做向量运算) dim_cols = [col for col in df.columns if col.startswith('dim')] target_arr = df_target[dim_cols].values compare_arr = df_compare[dim_cols].values compare_names = df_compare['Name'].values # 批量计算平方欧氏距离:利用Numpy广播,shape会变成(目标点数, 对比点数) # 平方欧氏距离 = sum((x_i - y_j)^2) 等价于 ||x||² + ||y||² - 2x·y(可以进一步优化计算) target_sq = np.sum(target_arr ** 2, axis=1, keepdims=True) compare_sq = np.sum(compare_arr ** 2, axis=1, keepdims=True) dot_product = np.dot(target_arr, compare_arr.T) distances_sq = target_sq + compare_sq.T - 2 * dot_product # 对每个目标点,找到距离最小的对比点索引 closest_indices = np.argmin(distances_sq, axis=1) # 映射回名称,赋值给目标集 df_target['closest_neighbour'] = compare_names[closest_indices] print(df_target)
关键优化点解释
距离计算的数学优化:
我们没有直接计算每个维度差的平方和,而是用了平方欧氏距离的等价公式:
$$d^2(x,y) = ||x||^2 + ||y||^2 - 2x·y$$
这个公式利用矩阵乘法(dot_product)批量计算,比逐维度循环快得多,尤其在高维场景下优势明显。批量处理目标点:
一次性计算所有目标点和对比点的距离矩阵,避免了逐个目标点循环,充分利用了CPU的向量化计算能力。避免全量排序:
用np.argmin直接找到每个目标点对应的最小距离索引,不需要对整个对比集排序,时间复杂度从O(nlogn)降到O(n),对于200万条数据来说,这个优化能节省大量时间。减少Pandas操作:
只在数据输入输出时用Pandas,核心计算全部用Numpy数组,避免了Pandas行操作的额外开销。
超大规模数据的进阶优化
如果你的数据量继续增长(比如超过千万级),还可以考虑:
- 用KD-Tree或Ball-Tree:Scikit-learn的
sklearn.neighbors.KDTree或BallTree专门针对最近邻搜索优化,适合高维数据,能把时间复杂度降到O(logn)。 - 分块计算:如果内存不够装下整个对比集,可以分块加载计算距离,最后合并结果。
- 利用GPU加速:用CuPy替代Numpy,或者用PyTorch/TensorFlow的GPU张量计算,进一步提升速度。
内容的提问来源于stack exchange,提问作者Vincent
相关产品推荐
相关产品推荐

