高效低内存K近邻搜索方案问询:适配多邻居场景
高效1维K近邻搜索:搞定百万级数据的内存与速度痛点
我完全懂你这个场景的痛苦——百万级数据找K近邻,全量计算直接爆内存,循环单条又慢到离谱。结合你的需求(1维、多邻居、低内存、快速度),我给你整理了问题分析和优化方案:
问题背景
你需要为search_for(100万条)里的每个点,在search_in(10万条)中找K个最近邻,核心诉求是低内存+快速度,重点支持多邻居(K>1),维度只需要1维。
之前的方案要么内存爆炸(全量距离矩阵要600GiB),要么速度拉胯(循环100万次单次超1小时),都没法满足生产需求。
现有方案的问题分析
1. 小数据集多维度KNN(内存爆炸版)
这个代码适合小数据、多维度,但全量生成100万×10万的距离矩阵直接把内存干爆,完全不适合大数据:
def knn_search(search_for, search_in, K = 1, return_col = ["ID"], col = 'A'): a_search_in = array(search_in[col]) a_search_for = array(search_for[col]) a = np.tile(a_search_for, [a_search_in.shape[0], 1]).T b = np.tile(a_search_in, [a_search_for.shape[0], 1]) t_diff = a - b diff = np.square(t_diff) idx = argsort(diff) if search_for.shape[0] == 1: return idx[:K] elif K == 1: return search_in.iloc[np.concatenate(idx[:,:K]), :][return_col] else: tmp = pd.DataFrame() for i in range(min(K, search_in.shape[0])): tmp = pd.concat([tmp.reset_index(drop=True), search_in.iloc[idx[:,i], :][[return_col]].reset_index(drop=True)], axis=1) return tmp
2. 1维1邻居高效版(不支持多邻居)
这个用searchsorted找插入位置,速度快内存低,但只能处理K=1的情况,没法扩展到多邻居:
def knn_search_1K_1D(search_for, search_in, return_col = ["ID"], col = 'A'): sort_search_in = search_in.sort_values(col).reset_index() idx = np.searchsorted(sort_search_in[col], search_for[col]) idx_pop = np.where(idx > len(sort_search_in) - 1, len(sort_search_in) - 1, idx) t = sort_search_in.iloc[idx_pop , :][[return_col]] search_for_nn = pd.concat([search_for.add_prefix('').reset_index(drop=True), t.add_prefix('nn_').reset_index(drop=True)], axis=1)
3. 1维多邻居低效版(耗时超1小时)
这个循环每个点计算全量距离排序,能支持K>1,但百万级数据下循环次数太多,单次搜索就超1小时,完全没法用:
def knn_search_nK_1D(search_for, search_in, K = 1, return_col = ["ID"], col = 'A'): t = [] #looping one point by one for i in range(search_for.shape[0]): y = search_in[col] x = search_for.iloc[i, :][col] nn = np.nanmean(search_in.iloc[np.argsort(np.abs(np.subtract(y, x)))[0:K], :][return_col]) t.append(nn) search_for_nn = search_for search_for_nn['nn_' + return_col] = t
示例数据与预期输出
search_for = pd.DataFrame({'ID': ["F", "G"], 'A' : [-1, 9]}) search_in = pd.DataFrame({'ID': ["A", "B", "C", "D", "E"], 'A' : [1, 2, 3, 4, 5 ]}) # 调用1邻居版本的预期输出 t = knn_search(search_for = search_for , search_in = search_in, K = 1, return_col = ['ID'], col = 'A') print(t) # 输出: # ID #0 A #1 E
优化方案:排序+滑动窗口的1维K近邻搜索
核心思路:先把search_in排序,用searchsorted找到每个点的插入位置,然后在插入位置前后的2K个候选点里找最近的K个。因为排序后,最近的K个点肯定在插入位置附近的小范围内,不需要遍历全部数据。
这个方案既避免了全量距离矩阵的内存问题,又把计算量降到最低,百万级数据几分钟就能搞定。
优化后的代码
import numpy as np import pandas as pd def fast_knn_1d(search_for, search_in, K=1, return_col=["ID"], col='A'): # 预处理:排序search_in,提取数值列和结果列 sorted_in = search_in.sort_values(col).reset_index(drop=True) sorted_vals = sorted_in[col].values sorted_return = sorted_in[return_col].values # 获取每个search_for点的插入位置 search_vals = search_for[col].values insert_pos = np.searchsorted(sorted_vals, search_vals) # 为每个点生成候选范围:从max(0, pos-K)到min(len(sorted_vals), pos+K) start = np.maximum(0, insert_pos - K) end = np.minimum(len(sorted_vals), insert_pos + K) # 批量处理所有点的候选距离计算 results = [] for idx, (s, e, val) in enumerate(zip(start, end, search_vals)): # 取候选区间的数值和对应返回列 candidates_vals = sorted_vals[s:e] candidates_return = sorted_return[s:e] # 计算距离,排序取前K个 distances = np.abs(candidates_vals - val) sorted_idx = np.argsort(distances)[:K] # 收集结果 selected = candidates_return[sorted_idx] results.append(selected.flatten()) # 整理结果为DataFrame,和原数据合并 # 处理列名:根据K和返回列数量生成 if K == 1: col_names = [f'nn_{rc}' for rc in return_col] else: col_names = [f'nn_{k+1}_{rc}' for k in range(K) for rc in return_col] result_df = pd.DataFrame(results, columns=col_names) final_df = pd.concat([search_for.reset_index(drop=True), result_df.reset_index(drop=True)], axis=1) return final_df
方案优势
- 内存友好:不需要生成全量距离矩阵,每个点只处理最多
2K个候选点,内存占用仅和K成正比,和数据总量无关。 - 速度极快:排序只做一次(O(N log N),10万条数据几毫秒),每个点的候选范围很小,计算距离和排序的时间可以忽略不计,百万级数据几分钟就能处理完。
- 完美适配多邻居:支持任意K值(只要K≤
search_in的长度),还能同时返回多个列的结果。
测试示例
search_for = pd.DataFrame({'ID': ["F", "G"], 'A' : [-1, 9]}) search_in = pd.DataFrame({'ID': ["A", "B", "C", "D", "E"], 'A' : [1, 2, 3, 4, 5 ]}) # 找2个最近邻 result = fast_knn_1d(search_for, search_in, K=2, return_col=['ID'], col='A') print(result) # 输出: # ID A nn_1_ID nn_2_ID #0 F -1 A B #1 G 9 E D
内容的提问来源于stack exchange,提问作者AAAA
相关产品推荐
相关产品推荐

