Scikit Learn - K最近邻
Scikit-Learn - K近邻 (KNN)
Section titled “Scikit-Learn - K近邻 (KNN)”本章探讨 K近邻 (K-Nearest Neighbors, KNN),这是一种基本的基于实例的学习算法,在 Scikit-Learn 中用于分类和回归任务。
基于邻居的学习方法是非参数的,这意味着它们不对底层数据分布做强假设。它们通常被称为“懒惰学习器”(lazy learners),因为它们通常没有一个明确的训练阶段来显式地构建模型。相反,它们存储整个训练数据集,并且仅在预测时才执行计算。
K近邻的核心原理很简单:
- 要分类一个新的数据点,找到与新点距离最近的预定义数量 (
k) 的训练样本。 - 对于分类:将这些
k个邻居中最频繁的类标签分配给新点(多数投票)。 - 对于回归:将这些
k个邻居的目标值的平均值(或中位数)分配给新点。
“距离”通常是欧氏距离 (Euclidean distance),但也可以使用其他度量。k 的选择(邻居的数量)和距离度量是关键的超参数 (hyperparameters)。
sklearn.neighbors 模块
Section titled “sklearn.neighbors 模块”Scikit-Learn 的 sklearn.neighbors 模块提供了无监督(例如,查找最近邻)和有监督(例如,KNN 分类/回归)的基于邻居的学习功能。这些类可以处理 NumPy 数组或 SciPy 稀疏矩阵 (sparse matrices) 作为输入。关键估计器包括 KNeighborsClassifier、KNeighborsRegressor 和 NearestNeighbors(用于无监督邻居搜索)。
最近邻搜索算法
Section titled “最近邻搜索算法”高效地找到最近邻对于 KNN 的性能至关重要,特别是在大型数据集上。Scikit-Learn 的基于邻居的估计器中的 algorithm 参数控制了使用的方法:
暴力搜索 (Brute Force) (algorithm='brute')
Section titled “暴力搜索 (Brute Force) (algorithm='brute')”最直接的方法是计算查询点到训练集中每个点的距离。对于单个查询点和 D 维度的 N 个训练样本,其复杂度为 O(DN)。虽然简单,但对于大型 N,计算成本很高。
**何时使用:**适用于小型数据集,或者当查询数量很少,构建树结构(如 K-D 树或 Ball Tree)的开销不值得时。
K-D 树 (K-D Tree) (algorithm='kd_tree')
Section titled “K-D 树 (K-D Tree) (algorithm='kd_tree')”K 维树是一种二叉树数据结构,它沿着数据轴递归地划分参数空间。这会创建嵌套的正交区域,从而比暴力搜索更快地进行邻居搜索。
优点:
- 构建速度快:沿着轴进行划分相对较快。
- 在低维度下查询高效:对于 D < ~20,平均查询时间可以达到 O(D log N)。
缺点:
- 维度诅咒 (Curse of dimensionality):在处理高维度数据 (D > ~20) 时性能显著下降,在最坏情况下接近 O(DN)。
Ball Tree (algorithm='ball_tree')
Section titled “Ball Tree (algorithm='ball_tree')”Ball Tree 是另一种基于树的数据结构,旨在解决 K-D 树在高维度下的效率低下问题。它递归地将数据划分为由质心 (centroid) 和半径 (radius) 定义的节点,形成嵌套的超球体。它利用三角不等式 (triangle inequality) 来剪枝搜索路径。
优点:
- 更适用于高维度:当 D > ~20 时通常优于 K-D 树。
- 能很好地处理各种距离度量。
缺点:
- 构建成本较高:划分超球体可能比 K-D 树的构建在计算上更密集。
自动选择 (Automatic Selection) (algorithm='auto')
Section titled “自动选择 (Automatic Selection) (algorithm='auto')”这通常是默认设置。Scikit-Learn 会尝试根据输入数据(例如,n_samples、n_features、稀疏性)和 metric 参数,在 fit 方法执行期间确定最佳算法(brute、kd_tree 或 ball_tree)。
选择合适的最近邻算法
Section titled “选择合适的最近邻算法”最佳算法的选择取决于几个因素:
- 样本数量 (N) 和维度 (D):
- **数据结构(稀疏性,内在维度):**基于树的算法 (Ball Tree, K-D Tree) 在稀疏数据或内在维度较低的数据上可能更快。暴力搜索的查询时间不受数据结构影响。
- **邻居数量 (k):**基于树的算法的查询时间可能会随
k的增加而增加。对于单个查询,暴力搜索在距离计算完成后受k的影响较小。 - **查询点数量:**对于少量查询,构建树的开销可能会使暴力搜索具有竞争力,甚至更快。当需要对静态数据集执行大量查询时,树结构更有益。
- **距离度量 (Distance Metric):**与 K-D 树(针对类似欧氏距离进行了优化)相比,Ball Tree 对不同的距离度量通常更具鲁棒性。
通常最好从 algorithm='auto' 开始,让 Scikit-Learn 进行选择。如果性能至关重要,可能需要对不同的算法进行基准测试。
示例:使用 KNeighborsClassifier
Section titled “示例:使用 KNeighborsClassifier”让我们演示如何在 KNeighborsClassifier 中使用这些算法。
from sklearn.datasets import load_irisfrom sklearn.model_selection import train_test_splitfrom sklearn.neighbors import KNeighborsClassifierfrom sklearn.preprocessing import StandardScalerfrom sklearn.pipeline import make_pipelinefrom sklearn.metrics import accuracy_score
# 加载 Iris 数据集iris = load_iris()X, y = iris.data, iris.target
# 分割数据X_train, X_test, y_train, y_test = train_test_split(X, y, test_size=0.3, random_state=42)
# 创建一个 Pipeline: StandardScaler + KNeighborsClassifier# 特征缩放对于像 KNN 这样的基于距离的算法很重要knn_pipeline = make_pipeline( StandardScaler(), KNeighborsClassifier(n_neighbors=5, algorithm='auto', metric='minkowski', p=2) # p=2 表示欧氏距离)
# 拟合模型knn_pipeline.fit(X_train, y_train)
# 进行预测y_pred = knn_pipeline.predict(X_test)
# 评估准确率accuracy = accuracy_score(y_test, y_pred)print(f"KNeighborsClassifier Accuracy: {accuracy:.4f}")
# 您可以尝试将 algorithm 更改为 'ball_tree', 'kd_tree', 或 'brute'# 例如,KNeighborsClassifier(n_neighbors=5, algorithm='kd_tree')# 并观察对于这个小型数据集,在拟合/预测时间上是否有任何显著差异。KNeighborsClassifier Accuracy: 1.0000KNN 的关键考虑因素:
- **选择
k:**较小的k可能导致决策边界不稳定(高方差),而较大的k可能使其过于平滑(高偏差)。通常通过交叉验证 (cross-validation) 来选择k。 - **距离度量 (Distance Metric):**欧氏距离 (Euclidean) 很常用,但根据数据,曼哈顿距离 (Manhattan)、闵可夫斯基距离 (Minkowski) 或余弦相似度 (Cosine similarity) 可能更合适。
- **特征缩放 (Feature Scaling):**至关重要,因为 KNN 依赖于距离。尺度较大的特征会主导距离计算。
- **计算成本:**对于大型数据集,预测可能会很慢,因为它涉及计算到所有(或许多,对于基于树的方法)训练样本的距离。
有关最近邻算法及其用法的更多详细信息,请参阅 Scikit-Learn 最近邻文档。