You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何使用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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.05.26 08:45:32