在R语言中高效查找仅基于历史数据的k近邻方法
增量式前k近邻搜索的效率优化方案
问题背景
有10000行×8列的数值型数据,需要为第1001到10000行的每一行,在该行之前的所有行中找到欧氏距离最近的前k个邻居。目前采用先通过RANN获取2k/5k个近邻再过滤的方式,冗余计算过多,速度太慢,需要优化。
现有慢实现(示例)
set.seed(123) data = matrix(rnorm(1000), nrow = 200, ncol = 5) result <- list() for (i in c(101:200)) { distances <- apply(data[1:(i-1),], 1, function(x) { dist(rbind(x, data[i, ])) }) neighbors <- sort(distances, index.return = TRUE)$ix[1:3] result[[i - 100]] <- neighbors }
优化建议及实现
1. 替换apply为向量化距离计算(快速见效)
原实现中apply循环计算每行距离是最大的性能瓶颈,改用矩阵向量化运算直接计算欧氏距离,效率能提升数倍:
set.seed(123) data = matrix(rnorm(1000), nrow = 200, ncol = 5) result <- list() k <- 3 for (i in 101:200) { # 向量化计算与前i-1行的欧氏距离 diff_mat <- data[1:(i-1), ] - data[i, ] distances <- sqrt(rowSums(diff_mat^2)) # 取前k个最小距离的索引 neighbors <- order(distances)[1:k] result[[i - 100]] <- neighbors }
2. 用Rcpp实现堆筛选(大幅降低时间复杂度)
当行数很大时(比如i=10000时要处理9999行),排序所有距离的时间成本很高。改用大顶堆维护前k个最小距离,无需排序所有数据,时间复杂度从O(n log n)降至O(n*d)(d为列数):
先编写Rcpp代码(保存为get_top_k.cpp):
#include <Rcpp.h> #include <queue> #include <vector> using namespace Rcpp; // [[Rcpp::export]] IntegerVector get_top_k_neighbors(NumericMatrix prev_data, NumericVector current_row, int k) { int n_rows = prev_data.nrow(); int n_cols = prev_data.ncol(); // 大顶堆:存储(距离平方, 行索引),堆顶是当前最大的距离平方 std::priority_queue<std::pair<double, int>> max_heap; for (int i = 0; i < n_rows; ++i) { double dist_sq = 0.0; for (int j = 0; j < n_cols; ++j) { double diff = prev_data(i, j) - current_row[j]; dist_sq += diff * diff; } if (max_heap.size() < k) { max_heap.push(std::make_pair(dist_sq, i + 1)); // R的行索引从1开始 } else if (dist_sq < max_heap.top().first) { max_heap.pop(); max_heap.push(std::make_pair(dist_sq, i + 1)); } } // 提取结果并反转(堆是从大到小,需要转为从小到大) IntegerVector res(k); for (int i = k - 1; i >= 0; --i) { res[i] = max_heap.top().second; max_heap.pop(); } return res; }
然后在R中调用:
sourceCpp("get_top_k.cpp") set.seed(123) data <- matrix(rnorm(10000*8), nrow = 10000, ncol = 8) k <- 5 start_idx <- 1001 result <- vector("list", length = nrow(data) - start_idx + 1) for (i in start_idx:nrow(data)) { neighbors <- get_top_k_neighbors(data[1:(i-1), ], data[i, ], k) result[[i - start_idx + 1]] <- neighbors }
3. 增量式空间索引(超大数据集适用)
如果数据集规模远超10万行,可以使用支持增量添加向量的空间索引库,比如rfaiss(FAISS的R绑定),它能动态维护索引,每次查询后将当前行加入索引,避免重复构建:
library(rfaiss) set.seed(123) data <- matrix(rnorm(10000*8), nrow = 10000, ncol = 8) k <- 5 start_idx <- 1001 result <- vector("list", length = nrow(data) - start_idx + 1) # 初始化IVF_FLAT索引(适合欧氏距离,可根据数据调整类型) index <- faiss_index_factory(ncol(data), "IVF_FLAT", metric = "L2") # 添加前1000行到索引 faiss_add(index, data[1:1000, ]) for (i in start_idx:nrow(data)) { # 查询前k个近邻 query_res <- faiss_search(index, data[i, , drop = FALSE], k) neighbors <- query_res$indices[1, ] + 1 # FAISS索引从0开始,转为R的1-based索引 result[[i - start_idx + 1]] <- neighbors # 将当前行加入索引 faiss_add(index, data[i, , drop = FALSE]) }
内容的提问来源于stack exchange,提问作者cccanhakan
相关产品推荐
相关产品推荐

