寻找最近邻
KNN 算法 - 寻找最近邻居
Section titled “KNN 算法 - 寻找最近邻居”K近邻(K-Nearest Neighbors, KNN)简介
Section titled “K近邻(K-Nearest Neighbors, KNN)简介”K近邻(KNN)是一种简单而有效的非参数监督机器学习算法,可用于分类和回归任务。它在分类方面尤其受欢迎。
KNN 的主要特点:
- 惰性学习(Lazy Learning): KNN 被认为是一种“惰性”学习器,因为它在单独的训练阶段不构建显式模型。相反,它存储整个训练数据集。
- 基于实例的学习(Instance-Based Learning): 对新数据点的预测是基于训练集中其最近邻居的特征进行的。
- 非参数(Non-Parametric): 它对底层数据分布不作任何假设(不像,例如,线性回归或高斯朴素贝叶斯)。
KNN 工作原理
Section titled “KNN 工作原理”核心思想是“特征相似性(feature similarity)”或近似度:新数据点被分配到其在特征空间中“K”个最近邻居中最常见的类别(或预测值)。
对新数据点进行分类的过程:
- 选择 K: 选择要考虑的邻居数量(K)。这是一个关键的超参数。
- 计算距离: 计算新数据点与训练数据集中每个点之间的距离。常见的距离度量(distance metrics)包括:
-
- 欧几里得距离(Euclidean Distance): 两点之间的直线距离(对连续变量最常见)。
sqrt(Σ(xᵢ - yᵢ)²)
- 欧几里得距离(Euclidean Distance): 两点之间的直线距离(对连续变量最常见)。
-
- 曼哈顿距离(Manhattan Distance): 它们笛卡尔坐标绝对差的总和。
Σ|xᵢ - yᵢ|
- 曼哈顿距离(Manhattan Distance): 它们笛卡尔坐标绝对差的总和。
-
- 海明距离(Hamming Distance): 用于分类变量(对应符号不同的位置数量)。
-
- 闵可夫斯基距离(Minkowski Distance): 欧几里得距离和曼哈顿距离的泛化。
- 确定邻居: 找到与新点计算距离最小的 K 个训练数据点(K 个最近邻居)。
- 多数投票(分类)(Majority Vote (Classification)): 确定这 K 个邻居中最常见的类别标签。将此多数类别标签分配给新数据点。
- 平均(回归)(Average (Regression)): 计算 K 个最近邻居的目标值的平均值(或中位数)。将此平均值作为新数据点的预测值。
选择合适的 K
Section titled “选择合适的 K”K 值显著影响模型的性能:
- 小的 K(例如,K=1): 模型对噪声和异常值高度敏感。决策边界可能复杂且不规则(高方差 variance,低偏差 bias)。容易过拟合。
- 大的 K: 模型变得更平滑,对噪声不那么敏感。决策边界变得更简单。然而,如果 K 过大,它可能会包含来自其他类别的邻居,可能导致点被错误分类(高偏差 bias,低方差 variance)。容易欠拟合(underfitting)。
- 选择 K: K 通常通过使用交叉验证(cross-validation)的超参数调优来选择。二元分类(binary classification)通常更喜欢奇数值的 K,以避免投票出现平局。
特征缩放(Feature Scaling)的重要性
Section titled “特征缩放(Feature Scaling)的重要性”由于 KNN 依赖于距离计算,值范围较大的特征可能会不成比例地影响距离度量。因此,在应用 KNN 之前进行特征缩放(feature scaling)(例如,使用 StandardScaler 进行标准化 Standardization 或使用 MinMaxScaler 进行归一化 Normalization)至关重要。
在 Python 中的实现(Scikit-learn)
Section titled “在 Python 中的实现(Scikit-learn)”让我们使用 Iris 数据集实现 KNN 进行分类。
KNN 作为分类器示例
Section titled “KNN 作为分类器示例”import numpy as npimport pandas as pdimport matplotlib.pyplot as pltimport seaborn as snsfrom sklearn.datasets import load_irisfrom sklearn.model_selection import train_test_split, cross_val_score, GridSearchCVfrom sklearn.preprocessing import StandardScalerfrom sklearn.neighbors import KNeighborsClassifierfrom sklearn.pipeline import Pipelinefrom sklearn.metrics import classification_report, confusion_matrix, accuracy_score
# 1. Load Iris datasetiris = load_iris()X = iris.datay = iris.targetfeature_names = iris.feature_namesclass_names = iris.target_names
# 2. Split data into training and test setsX_train, X_test, y_train, y_test = train_test_split(X, y, test_size=0.3, random_state=42, stratify=y)
# 3. Create a pipeline: StandardScaler + KNeighborsClassifier# Pipeline ensures scaling is applied correctly during cross-validation and predictionpipe = Pipeline([ ('scaler', StandardScaler()), ('knn', KNeighborsClassifier())])
# 4. --- Hyperparameter Tuning (Finding the best K) ---# Define a range of K values to testparam_grid = {'knn__n_neighbors': range(1, 26)} # Test K from 1 to 25
# Use GridSearchCV for exhaustive search with cross-validation (e.g., 5-fold)grid_search = GridSearchCV(pipe, param_grid, cv=5, scoring='accuracy')grid_search.fit(X_train, y_train)
# Best K foundbest_k = grid_search.best_params_['knn__n_neighbors']print(f"Best K found by GridSearchCV: {best_k}")print(f"Best cross-validation accuracy: {grid_search.best_score_:.4f}")
# 5. Train the final model using the best K on the full training datafinal_model = Pipeline([ ('scaler', StandardScaler()), ('knn', KNeighborsClassifier(n_neighbors=best_k))])final_model.fit(X_train, y_train)
# 6. Make predictions on the test sety_pred = final_model.predict(X_test)
# 7. Evaluate the final modelprint("\n--- Final Model Evaluation on Test Set ---")conf_mat = confusion_matrix(y_test, y_pred)acc = accuracy_score(y_test, y_pred)class_rep = classification_report(y_test, y_pred, target_names=class_names)
print(f"Test Accuracy: {acc:.4f}")print("\nConfusion Matrix:")# Using seaborn for a nicer heatmapplt.figure(figsize=(6, 4))sns.heatmap(conf_mat, annot=True, fmt='d', cmap='Blues', xticklabels=class_names, yticklabels=class_names)plt.xlabel('Predicted Label')plt.ylabel('True Label')plt.title('Confusion Matrix')plt.show()
print("\nClassification Report:")print(class_rep)输出(示例)
Section titled “输出(示例)”Best K found by GridSearchCV: 3Best cross-validation accuracy: 0.9619
--- Final Model Evaluation on Test Set ---Test Accuracy: 0.9778
Confusion Matrix:(Seaborn heatmap plot showing the confusion matrix)
Classification Report: precision recall f1-score support
setosa 1.00 1.00 1.00 15 versicolor 1.00 0.93 0.97 15 virginica 0.94 1.00 0.97 15
accuracy 0.98 45 macro avg 0.98 0.98 0.98 45weighted avg 0.98 0.98 0.98 45此示例演示了加载数据、分割数据、使用 Pipeline 进行缩放和分类、使用 GridSearchCV 和交叉验证找到最优 K 值、使用最佳 K 值训练最终模型以及评估其在未见过测试集上的性能。结果显示,对于 Iris 数据集,K=3 时获得了高准确率。
KNN 作为回归器示例
Section titled “KNN 作为回归器示例”KNN 也可以预测连续值。这里是使用 KNeighborsRegressor 在 Diabetes 数据集上的一个简短示例。
import numpy as npfrom sklearn.datasets import load_diabetesfrom sklearn.model_selection import train_test_split, GridSearchCVfrom sklearn.preprocessing import StandardScalerfrom sklearn.neighbors import KNeighborsRegressorfrom sklearn.pipeline import Pipelinefrom sklearn.metrics import mean_squared_error, r2_score
# 1. Load Diabetes dataset (regression task)diabetes = load_diabetes()X, y = diabetes.data, diabetes.target
# 2. Split dataX_train, X_test, y_train, y_test = train_test_split(X, y, test_size=0.3, random_state=42)
# 3. Create pipelinepipe_reg = Pipeline([ ('scaler', StandardScaler()), ('knn_reg', KNeighborsRegressor())])
# 4. Tune Kparam_grid_reg = {'knn_reg__n_neighbors': range(1, 31)} # Test K from 1 to 30grid_search_reg = GridSearchCV(pipe_reg, param_grid_reg, cv=5, scoring='neg_mean_squared_error')grid_search_reg.fit(X_train, y_train)
best_k_reg = grid_search_reg.best_params_['knn_reg__n_neighbors']print(f"Best K for Regression: {best_k_reg}")print(f"Best CV Negative MSE: {grid_search_reg.best_score_:.2f}")
# 5. Final Model Training & Evaluationfinal_model_reg = Pipeline([ ('scaler', StandardScaler()), ('knn_reg', KNeighborsRegressor(n_neighbors=best_k_reg))])final_model_reg.fit(X_train, y_train)y_pred_reg = final_model_reg.predict(X_test)
mse = mean_squared_error(y_test, y_pred_reg)r2 = r2_score(y_test, y_pred_reg)
print("\n--- Regression Model Evaluation on Test Set ---")print(f"Mean Squared Error (MSE): {mse:.2f}")print(f"R-squared (R2): {r2:.4f}")输出(示例)
Section titled “输出(示例)”Best K for Regression: 11Best CV Negative MSE: -3380.73
--- Regression Model Evaluation on Test Set ---Mean Squared Error (MSE): 3027.05R-squared (R2): 0.4529这显示了 KNN 用于回归,预测连续目标值。通过交叉验证找到了最佳 K 值,并使用 MSE 和 R-squared 评估了模型的性能。
KNN 的优缺点
Section titled “KNN 的优缺点”- 简单直观: 易于理解和实现。
- 无训练阶段: 瞬间学习(只需存储数据),因此“训练”速度快。
- 易于适应: 可以轻松适应新的训练数据,无需重新训练复杂模型。
- 非参数: 对数据分布不作任何假设。
- 多用途: 可用于分类和回归。
- 预测计算成本高昂: 对于每个预测都要计算与所有训练点的距离可能很慢,尤其是在大型数据集上(可以使用 BallTree 或 KDTree 等算法缓解,Scikit-learn 通常默认使用这些算法)。
- 内存使用率高: 需要存储整个训练数据集。
- 对 K 敏感: 性能高度依赖于 K 值的选择。
- 对特征缩放敏感: 需要对特征进行适当的缩放。
- 对无关特征敏感(“维度灾难”): 在高维空间中,随着距离变得不那么有意义,以及无关特征可能主导距离计算,性能会下降(curse of dimensionality)。
KNN 的应用
Section titled “KNN 的应用”KNN 常用于:
- 推荐系统: 寻找相似的物品或用户。
- 图像识别: 根据与已知图像的相似性对图像进行分类。
- 异常检测: 识别远离其邻居的点。
- 金融建模: 根据相似的历史模式预测信用评级或股票价格。
- 医学诊断: 基于相似病例协助诊断(尽管需要仔细验证)。
- Scikit-learn 最近邻居文档: https://scikit-learn.org/stable/modules/neighbors.html
- StatQuest: K近邻解释: https://statquest.org/video-index/ (搜索 KNN)