如何使用Matplotlib高效连接散点图中的k最近邻点?
优化k近邻图构建的效率问题
嘿,我注意到你现在用pdist加squareform来构建k近邻图,这个方案确实能实现需求,但架不住数据量变大后效率拉胯——毕竟pdist要计算所有点对的距离,这可是O(n²)的复杂度,数据点多了内存和计算时间都会直接爆炸。咱们来聊聊怎么优化这个过程~
先说说现有方案的问题
- 你当前的代码先计算了全量的距离矩阵,对于n=250这种小数据量来说还好,但如果n涨到几千甚至上万,光是存储这个n×n的矩阵就会吃掉大量内存,更别说计算时间了
- 实际上我们只需要每个点的k个最近邻,完全没必要计算所有点对的距离
方案1:用Scikit-learn的NearestNeighbors(最省心高效)
这个库专门做近邻搜索,底层用了KD-Tree或Ball-Tree这类优化过的结构,不需要计算全量距离矩阵,大数据量下优势特别明显。直接上代码:
import numpy as np from sklearn.neighbors import NearestNeighbors import matplotlib.pyplot as plt # 生成示例数据 X = np.random.random(500).reshape((250, 2)) k = 4 # 初始化近邻模型,设置找k+1个近邻(因为会包含点自身) nn_model = NearestNeighbors(n_neighbors=k+1, metric='euclidean') nn_model.fit(X) # 获取每个点的近邻索引,切掉第一个(自身) _, neighbor_indices = nn_model.kneighbors(X) neighbor_indices = neighbor_indices[:, 1:] # 只保留k个真正的近邻 # 绘制k近邻图 plt.scatter(X[:, 0], X[:, 1], s=50, c='#3498db') for idx in range(X.shape[0]): for neighbor_idx in neighbor_indices[idx]: # 绘制当前点到每个近邻的连线 plt.plot( [X[idx, 0], X[neighbor_idx, 0]], [X[idx, 1], X[neighbor_idx, 1]], c='#95a5a6', linewidth=0.5 ) plt.title(f'k={k} Nearest Neighbor Graph') plt.show()
这个方法在n=10000的时候,比全量距离矩阵的方法快好几个数量级,而且内存占用也小很多。
方案2:纯Numpy实现(无需额外依赖)
如果不想引入Scikit-learn,用纯Numpy也能优化找近邻的步骤——核心是用argpartition代替全排序,它能在O(n log k)的时间内找到每个点的k个最小距离的索引,比argsort的O(n log n)快很多:
import numpy as np from matplotlib import pyplot as plt X = np.random.random(500).reshape((250, 2)) k = 4 # 计算全量距离矩阵(这一步还是O(n²),但找近邻的步骤优化了) distances = np.sqrt(((X[:, np.newaxis] - X) ** 2).sum(axis=2)) # 对每个行(每个点的距离),找到k+1个最小距离的索引(包含自身) # argpartition会把最小的k+1个元素放到数组前半部分,不需要完全排序 candidate_indices = np.argpartition(distances, k+1, axis=1)[:, :k+1] # 过滤掉自身,并确保得到的是真正的k个最近邻(argpartition只是分区,不是排序) final_neighbors = [] for idx in range(X.shape[0]): # 获取当前点的候选近邻索引和对应的距离 candidates = candidate_indices[idx] candidate_dists = distances[idx, candidates] # 对候选距离排序,得到真正的近邻顺序 sorted_idx = np.argsort(candidate_dists) # 去掉自身,取前k个近邻 nearest = candidates[sorted_idx][1:k+1] final_neighbors.append(nearest) final_neighbors = np.array(final_neighbors) # 绘图 plt.scatter(X[:, 0], X[:, 1], s=50, c='#3498db') for idx in range(X.shape[0]): for neighbor_idx in final_neighbors[idx]: plt.plot( [X[idx, 0], X[neighbor_idx, 0]], [X[idx, 1], X[neighbor_idx, 1]], c='#95a5a6', linewidth=0.5 ) plt.title(f'k={k} Nearest Neighbor Graph (Numpy Only)') plt.show()
这个方案适合不能引入第三方库的场景,虽然计算距离矩阵还是O(n²),但找近邻的步骤已经优化了,比你原来的代码快不少。
总结一下
- 小数据量(n<1000):两种方案都可以,纯Numpy的也能轻松搞定
- 大数据量(n>1000):优先用Scikit-learn的
NearestNeighbors,避免计算全量距离矩阵,既省内存又省时间
内容的提问来源于stack exchange,提问作者nhoeft
相关产品推荐
相关产品推荐

