KNN 算法简介
KNN 算法简介
代码地址:https://github.com/ForceInjection/hands-on-ML/blob/main/nju_software/knn.ipynb
一、引言
K 近邻(K-Nearest Neighbors,KNN)是一种基本的机器学习算法,广泛应用于分类和回归任务。它的核心思想是:“物以类聚”,即如果一个数据点在特征空间中与某些已知类别的数据点靠得很近,那么它很可能属于这些数据点所属的类别。
在机器学习的广阔领域中,K 近邻算法以其独特的 “惰性学习” 特质脱颖而出。所谓惰性学习,即算法在训练阶段几乎不进行模型构建,而是将所有数据都 “原样” 保存,等到预测时才根据需要进行计算。这种策略虽看似 “偷懒”,却在实际应用中大放异彩。
从电影推荐系统到手写数字识别,KNN 算法都展现出其强大的通用性。例如:
在电影推荐场景中,通过分析与目标用户兴趣相似的 K个邻居的观影喜好,为用户精准筛选出潜在喜爱的电影;在手写数字识别任务里,基于数字图像的特征,找出训练集中最相近的 K个样本,从而判断新数字的类别。
其核心特点在于无需复杂的训练过程,逻辑直观易懂,即使初学者也能快速上手,为各类数据挖掘任务提供有效解决方案。
二、算法原理
1. 核心思想
KNN 算法的基本假设是,具有相似特征的对象在空间中彼此靠近。简单来说,一个未知类别的样本,其 “邻居” 的类别能在很大程度上反映出自身所属类别。在数学上,通过计算样本 x 的条件概率密度函数来确定其类别归属,公式表达为:
其中, 表示样本 x 的 K 个最近邻样本集合, 是指示函数,若邻居样本 的类别 属于类别 则取 1,否则取 0。通俗地讲,就是统计 K 个邻居中各类别出现的次数,将出现次数最多的类别作为样本 x 的预测类别。
2. 关键三要素
距离度量 :距离度量方式直接决定了样本之间 “相似程度” 的量化。常见的有欧氏距离、曼哈顿距离和余弦相似度,具体对比如下:
度量方式 公式 适用场景 欧氏距离 连续特征 曼哈顿距离 网格路径 余弦相似度 文本数据 欧氏距离衡量的是两点在 n 维空间中的直线距离,适用于具有连续数值特征的数据;曼哈顿距离则是在网格状路径中计算两点之间的绝对距离之和,常用于城市街区等类似场景;余弦相似度重点关注两个向量在方向上的相似性,文本数据因高频词汇的特性,使用余弦相似度能更好地反映文本间语义关联。
K 值选择 :
K值的选取对模型性能至关重要。较小的K值会使模型更贴合数据,但容易受噪声点影响导致过拟合;较大的K值则会使模型更平滑,可能降低模型的复杂度同时提升泛化能力。一般采用肘部法则来确定最佳K值,绘制准确率 -K值曲线,曲线的肘部位置(误差开始趋于稳定的地方)即为较优K值。同时,为避免出现平票情况,K值通常选择奇数。决策规则 :在分类任务中,多数表决法是最常用的决策规则。即统计
K个邻居中每个类别出现的次数,将出现次数最多的类别作为预测类别;也可引入权重机制,依据邻居到样本点的距离给予不同投票权重,距离近的邻居权重更高。在回归任务中,预测值为K个邻居目标值的均值或加权均值,权重同样可基于距离设置。
3. 算法复杂度分析
KNN 算法在 训练阶段几乎没有计算开销,主要计算成本集中在 预测阶段。对于一个包含 个样本,每个样本具有 维特征的数据集,预测时的计算复杂度如下:
计算距离:
需要计算待预测样本与所有训练样本的距离,时间复杂度为:筛选 K 个最近邻:
采用排序或堆结构进行筛选,时间复杂度为:
直接排序:(不推荐) 维护小顶堆:(常见优化方案) 朴素筛选(最坏情况):
因此,整体时间复杂度约为:
在实际应用中,如果数据规模庞大(例如百万级样本),KNN 的计算开销较大,因此常用 KD-Tree、Ball-Tree 或 近似最近邻搜索(ANN) 来优化查询效率。
空间复杂度 主要由存储所有训练样本的数据决定,大小约为:
这意味着数据量较大时,存储成本也是一个需要考虑的因素。
在实际应用中,KNN 适用于 小样本、低维数据集,但对于高维和大规模数据,建议结合索引结构(如 Faiss、HNSW)进行优化。
三、完整示例:鸢尾花分类
1. 环境准备
要运行 KNN 算法进行鸢尾花分类,需先准备好相应的 Python 环境与库。以下为必需的库导入代码:
ounter(lineounter(lineounter(lineounter(lineounter(lineounter(lineounter(lineimport matplotlib.pyplot as pltimport seaborn as snsfrom sklearn.datasets import load_irisfrom sklearn.neighbors import KNeighborsClassifierfrom sklearn.model_selection import train_test_splitfrom sklearn.preprocessing import StandardScalerfrom sklearn.metrics import classification_report
2. 数据预处理
加载鸢尾花数据集,并对其进行预处理,以下是关键代码:
ounter(lineounter(lineounter(lineounter(lineounter(lineounter(lineounter(lineounter(lineounter(line# 加载数据并创建DFiris = load_iris()X, y = iris.data[:, :2], iris.target # 选用前两个特征便于可视化# 标准化处理(演示正确数据分割方法)X_train, X_test, y_train, y_test = train_test_split(X, y, test_size=0.3, random_state=42)scaler = StandardScaler().fit(X_train)X_train = scaler.transform(X_train)X_test = scaler.transform(X_test)
此处选取了鸢尾花数据集的前两个特征,便于后续可视化展示。采用标准差标准化对数据进行处理,消除不同特征量纲对距离计算的影响,确保模型性能的稳定性。
3. 模型训练与评估
通过循环遍历不同的 K 值,寻找最优的 K 值以构建最佳模型,代码如下:
ounter(lineounter(lineounter(lineounter(lineounter(lineounter(lineounter(lineounter(lineounter(lineounter(lineounter(lineounter(line# 动态K值选择演示best_k = 0best_score = 0for k in range(1, 20):knn = KNeighborsClassifier(n_neighbors=k)knn.fit(X_train, y_train)score = knn.score(X_test, y_test)if score > best_score:best_score = scorebest_k = kprint(f'Optimal K: {best_k} (Accuracy: {best_score:.2%})')
该代码片段通过不断调整 K 值,训练模型并计算模型在测试集上的准确率,最终输出最优 K 值及对应的准确率。
4. 决策边界可视化(二维展示)
为了直观地呈现 KNN 模型的决策边界,利用以下代码绘制图形:
ounter(lineounter(lineounter(lineounter(lineounter(lineounter(lineounter(lineounter(lineounter(lineounter(lineounter(lineounter(lineounter(line# 创建网格数据xx, yy = np.meshgrid(np.linspace(-2.5, 2.5, 200), np.linspace(-2.5, 2.5, 200))Z = knn.predict(np.c_[xx.ravel(), yy.ravel()]).reshape(xx.shape)# 绘制图形plt.figure(figsize=(10,6))plt.contourf(xx, yy, Z, alpha=0.4)sns.scatterplot(x=X_train[:,0], y=X_train[:,1], hue=iris.target_names[y_train],palette='Dark2', edgecolor='black')plt.title("KNN Decision Boundaries (k=3)")plt.xlabel(iris.feature_names[0])plt.ylabel(iris.feature_names[1])plt.show()
这段代码生成了一个二维网格,计算网格中每个点的预测类别,绘制出决策边界,并将训练数据点按类别着色绘制其上,清晰地展示了不同类别之间的划分区域。
四、算法优化策略
1. 维度灾难解决方案
随着特征维度的增加,数据在高维空间中变得稀疏,计算距离时会受到大量无关特征的干扰,导致 “维度灾难”。为此,可采用 PCA 降维,保留数据中 90% 的方差,降低特征维度;或利用互信息法进行特征选择,筛选出与目标变量相关性最高的特征,减轻维度对 KNN 算法性能的影响。
2. 效率提升方案
当数据量较大时,为提升 KNN 算法的查询效率,可以采用 KD-Tree 数据结构,其时间复杂度可降至 ,通过构建树形结构实现快速的最近邻搜索;也可使用近似最近邻(LSH 哈希)方法,在牺牲少量精度的前提下,大幅提升大规模数据的查询速度。
3. 样本不平衡处理
在实际数据中,若不同类别样本数量差异悬殊,普通 KNN 算法的预测结果会偏向样本数量多的类别。此时,可通过加权投票的方式,赋予距离近的邻居更高的权重;或采用 SMOTE 过采样技术,合成少数类样本,平衡数据集类别分布,提升模型对少数类的识别能力。
五、总结
KNN 算法凭借其简单直观的特点,在小规模数据、对解释性要求较高的场景下有着广泛的应用,如客户分类、模式识别等。然而,它也并非完美无缺,面对大数据和高维数据时存在一定的局限性。未来,随着深度学习技术的发展,将 KNN 算法与深度特征提取相结合,有望进一步提升其性能。对于学习者而言,从深入研究 KD-Tree 的源码实现入手,是理解 KNN 算法本质、掌握优化技巧的有效路径,这有助于更好地运用和改进 KNN 算法,使其在更广泛的领域中发挥作用。