【机器学习|DAY03】K近邻算法(KNN)笔记

文章目录

  • K近邻(KNN)
    • [1. KNN算法核心思想与步骤](#1. KNN算法核心思想与步骤)
      • [1.1 算法思想](#1.1 算法思想)
      • [1.2 具体算法步骤](#1.2 具体算法步骤)
    • [2. 相似性度量:KNN的"距离"是什么?](#2. 相似性度量:KNN的“距离”是什么?)
      • [2.1 欧氏距离(Euclidean Distance)](#2.1 欧氏距离(Euclidean Distance))
      • [2.2 曼哈顿距离(Manhattan Distance)](#2.2 曼哈顿距离(Manhattan Distance))
      • [2.3 切比雪夫距离(Chebyshev Distance)](#2.3 切比雪夫距离(Chebyshev Distance))
      • [2.4 闵可夫斯基距离(Minkowski Distance)](#2.4 闵可夫斯基距离(Minkowski Distance))
      • [2.5 距离对比与选择](#2.5 距离对比与选择)
    • [3. 高效检索:KD树加速最近邻搜索](#3. 高效检索:KD树加速最近邻搜索)
      • [3.1 为什么需要KD树?](#3.1 为什么需要KD树?)
      • [3.2 KD树的构建:用平面反复切割空间](#3.2 KD树的构建:用平面反复切割空间)
      • [3.3 KD树的搜索过程:回溯与"画圆"](#3.3 KD树的搜索过程:回溯与“画圆”)
      • [3.4 sklearn中的算法选择](#3.4 sklearn中的算法选择)
    • [4. K值的选择:过拟合与欠拟合的权衡](#4. K值的选择:过拟合与欠拟合的权衡)
      • [4.1 K值过小 : 过拟合](#4.1 K值过小 : 过拟合)
      • [4.2 K值过大:欠拟合](#4.2 K值过大:欠拟合)
      • [4.3 如何找到合适的K?](#4.3 如何找到合适的K?)
    • [5. sklearn API 实战](#5. sklearn API 实战)
      • [5.1 主要类与模块](#5.1 主要类与模块)
      • [5.2 关键参数详解](#5.2 关键参数详解)
      • [5.3 代码示例](#5.3 代码示例)
    • [6. KNN算法优缺点总结](#6. KNN算法优缺点总结)
    • [7. 机器学习知识补充:超参数的确定](#7. 机器学习知识补充:超参数的确定)
      • [7.1 交叉验证(Cross-Validation)](#7.1 交叉验证(Cross-Validation))
      • [7.2 网格搜索(Grid Search)](#7.2 网格搜索(Grid Search))
      • [7.3 sklearn实现:GridSearchCV](#7.3 sklearn实现:GridSearchCV)

K近邻(KNN)


1. KNN算法核心思想与步骤

1.1 算法思想

KNN(K-Nearest Neighbors)是一种直观且经典的监督学习算法 。它的核心思想简单来说:一个样本的标签,应该和它在特征空间里最相似的 K 个邻居的标签一致

1.2 具体算法步骤

  1. 计算距离:计算待预测样本与训练集中每一个样本的相似性(通常用距离来衡量)。
  2. 排序取K:将计算出的所有距离递增排序,选出相似度最高(距离最小)的 K 个样本。
  3. 投票/平均
    • 分类问题:统计这 K 个样本中每个类别的出现次数,取出现次数最多的类别作为预测结果(多数表决)。
    • 回归问题:计算这 K 个样本目标值的平均值,作为预测结果。
  4. 输出结果

疑惑:KNN在预测时要依赖训练集,那它算不算监督学习?

答案是是的 。 KNN在训练阶段实际上只是"记住"了整个训练集(没有显式的模型训练过程),当新样本到来时,才去寻找训练集中的近邻。它必须使用带标签的训练数据,因此是典型的监督学习。预测时并不需要测试集的真实标签,只是利用训练集的标签来推断。


2. 相似性度量:KNN的"距离"是什么?

"距离"是KNN判断相似性的核心。在原理上,距离公式是可以自定义的,但在scikit-learn中默认使用欧氏距离,也可以通过修改参数选择其他距离。下面介绍几种常用的距离公式及其直观理解。

2.1 欧氏距离(Euclidean Distance)

公式

d = ∑ i = 1 n ( x i − y i ) 2 d = \sqrt{\sum_{i=1}^{n}(x_i - y_i)^2} d=i=1∑n(xi−yi)2

通俗理解:就是我们最熟悉的"两点之间的直线距离"。可以沿着任意方向移动,所以我们可以理解为朝着目标点的方向移动的距离。

2.2 曼哈顿距离(Manhattan Distance)

公式

d = ∑ i = 1 n ∣ x i − y i ∣ d = \sum_{i=1}^{n}|x_i - y_i| d=i=1∑n∣xi−yi∣

通俗理解 :好比你在城市里开车,只能沿着纵横的街道走,不能斜穿街区。也就是说可以沿着x方向和y方向移动,所以从一点到另一点必须走的水平方向距离加上垂直方向距离之和。

来源 :得名于曼哈顿的街区网格布局,适用于特征相互独立、各维度差异需同等对待的场景。

2.3 切比雪夫距离(Chebyshev Distance)

公式

d = max ⁡ i ( ∣ x i − y i ∣ ) d = \max_i(|x_i - y_i|) d=imax(∣xi−yi∣)

通俗理解 :假设你是一个国际象棋里的王,每次可以向八个方向(上下左右、四个斜角)移动一格。也就是说可以沿着八个方向移动,从一点到另一点的距离就取决于各坐标差值中最大的那个。因为在你移动的过程中,斜向移动可以同时减少两个坐标的差值,所以步数由最大的差距决定,较小差距早就被"顺路"处理了。

来源 :常用于网格、物流调度等场景,衡量某一维度的最大偏差。

2.4 闵可夫斯基距离(Minkowski Distance)

公式

d = ( ∑ i = 1 n ∣ x i − y i ∣ p ) 1 / p d = \left( \sum_{i=1}^{n}|x_i - y_i|^p \right)^{1/p} d=(i=1∑n∣xi−yi∣p)1/p

这是上面几种距离的统一定义形式。当 (p=1) 时,就是曼哈顿距离;(p=2) 时是欧氏距离;(p \to \infty) 时就是切比雪夫距离。

2.5 距离对比与选择

距离类型 特点 适用场景
欧氏距离 直线距离,各方向等权重 大多数连续特征数据集
曼哈顿距离 坐标轴方向累加,对异常值较不敏感 特征间差异需要独立累加时,如文档分析
切比雪夫距离 只看最大差值,忽略其他维度微小变动 强调极端维度差异的场合
闵可夫斯基距离 可以通过p调节侧重 需要灵活调整距离度量时

在sklearn中,通过 metric 参数即可自由切换,同时也可以传入自定义的距离函数。


3. 高效检索:KD树加速最近邻搜索

如果每预测一个点都要和全部训练样本算一次距离,当数据量很大时会非常慢。KD树(K-Dimensional Tree)正是为了加速这一过程而生的数据结构

3.1 为什么需要KD树?

暴力计算的时间复杂度为 O(n×d),n为样本数,d为维度。KD树通过提前将空间分层划分,可以将最近邻搜索的平均复杂度降到 O(log n)。它本质上是一种"空间换时间"的策略。

3.2 KD树的构建:用平面反复切割空间

  1. 在R维空间中,每次选择一个维度(例如先从第1维开始)。
  2. 找到当前样本在该维度上的中位值 ,用一个垂直于该坐标轴的超平面将空间一分为二。
  3. 左子节点存放该维度值 ≤ 中位值的样本,右子节点存放 > 中位值的样本。
  4. 切换到下一维度,对两个子空间递归执行同样的分割,直到每个叶子节点只包含少数样本(或达到预定深度)。

形象比喻:像切蛋糕,第一刀竖直切成左右两块,接着把左块横着切,右块也横着切......每一刀都是在一个维度上取中间点,不断将空间划分为更小的超矩形。这样树中的每个非叶节点就是一个"切割点",叶子节点就是不能再分的"小空间"。

复制代码
         [按维度1分割]
        /            \
  (左空间)        (右空间)
   [按维度2分割]    [按维度2分割]
    /     \          /     \
叶子A  叶子B     叶子C  叶子D

例子理解:使用二维数据点(每层切换划分维度,第一层按 x,第二层按 y,以此类推):

( 2 , 3 ) , ( 5 , 4 ) , ( 9 , 6 ) , ( 4 , 7 ) , ( 8 , 1 ) , ( 7 , 2 ) (2,3),\ (5,4),\ (9,6),\ (4,7),\ (8,1),\ (7,2) (2,3), (5,4), (9,6), (4,7), (8,1), (7,2)

  • 根节点 (7,2),按 x 划分,分割值 7。
  • 左子树节点 (5,4),按 y 划分,分割值 4;其左子 (2,3),右子 (4,7)
  • 右子树节点 (9,6),按 y 划分,分割值 6;其左子 (8,1)

3.3 KD树的搜索过程:回溯与"画圆"

当你查询一个测试点时,先像插入一样从根节点沿分支走到对应的叶子节点 ,把这个叶子里的点当作"当前最近邻"。

最近邻不一定就在该叶子空间内 !于是我们执行回溯,向上回溯父节点(分割点),对路径上的每一个节点,执行以下操作:

  • 计算与该节点所存数据点的距离。若该距离小于当前最近距离,则更新最近点和最近距离(r) 。
  • 当前节点在维度 d上把空间一分为二,Q 落在一侧,另一侧未被搜索。所以以查询点为圆心、(r) 为半径画一个超球面
    • 检查该超球面是否与回溯路径上的分割超平面相交。如果相交,说明另一侧分支中可能存在更近的点,则必须进入该分支重新搜索。
    • 若不相交,则另一侧不可能有更近的点,直接剪枝,不搜索另一侧。
  • 继续向上回溯,直到根节点处理完毕,搜索结束。最终得到的最近点即为查询点的最近邻

为什么这样能找到真实最近邻? 因为如果超球面与某分割面相交,就代表可能存在落在分割面另一侧、但离查询点比 (r) 更近的点;若不相交,则该侧所有点距离一定大于 (r),可以直接剪枝。这种机制保证了查找效率。

例子理解:设查询点 Q = ( 2 , 4.5 ) Q = (2, 4.5) Q=(2,4.5),搜索其最近邻。

第一步:向下到达叶节点

  • 从根 (7,2) x 开始: Q . x = 2 < 7 Q.x = 2 < 7 Q.x=2<7 → 进入左子树,到达 (5,4) y
  • (5,4) 按 y 划分: Q . y = 4.5 > 4 Q.y = 4.5 > 4 Q.y=4.5>4 → 进入右子树,到达叶节点 (4,7)
  • (4,7) 设为当前最近点,距离:
    d = ( 2 − 4 ) 2 + ( 4.5 − 7 ) 2 = 4 + 6.25 ≈ 3.201 d = \sqrt{(2-4)^2 + (4.5-7)^2} = \sqrt{4 + 6.25} \approx 3.201 d=(2−4)2+(4.5−7)2 =4+6.25 ≈3.201

第二步:回溯到 (5,4)

  • 计算 Q Q Q 到 (5,4) 的距离:
    d = ( 2 − 5 ) 2 + ( 4.5 − 4 ) 2 = 9 + 0.25 ≈ 3.041 d = \sqrt{(2-5)^2 + (4.5-4)^2} = \sqrt{9 + 0.25} \approx 3.041 d=(2−5)2+(4.5−4)2 =9+0.25 ≈3.041
    3.041 < 3.201 3.041 < 3.201 3.041<3.201,更新最近点为 (5,4),最近距离 ≈ 3.041。
  • 判断另一子树(即 (5,4) 的左子树,包含 (2,3)):
    (5,4) 的分割维度是 y,分割值 4。 Q Q Q 在分割面的右侧(y=4.5 > 4),另一侧是 y < 4 的半空间。
    到分割面的距离 = ∣ 4.5 − 4 ∣ = 0.5 |4.5 - 4| = 0.5 ∣4.5−4∣=0.5。
    0.5 < 3.041 0.5 < 3.041 0.5<3.041,说明另一侧可能存在更近的点,必须搜索
    进入 (5,4) 的左子节点 (2,3)

第三步:搜索 (2,3) 并返回

  • (2,3) 为叶节点,计算 Q Q Q 到它的距离:
    d = ( 2 − 2 ) 2 + ( 4.5 − 3 ) 2 = 1.5 d = \sqrt{(2-2)^2 + (4.5-3)^2} = 1.5 d=(2−2)2+(4.5−3)2 =1.5
    1.5 < 3.041 1.5 < 3.041 1.5<3.041,更新最近点为 (2,3),最近距离 = 1.5。
  • 回溯至 (5,4)(已处理完毕)。

第四步:回溯到根 (7,2)

  • 计算 Q Q Q 到 (7,2) 的距离:
    d = ( 2 − 7 ) 2 + ( 4.5 − 2 ) 2 = 25 + 6.25 ≈ 5.590 d = \sqrt{(2-7)^2 + (4.5-2)^2} = \sqrt{25 + 6.25} \approx 5.590 d=(2−7)2+(4.5−2)2 =25+6.25 ≈5.590
    5.590 > 1.5 5.590 > 1.5 5.590>1.5,不更新最近点。
  • 判断另一子树(根节点的右子树,包含 (9,6)(8,1)):
    根按 x 划分,分割值 7。 Q . x = 2 Q.x=2 Q.x=2,在分割面左侧,另一侧是 x > 7 的半空间。
    到分割面的距离 = ∣ 2 − 7 ∣ = 5 |2 - 7| = 5 ∣2−7∣=5。
    5 ≥ 1.5 5 \ge 1.5 5≥1.5,另一侧不可能有比 1.5 更近的点,剪枝,不搜索右子树。

第五步:回溯结束

  • 当前最近点为 (2,3),最近距离 1.5,即为最终结果。

3.4 sklearn中的算法选择

在sklearn的KNN中,通过参数 algorithm 控制搜索算法:

  • 'auto':自动根据数据选择最合适的算法。
  • 'ball_tree':球树,高维数据下效果更好。
  • 'kd_tree':KD树,适合低维数据(一般<20维)。
  • 'brute':暴力搜索,直接计算所有距离。

通常保留 'auto' 即可,它会智能选择。可以看到,距离和排序实现并不是固定的 ,我们可以通过 metricalgorithm 参数灵活调整。


4. K值的选择:过拟合与欠拟合的权衡

K是KNN里最重要的超参数,它直接控制着模型的复杂度。理解它对偏差与方差的影响是关键。

4.1 K值过小 : 过拟合

当 K=1 时,预测只依赖最近的那个邻居。这意味着:

  • 模型变得非常复杂,完全贴合训练数据。
  • 对噪声点和异常点极度敏感,每一个孤立样本都可能形成小的"决策孤岛"。
  • 训练集上表现好,测试集上表现差 → 过拟合

4.2 K值过大:欠拟合

当 K 接近训练集总样本数时,分类将总是预测频率最高的类别,回归就是全局平均值。

  • 模型过于简单,很多无关的远距离样本也参与投票/平均,导致决策边界过于平滑。
  • 学习的"特征"太少,无法捕捉数据的真实结构
  • 训练集和测试集表现都差 → 欠拟合

"K过大时,因为加了远距离的样本作平均,难道不是引入了噪声让模型复杂吗?"

实际上,这些远距离点把预测值向全局均值拉扯,抑制了决策边界的复杂性,使得模型的偏差增大 ,整体表达能力下降。噪声虽然进入了投票,但其效果是被平均掉的,更多表现为欠拟合而非过拟合。过拟合的根源在于模型对训练集中的细微信号和噪声过于忠实,K越小越忠实;K越大越"糊化"信息,导致欠拟合。

4.3 如何找到合适的K?

选择K没有绝对公式,一般通过交叉验证网格搜索来确定(详见第7节)。常见经验:

  • K一般选奇数(避免投票平局)。
  • 从较小的K开始尝试,观察验证曲线。
  • 数据量越大,K可相对选大一些;数据维度高时则要小心维度灾难。

5. sklearn API 实战

5.1 主要类与模块

  • sklearn.neighbors.KNeighborsClassifier:用于分类。
  • sklearn.neighbors.KNeighborsRegressor:用于回归。

5.2 关键参数详解

以分类器为例:

python 复制代码
KNeighborsClassifier(
    n_neighbors=5,          # K值,默认5
    weights='uniform',      # 'uniform': 所有邻居权重相同;'distance': 距离倒数加权(更近的邻居影响更大)
    algorithm='auto',       # 搜索算法:'ball_tree', 'kd_tree', 'brute', 'auto'
    metric='minkowski',     # 距离度量,默认闵可夫斯基(p=2即欧氏距离)
    p=2,                    # 闵可夫斯基距离的参数p
    n_jobs=None             # 并行搜索的CPU核数
)

常用方法

  • fit(X, y):训练模型,KNN在这里实际上只是存储数据。
  • predict(X):对测试样本进行预测。
  • kneighbors(X, n_neighbors=None, return_distance=True):返回每个测试样本的K个最近邻居的距离及索引。
  • score(X, y):返回预测的准确率(分类)或R²(回归)。

5.3 代码示例

python 复制代码
from sklearn.neighbors import KNeighborsClassifier
from sklearn.datasets import load_iris
from sklearn.model_selection import train_test_split

# 数据准备
iris = load_iris()
X_train, X_test, y_train, y_test = train_test_split(
    iris.data, iris.target, test_size=0.2, random_state=42
)

# 建立KNN分类器,K=3,使用距离加权
knn = KNeighborsClassifier(n_neighbors=3, weights='distance')
knn.fit(X_train, y_train)               # 训练(存储数据)
y_pred = knn.predict(X_test)            # 预测
accuracy = knn.score(X_test, y_test)    # 评估
print(f"Test accuracy: {accuracy:.3f}")

6. KNN算法优缺点总结

优点 原因
简单直观,易于理解实现 原理完全基于距离和投票,无需复杂数学推导
无需训练过程(懒惰学习) 只需存储数据,新数据到来时即时计算
天然支持多分类 多数表决机制自然支持任意类别数
能处理非线性问题 决策边界可以是任意复杂形状,不假设数据分布
对异常值不敏感(K较大时) K值增大时,个别离群点被多数邻居稀释
缺点 原因
计算开销大,预测慢 每次预测都需要遍历或搜索大量训练样本
内存消耗高 需要存储全部训练数据,数据量大时不适用
维度灾难 高维空间中距离区分度下降,几乎所有点都差不多远
对特征缩放敏感 距离计算受量纲影响大,必须做标准化或归一化
样本不平衡问题 多数类容易主导投票结果,需要适当加权或重采样

7. 机器学习知识补充:超参数的确定

K就是一个典型的超参数,需要人为设定。如何系统化地选出最优K?这就要用到交叉验证和网格搜索。

7.1 交叉验证(Cross-Validation)

思想 :将训练数据分成n份,依次用其中1份作为验证集,其余n-1份作为训练集,重复n次训练和验证,得到n个评估分数,最终取平均作为模型的性能指标。这就是n折交叉验证

为什么要这样做?

  • 单次划分的训练/验证集可能带有偶然性,交叉验证通过多次"换着来"使得评估结果更加稳定、可信
  • 能有效检测模型是否稳定,防止因某次数据分割的"运气"而得到误导性的好成绩。
  • 单独使用交叉验证也是有意义的,它可以用来评估给定超参数下模型的泛化能力,而不仅仅是服务于网格搜索。

7.2 网格搜索(Grid Search)

当有多个超参数需要调节时,手动尝试每一种组合非常繁琐。网格搜索的做法是:

  1. 为每个超参数指定一个候选值的列表。
  2. 生成所有参数组合的"网格"。
  3. 对每一组组合,使用交叉验证评估模型表现。
  4. 选出交叉验证平均得分最高的那一组参数作为最终选择。

比如我们为KNN设定候选 n_neighbors=[3,5,7,9]weights=['uniform','distance'],网格搜索会遍历这 4×2=8 种组合,找出最优搭配。

7.3 sklearn实现:GridSearchCV

python 复制代码
from sklearn.model_selection import GridSearchCV
from sklearn.neighbors import KNeighborsClassifier

# 定义KNN模型
knn = KNeighborsClassifier()

# 超参数网格
param_grid = {
    'n_neighbors': [3, 5, 7, 9, 11],
    'weights': ['uniform', 'distance'],
    'metric': ['euclidean', 'manhattan']
}

# 5折交叉验证的网格搜索
grid_search = GridSearchCV(knn, param_grid, cv=5, scoring='accuracy')
grid_search.fit(X_train, y_train)

print("最佳参数组合:", grid_search.best_params_)
print("最佳交叉验证分数:", grid_search.best_score_)

# 使用最佳模型预测
best_knn = grid_search.best_estimator_
test_score = best_knn.score(X_test, y_test)
print("测试集准确率:", test_score)

常用参数说明

  • estimator:待调参的模型。
  • param_grid参数字典,键为参数名,值为要尝试的列表
  • cv:交叉验证折数(默认为5)。
  • scoring:评估指标,如 'accuracy''f1' 等。
  • 训练完后,best_params_best_score_ 可查看最优结果。

希望这篇笔记能帮你理清 KNN。KNN作为最"佛系"的机器学习算法,不训练模型,靠"邻里关系"做决策。理解它的距离度量、KD树加速原理、K值对偏差方差的影响,以及如何用交叉验证和网格搜索调优。

以上为个人学习总结,旨在梳理个人理解。如有疏漏或不当之处,欢迎指正与交流。如果文章对你有帮助,别忘了点个赞、留个言,让更多的小伙伴看到~ 我们下篇再见!

相关推荐
圣光SG11 小时前
Servlet学习笔记
笔记·学习·servlet
奋发向前wcx11 小时前
y1,y2总复习笔记5 2026.7.19
数据结构·笔记·算法
@Mike@13 小时前
02-数据库学习笔记(SQL引擎)
数据库·笔记·学习
六点_dn13 小时前
RabbitMQ学习笔记-定义与作用
笔记·学习·rabbitmq
人生百态,人生如梦15 小时前
情感交互仿生人从技术到落地构想3——技术交流贴(2026.7)
人工智能·机器学习·人机交互·交互·具身智能
香辣牛肉饭16 小时前
【算法】动态规划 最长公共子序列(LCS)
经验分享·笔记·算法·动态规划
迷途呀16 小时前
Python:函数中的参数类型
开发语言·笔记·python·langchain·nlp
让学习成为一种生活方式16 小时前
机器学习酶功能预测--Nature
人工智能·机器学习
依然范特东17 小时前
动手学深度学习笔记--卷积层、微调
人工智能·笔记·深度学习