Skip to content

K 近邻与 kd 树

K 近邻(K-Nearest Neighbors, KNN)是一种基于实例的学习方法。它不显式训练参数,而是在预测时寻找距离测试样本最近的 k 个训练样本,再根据这些邻居投票或平均。

概念详解

分类任务中,KNN 的预测规则为:

y^=argmaxcxiNk(x)I(yi=c)

其中 Nk(x) 表示距离 x 最近的 k 个样本集合。回归任务中通常取均值:

y^=1kxiNk(x)yi

常用欧氏距离为:

d(xi,xj)=m=1d(ximxjm)2

当样本很多时,逐个计算距离代价较高。kd 树通过递归划分空间,加速最近邻搜索。

数学推导

KNN 可以看作在局部区域内估计条件概率:

P(Y=c|X=x)1kxiNk(x)I(yi=c)

因此最大后验分类为:

y^=argmaxcP(Y=c|X=x)

应用代码

python
import numpy as np
from collections import Counter

class KNNClassifier:
    def __init__(self, k=3):
        self.k = k

    def fit(self, X, y):
        self.X = np.asarray(X)
        self.y = np.asarray(y)

    def predict_one(self, x):
        dist = np.sqrt(((self.X - x) ** 2).sum(axis=1))
        idx = np.argsort(dist)[:self.k]
        votes = Counter(self.y[idx])
        return votes.most_common(1)[0][0]

    def predict(self, X):
        return np.array([self.predict_one(x) for x in X])

X = np.array([[0, 0], [1, 1], [1, 0], [5, 5], [6, 5], [5, 6]])
y = np.array([0, 0, 0, 1, 1, 1])

clf = KNNClassifier(k=3)
clf.fit(X, y)
print(clf.predict([[0.8, 0.7], [5.2, 5.4]]))

小结

KNN 简单直观,不需要训练过程,但预测速度和特征尺度非常关键。使用前通常需要标准化特征。