K 近邻与 kd 树
K 近邻(K-Nearest Neighbors, KNN)是一种基于实例的学习方法。它不显式训练参数,而是在预测时寻找距离测试样本最近的
概念详解
分类任务中,KNN 的预测规则为:
其中
常用欧氏距离为:
当样本很多时,逐个计算距离代价较高。kd 树通过递归划分空间,加速最近邻搜索。
数学推导
KNN 可以看作在局部区域内估计条件概率:
因此最大后验分类为:
应用代码
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 简单直观,不需要训练过程,但预测速度和特征尺度非常关键。使用前通常需要标准化特征。