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

如何在Python中基于scipy.cdist生成的矩形距离矩阵计算k近邻?

基于矩形距离矩阵计算跨样本集的k近邻

问题背景

我想用sklearn、scipy或numpy,基于scipy.spatial.distance.cdist输出的矩形距离矩阵计算k近邻。尝试将该矩阵传入kneighbors_graph和KNeighborsTransformer并设置metric="precomputed"时失败,具体代码及报错如下:

代码示例

from scipy.spatial.distance import cdist
from sklearn.datasets import make_classification
from sklearn.neighbors import kneighbors_graph, KNeighborsTransformer

X, _ = make_classification(n_samples=15, n_features=4, n_classes=2, n_clusters_per_class=1, random_state=0)
A = X[:10,:]
B = X[10:,:]
A.shape, B.shape
# ((10, 4), (5, 4))

# 矩形距离矩阵:A中每个样本到B中每个样本的距离
dist = cdist(A,B)
dist.shape
# (10, 5)

n_neighbors=3
kneighbors_graph(dist, n_neighbors=n_neighbors, metric="precomputed")

报错信息

ValueError                                Traceback (most recent call last)
Cell In[165], line 17
      14 # (10, 5)
      16 n_neighbors=3
---> 17 kneighbors_graph(dist, n_neighbors=n_neighbors, metric="precomputed")

File ~/miniconda3/envs/soothsayer_env/lib/python3.9/site-packages/sklearn/neighbors/_graph.py:117, in kneighbors_graph(X, n_neighbors, mode, metric, p, metric_params, include_self, n_jobs)
     50 """Compute the (weighted) graph of k-Neighbors for points in X.
     51 
     52 Read more in the :ref:`User Guide <unsupervised_neighbors>`.
   (...)
    114        [1., 0., 1.]])
    115 """
    116 if not isinstance(X, KNeighborsMixin):
--> 117     X = NearestNeighbors(
    118         n_neighbors=n_neighbors,
    119         metric=metric,
    120         p=p,
    121         metric_params=metric_params,
    122         n_jobs=n_jobs,
    123     ).fit(X)
    124 else:
    125     _check_params(X, metric, p, metric_params)

File ~/miniconda3/envs/soothsayer_env/lib/python3.9/site-packages/sklearn/neighbors/_unsupervised.py:176, in NearestNeighbors.fit(self, X, y)
    159 """Fit the nearest neighbors estimator from the training dataset.
    160 
    161 Parameters
   (...)
    173     The fitted nearest neighbors estimator.
    174 """
    175 self._validate_params()
--> 176 return self._fit(X)

File ~/miniconda3/envs/soothsayer_env/lib/python3.9/site-packages/sklearn/neighbors/_base.py:545, in NeighborsBase._fit(self, X, y)
    543     # Precomputed matrix X must be squared
    544     if X.shape[0] != X.shape[1]:
--> 545         raise ValueError(
    546             "Precomputed matrix must be square."
    547             " Input is a {}x{} matrix.".format(X.shape[0], X.shape[1])
    548         )
    549     self.n_features_in_ = X.shape[1]
    551 n_samples = X.shape[0]

ValueError: Precomputed matrix must be square. Input is a 10x5 matrix.

原因分析

sklearn的kneighbors_graph和KNeighborsTransformer在metric="precomputed"模式下,要求输入的距离矩阵必须是方阵(即同一数据集内样本两两之间的距离)。而cdist(A,B)生成的是矩形矩阵,描述的是两个不同数据集A和B之间的样本距离,不符合该模式的输入要求。

解决方案

方案1:用numpy手动提取k近邻

直接利用numpy的排序函数,从矩形距离矩阵中快速筛选每个A样本对应的B中k个最近邻:

import numpy as np

n_neighbors = 3
# 获取每个A样本对应的B中k个近邻的索引
neighbor_indices = np.argpartition(dist, n_neighbors, axis=1)[:, :n_neighbors]
# 获取对应的距离值
neighbor_distances = np.take_along_axis(dist, neighbor_indices, axis=1)

# 输出结果
print("每个A样本的k近邻在B中的索引:")
print(neighbor_indices)
print("\n对应的距离值:")
print(neighbor_distances)

方案2:用sklearn NearestNeighbors API处理

先拟合目标数据集B,再将预计算的A→B距离矩阵传入kneighbors方法:

from sklearn.neighbors import NearestNeighbors

# 拟合B数据集
nn = NearestNeighbors(n_neighbors=n_neighbors).fit(B)
# 传入预计算的距离矩阵获取近邻信息
distances, indices = nn.kneighbors(X=dist, n_neighbors=n_neighbors, metric="precomputed")

print("每个A样本的k近邻在B中的索引:")
print(indices)
print("\n对应的距离值:")
print(distances)

方案3:构建邻接矩阵(替代kneighbors_graph)

如果需要生成类似kneighbors_graph的稀疏邻接矩阵,可手动构建:

from scipy.sparse import csr_matrix

rows = np.repeat(np.arange(A.shape[0]), n_neighbors)
cols = indices.ravel()
data = np.ones_like(cols)  # 若需加权,替换为distances.ravel()

# 构建A→B的邻接矩阵,形状为(10,5)
adjacency_matrix = csr_matrix((data, (rows, cols)), shape=(A.shape[0], B.shape[0]))
print("邻接矩阵:")
print(adjacency_matrix.toarray())

内容的提问来源于stack exchange,提问作者O.rka

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.23 08:05:54