KNN算法实战:从鸢尾花分类到手写数字识别

📅 2026/8/17 18:09:24
KNN算法实战:从鸢尾花分类到手写数字识别
1. 项目概述K最近邻K-Nearest Neighbors简称KNN算法是机器学习领域最基础也最经典的算法之一。作为监督学习中的分类算法KNN以其简单直观、无需训练过程的特性成为无数机器学习初学者的第一个实战案例。今天我们就用两个经典数据集——鸢尾花分类和手写数字识别带大家真正动手实现KNN算法。提示本文假设读者已经了解Python基础语法和机器学习基本概念。如果尚未安装Python环境建议先配置AnacondaJupyter Notebook的开发环境。KNN算法的核心思想可以用一句俗语概括近朱者赤近墨者黑。当我们需要对一个新样本进行分类时只需找到训练集中与它最接近的K个邻居根据这些邻居的类别投票决定新样本的类别。这种懒惰学习Lazy Learning的特性使得KNN算法实现简单但计算复杂度会随着数据规模增大而显著增加。2. 核心原理与数学基础2.1 KNN算法三要素KNN算法的实现主要依赖三个关键要素距离度量常用欧氏距离Euclidean Distance对于二维空间中的两点(x1,y1)和(x2,y2)其距离计算公式为distance sqrt((x2-x1)^2 (y2-y1)^2)对于更高维度的数据公式可自然扩展。在文本分类等场景中也常使用曼哈顿距离或余弦相似度。K值选择K是算法中的超参数表示考虑最近邻的数量。K值过小容易过拟合对噪声敏感K值过大会使分类边界模糊。通常通过交叉验证确定最佳K值。分类决策规则一般采用多数表决法即K个邻居中出现次数最多的类别作为预测结果。也可以根据距离加权投票近距离的邻居拥有更大权重。2.2 算法流程分解一个完整的KNN分类流程包括以下步骤数据准备加载数据集划分训练集和测试集特征标准化对数据进行归一化处理重要距离计算测试样本与所有训练样本的距离排序找邻居按距离升序排列选取前K个投票决策统计K个邻居的类别分布结果输出将得票最多的类别作为预测结果性能评估计算准确率等指标3. 鸢尾花分类实战3.1 数据集介绍鸢尾花数据集Iris是机器学习领域的Hello World包含150个样本每个样本有4个特征花萼长度sepal length花萼宽度sepal width花瓣长度petal length花瓣宽度petal width目标变量是鸢尾花的三个品种Iris SetosaIris VersicolourIris Virginica3.2 代码实现步骤# 导入必要库 from sklearn.datasets import load_iris from sklearn.model_selection import train_test_split from sklearn.preprocessing import StandardScaler from sklearn.neighbors import KNeighborsClassifier from sklearn.metrics import accuracy_score # 加载数据 iris load_iris() X, y iris.data, iris.target # 数据分割 X_train, X_test, y_train, y_test train_test_split(X, y, test_size0.3, random_state42) # 特征标准化 scaler StandardScaler() X_train scaler.fit_transform(X_train) X_test scaler.transform(X_test) # 注意使用训练集的参数转换测试集 # 创建KNN模型 knn KNeighborsClassifier(n_neighbors5) # 训练模型 knn.fit(X_train, y_train) # 预测测试集 y_pred knn.predict(X_test) # 评估准确率 accuracy accuracy_score(y_test, y_pred) print(f模型准确率: {accuracy:.2f})3.3 关键问题与调优特征标准化的重要性不同特征的量纲差异会导致距离计算偏向大数值特征标准化使所有特征具有相同的重要性常用方法Z-score标准化StandardScaler或MinMax缩放K值选择实验 通过交叉验证寻找最佳K值from sklearn.model_selection import cross_val_score k_range range(1, 31) k_scores [] for k in k_range: knn KNeighborsClassifier(n_neighborsk) scores cross_val_score(knn, X_train, y_train, cv5, scoringaccuracy) k_scores.append(scores.mean()) # 绘制K值与准确率关系图 import matplotlib.pyplot as plt plt.plot(k_range, k_scores) plt.xlabel(K值) plt.ylabel(交叉验证准确率) plt.show()可视化决策边界 由于鸢尾花有4个特征我们可以选择两个主要特征进行降维可视化from matplotlib.colors import ListedColormap # 选择前两个特征 X_2d X_train[:, :2] # 创建网格点 h 0.02 # 步长 x_min, x_max X_2d[:, 0].min() - 1, X_2d[:, 0].max() 1 y_min, y_max X_2d[:, 1].min() - 1, X_2d[:, 1].max() 1 xx, yy np.meshgrid(np.arange(x_min, x_max, h), np.arange(y_min, y_max, h)) # 训练仅使用两个特征的KNN knn_2d KNeighborsClassifier(n_neighbors5) knn_2d.fit(X_2d, y_train) # 预测网格点 Z knn_2d.predict(np.c_[xx.ravel(), yy.ravel()]) Z Z.reshape(xx.shape) # 绘制决策边界 cmap_light ListedColormap([#FFAAAA, #AAFFAA, #AAAAFF]) plt.contourf(xx, yy, Z, cmapcmap_light, alpha0.8) # 绘制训练点 plt.scatter(X_2d[:, 0], X_2d[:, 1], cy_train, edgecolork, s20) plt.xlabel(标准化花萼长度) plt.ylabel(标准化花萼宽度) plt.title(KNN决策边界(K5)) plt.show()4. 手写数字识别实战4.1 MNIST数据集简介MNIST数据集包含70,000张手写数字(0-9)的28x28像素灰度图像是计算机视觉领域的经典入门数据集。每个像素点的值范围是0-255表示灰度强度。4.2 完整实现代码from sklearn.datasets import fetch_openml from sklearn.model_selection import train_test_split from sklearn.preprocessing import MinMaxScaler from sklearn.neighbors import KNeighborsClassifier from sklearn.metrics import accuracy_score, confusion_matrix import matplotlib.pyplot as plt import numpy as np # 加载数据 mnist fetch_openml(mnist_784, version1) X, y mnist.data, mnist.target # 数据预览 plt.figure(figsize(10,5)) for i in range(20): plt.subplot(2,10,i1) plt.imshow(X.iloc[i].values.reshape(28,28), cmapgray) plt.title(fLabel: {y[i]}) plt.axis(off) plt.show() # 数据分割 - 使用前10000个样本加速演示 X X[:10000] y y[:10000] X_train, X_test, y_train, y_test train_test_split(X, y, test_size0.2, random_state42) # 特征缩放 scaler MinMaxScaler() X_train scaler.fit_transform(X_train) X_test scaler.transform(X_test) # 创建KNN模型 knn KNeighborsClassifier(n_neighbors5, n_jobs-1) # n_jobs-1使用所有CPU核心 # 训练模型 knn.fit(X_train, y_train) # 预测测试集 y_pred knn.predict(X_test) # 评估模型 accuracy accuracy_score(y_test, y_pred) print(f模型准确率: {accuracy:.2f}) # 混淆矩阵 cm confusion_matrix(y_test, y_pred) plt.figure(figsize(10,8)) plt.imshow(cm, cmapBlues) plt.colorbar() plt.xlabel(预测标签) plt.ylabel(真实标签) plt.title(混淆矩阵) plt.show()4.3 性能优化技巧降维处理原始784维特征(28x28)计算距离耗时严重使用PCA降维保留95%方差from sklearn.decomposition import PCA pca PCA(n_components0.95) # 保留95%方差 X_train_pca pca.fit_transform(X_train) X_test_pca pca.transform(X_test) print(f原始维度: {X_train.shape[1]}) print(f降维后维度: {X_train_pca.shape[1]})近似最近邻算法当数据量很大时使用近似算法加速from sklearn.neighbors import NearestNeighbors # 使用BallTree算法 nn NearestNeighbors(n_neighbors5, algorithmball_tree) nn.fit(X_train_pca) # 查询测试样本的邻居 distances, indices nn.kneighbors(X_test_pca) # 手动实现投票 from collections import Counter y_pred [] for idx in indices: votes y_train.iloc[idx] pred Counter(votes).most_common(1)[0][0] y_pred.append(pred) accuracy accuracy_score(y_test, y_pred) print(f近似KNN准确率: {accuracy:.2f})距离加权投票近距离的邻居应该有更大的投票权重# 自定义权重函数 def inverse_distance(weights): return 1 / (weights 1e-6) # 避免除以零 weighted_knn KNeighborsClassifier(n_neighbors5, weightsinverse_distance) weighted_knn.fit(X_train_pca, y_train) y_pred_weighted weighted_knn.predict(X_test_pca) print(f加权KNN准确率: {accuracy_score(y_test, y_pred_weighted):.2f})5. 常见问题与解决方案5.1 计算效率问题问题表现数据集较大时预测速度很慢内存消耗高解决方案使用KD树或BallTree数据结构加速邻居搜索knn KNeighborsClassifier(algorithmkd_tree) # 或ball_tree对大数据集使用近似最近邻算法如LSH降维处理减少特征数量考虑使用GPU加速库如cuML5.2 类别不平衡问题问题表现某些类别样本数远多于其他类别多数表决法会偏向多数类解决方案使用距离加权投票对多数类进行欠采样或对少数类过采样调整类别权重参数knn KNeighborsClassifier(weightsdistance)5.3 高维灾难问题问题表现特征维度很高时所有样本的距离趋于相似分类性能下降解决方案特征选择去除无关特征使用PCA等降维方法考虑使用更适合高维数据的算法如SVM5.4 参数调优技巧K值选择从Ksqrt(N)开始尝试N为训练样本数使用网格搜索交叉验证from sklearn.model_selection import GridSearchCV param_grid {n_neighbors: range(1, 20)} grid GridSearchCV(KNeighborsClassifier(), param_grid, cv5) grid.fit(X_train, y_train) print(f最佳K值: {grid.best_params_[n_neighbors]})距离度量选择欧氏距离默认适用于连续特征曼哈顿距离对异常值更鲁棒余弦相似度适用于文本数据6. 项目扩展与进阶方向6.1 自定义距离度量在某些特定场景可能需要自定义距离函数。例如对于图像数据可以尝试以下距离# 自定义距离函数示例直方图相交距离 def histogram_intersection(a, b): return np.minimum(a, b).sum() # 使用自定义距离的KNN custom_knn KNeighborsClassifier(n_neighbors5, metrichistogram_intersection) custom_knn.fit(X_train, y_train)6.2 多输出KNNKNN也可以用于多输出任务每个样本有多个目标变量from sklearn.datasets import make_regression from sklearn.neighbors import KNeighborsRegressor # 生成多输出回归数据 X, y make_regression(n_samples1000, n_features10, n_targets2) # 多输出KNN回归 knn_reg KNeighborsRegressor(n_neighbors5) knn_reg.fit(X, y)6.3 在线学习实现标准KNN不支持增量学习但可以通过以下方式实现class OnlineKNN: def __init__(self, k5): self.k k self.X None self.y None def partial_fit(self, X_new, y_new): if self.X is None: self.X X_new self.y y_new else: self.X np.vstack([self.X, X_new]) self.y np.concatenate([self.y, y_new]) def predict(self, X_test): from sklearn.neighbors import NearestNeighbors nn NearestNeighbors(n_neighborsself.k) nn.fit(self.X) distances, indices nn.kneighbors(X_test) predictions [] for idx in indices: votes self.y[idx] pred Counter(votes).most_common(1)[0][0] predictions.append(pred) return np.array(predictions)6.4 与其他算法结合KNN可以与其他算法结合构建更强大的模型KNN特征工程使用KNN提取样本邻居的统计特征作为新特征例如计算每个样本的K个最近邻的类别分布集成学习方法构建多个不同参数的KNN模型进行投票例如使用不同的K值和距离度量from sklearn.ensemble import VotingClassifier knn1 KNeighborsClassifier(n_neighbors5) knn2 KNeighborsClassifier(n_neighbors10, weightsdistance) knn3 KNeighborsClassifier(n_neighbors7, metricmanhattan) ensemble VotingClassifier( estimators[(knn5, knn1), (knn10, knn2), (knn7, knn3)], votinghard) ensemble.fit(X_train, y_train)在实际项目中KNN虽然简单但在特征工程良好、数据规模适中的情况下往往能取得出人意料的好效果。特别是在需要快速验证想法或建立基线模型的场景中KNN因其实现简单、无需复杂调参的优势仍然是机器学习工具箱中的重要成员。