如何用Matplotlib高效连接散点图中的k近邻节点
优化k近邻图构建的高效方案
嘿,你的这个k近邻图构建方法虽然能跑,但确实效率不高——尤其是如果以后数据集规模超过几百个点的话,问题会更明显。咱们先说说原方案的瓶颈在哪,再用更高效的方法来优化它!
原方案里用pdist生成所有点对的距离,再转成方阵,本质是做了**O(n²)**的计算,还得存一个n×n的距离矩阵。对于250个点来说这点计算不算啥,但如果数据量涨到几千甚至上万,内存和时间都会直接崩掉——毕竟你其实只需要每个点的k个近邻,根本不需要所有点对的距离!
下面给你两种优化思路,从简单到进阶:
方法一:用Scipy的KDTree快速找精确近邻
Scipy自带的KDTree专门针对低维数据的近邻搜索做了优化,能把时间复杂度降到O(n log n),而且完全不需要存全量距离矩阵,内存友好太多。
优化后的代码:
import numpy as np from scipy.spatial import KDTree import matplotlib.pyplot as plt # 生成示例数据 X = np.random.random(500).reshape((250, 2)) k = 4 # 构建KDTree,查询每个点的k+1个近邻(第一个是点自己,所以后面要跳过) tree = KDTree(X) # query返回两个数组:每个点到近邻的距离,以及近邻的索引 distances, indices = tree.query(X, k=k+1) # 先画所有数据点 plt.scatter(X[:, 0], X[:, 1], c='cornflowerblue', s=20) # 遍历每个点,画到它k个近邻的连线 for i in range(X.shape[0]): # 跳过自身,取后面k个真正的近邻 for neighbor_idx in indices[i][1:]: plt.plot( [X[i, 0], X[neighbor_idx, 0]], [X[i, 1], X[neighbor_idx, 1]], c='lightgray', linewidth=0.5 ) plt.title(f'精确k近邻图 (k={k})') plt.show()
为啥这个更高效?
KDTree通过空间划分的方式,把数据分成一个个小区域,查询近邻的时候不用遍历所有点,每个点的k近邻查询只需要**O(log n)**的时间。而且不用存n×n的距离矩阵,内存占用直接从O(n²)降到O(n),大数据集下差异特别明显。
方法二:用Annoy做近似近邻搜索(超大数据集场景)
如果你的数据量特别大(比如超过1万个点),甚至可以用近似近邻搜索工具Annoy(先装一下:pip install annoy)。它比KDTree更快,虽然是近似结果,但在可视化这种对精度要求没那么苛刻的场景下,完全够用。
示例代码:
import numpy as np import annoy import matplotlib.pyplot as plt X = np.random.random(10000).reshape((5000, 2)) k = 4 # 构建Annoy索引,指定维度和距离度量 annoy_index = annoy.AnnoyIndex(2, 'euclidean') for i in range(X.shape[0]): annoy_index.add_item(i, X[i]) # build参数是树的数量,越多精度越高,构建时间也越长 annoy_index.build(10) # 绘制散点图(点调小一点,避免太挤) plt.scatter(X[:, 0], X[:, 1], c='cornflowerblue', s=5) # 遍历每个点画连线 for i in range(X.shape[0]): # 取k+1个近邻,跳过第一个(自己) neighbor_indices = annoy_index.get_nns_by_item(i, k+1)[1:] for neighbor_idx in neighbor_indices: plt.plot( [X[i, 0], X[neighbor_idx, 0]], [X[i, 1], X[neighbor_idx, 1]], c='lightgray', linewidth=0.2 ) plt.title(f'近似k近邻图 (k={k})') plt.show()
小提醒:
- Annoy的结果是近似的,如果必须要精确的近邻,还是用KDTree
build方法的树数量可以根据需求调整,比如要更高精度就设20,要更快就设5
对比原方案的核心优势
| 维度 | 原方案 | KDTree方案 | Annoy方案 |
|---|---|---|---|
| 时间复杂度 | O(n²) | O(n log n) | 接近O(n) |
| 内存占用 | O(n²)(大矩阵) | O(n)(仅存索引) | O(n)(索引更紧凑) |
| 支持数据规模 | 几百个点 | 几千到几万个点 | 十万级以上点 |
内容的提问来源于stack exchange,提问作者nhoeft
相关产品推荐
相关产品推荐

