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

使用Scikit KNeighborsClassifier获取预测结果与距离

这问题我之前也碰到过,确实直接用predict()和kneighbors()分开调用有点麻烦,还要额外查询训练数据的话更冗余。给你两个可行的方案,按需选择:

方案1:自定义包装函数,一次性获取标签、距离及对应邻居标签

这个方案不用额外查询训练数据,核心是提前把训练集的标签存在内存里,直接通过索引匹配:

from sklearn.neighbors import KNeighborsClassifier
import numpy as np

# 先完成模型训练
X_train = np.array([[1, 2], [3, 4], [5, 6], [7, 8]])
y_train = np.array([0, 0, 1, 1])
knn = KNeighborsClassifier(n_neighbors=2)
knn.fit(X_train, y_train)

# 自定义函数,返回预测标签、k个邻居的距离、k个邻居的标签
def predict_with_distances(model, X_test):
    y_pred = model.predict(X_test)
    distances, indices = model.kneighbors(X_test)
    # 直接用索引取已存的训练集标签,无需额外查询
    neighbor_labels = y_train[indices]
    return y_pred, distances, neighbor_labels

# 测试调用
X_test = np.array([[2, 3], [6, 7]])
y_pred, distances, neighbor_labels = predict_with_distances(knn, X_test)
print("预测标签:", y_pred)
print("对应邻居距离:", distances)
print("对应邻居标签:", neighbor_labels)
方案2:继承模型类,添加自定义预测方法

如果想让模型本身支持这个功能,可以继承KNeighborsClassifier,利用模型内部存储的训练标签(self._y是Scikit-learn训练后自动保存的属性):

from sklearn.neighbors import KNeighborsClassifier
import numpy as np

class KNNWithDistances(KNeighborsClassifier):
    def predict_with_distances(self, X):
        y_pred = self.predict(X)
        distances, indices = self.kneighbors(X)
        # 直接调用模型内部存储的训练标签
        neighbor_labels = self._y[indices]
        return y_pred, distances, neighbor_labels

# 使用示例
knn = KNNWithDistances(n_neighbors=2)
knn.fit(X_train, y_train)
y_pred, distances, neighbor_labels = knn.predict_with_distances(X_test)
补充:获取预测类别对应的代表性距离

如果你只需要和预测标签相关的距离(比如同类邻居的平均距离),可以在上面的基础上再做一层处理:

def predict_with_class_distance(model, X_test):
    y_pred, distances, neighbor_labels = predict_with_distances(model, X_test)
    class_distances = []
    for pred, dists, labels in zip(y_pred, distances, neighbor_labels):
        # 筛选出和预测标签一致的邻居距离
        same_class_dists = dists[labels == pred]
        # 这里取平均,你也可以换成最小/最大距离
        avg_dist = np.mean(same_class_dists)
        class_distances.append(avg_dist)
    return y_pred, np.array(class_distances)

内容的提问来源于stack exchange,提问作者Guyom

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.21 07:21:15