【算法笔记】KNN(K近邻算法)¶
约 2093 个字 58 行代码 1 张图片 预计阅读时间 8 分钟
KNN(K-Nearest Neighbors,K近邻算法)是机器学习领域入门级的经典算法,因其原理直观、实现简单,常作为监督学习的入门案例。它不依赖复杂的数学建模,核心靠“相似性投票”决策,既能处理分类问题,也能应对回归任务,在推荐系统、图像识别、文本分类等场景中都有广泛应用。本文将从原理、流程、参数、特点到实战要点,全面拆解KNN算法。
一、算法核心概述¶
KNN属于监督学习算法,同时具备“惰性学习”和“非参数模型”的属性,这两个特性是理解它的关键:
-
惰性学习(Lazy Learning):与逻辑回归、决策树等“急切学习”算法不同,KNN在训练阶段不做任何模型训练和参数拟合,仅存储训练集数据。直到有新样本需要预测时,才开始计算距离、找近邻、做决策,相当于“临阵磨枪”式学习。
-
非参数模型(Non-parametric Model):无需假设数据服从某种固定分布(如正态分布),也不依赖预设的参数结构,完全由数据本身决定模型形态,对复杂数据的适应性更强。
其核心假设是:相似的样本在特征空间中会彼此靠近。就像生活中“物以类聚,人以群分”,通过新样本周围的“邻居”特征,就能推断出它的类别或数值。
二、核心原理与执行流程¶
1. 核心思想¶
对于待预测的新样本,先计算它与训练集中所有样本的“距离”(衡量特征相似度),然后筛选出距离最近的K个样本(即“K近邻”),最后根据这K个近邻的标签(分类任务)或数值(回归任务),通过特定规则得出新样本的预测结果。
2. 完整执行步骤¶
KNN的执行流程清晰易懂,可拆解为三步,每一步都有关键细节需要注意:
-
计算距离(相似度度量):距离是KNN的核心度量指标,距离越小表示样本越相似。常用的距离计算方式有两种: 注意:计算距离前需对特征做标准化(如归一化、标准化),避免数值范围大的特征主导距离计算结果(例如“身高(cm)”和“体重(kg)”,若不标准化,身高的数值差异会掩盖体重的影响)。
- 欧式距离(Euclidean Distance):最常用的距离度量,适用于连续型特征,计算两点在特征空间中的直线距离,公式为:
\[d(x,y)=\sqrt{\sum_{i=1}^{n}(x_i - y_i)^2}\](n为特征维度,\(x_i\)、\(y_i\)分别为两个样本的第i个特征值)。
- 曼哈顿距离(Manhattan Distance):适用于特征值为离散型或存在异常值的场景,计算两点在特征空间中的“直角距离”,公式为:
\[d(x,y)=\sum_{i=1}^{n}|x_i - y_i|\] -
筛选K个近邻:将新样本与训练集所有样本的距离按从小到大排序,取前K个距离最近的样本作为近邻。这里的K是算法的核心参数,取值直接影响预测结果。
-
执行决策规则:根据任务类型(分类/回归)采用不同规则:
-
分类任务:采用“多数投票法”,即K个近邻中出现次数最多的标签,作为新样本的预测标签。若K取偶数,可能出现投票平局,因此通常优先取奇数。
-
回归任务:采用“均值法”或“加权均值法”。均值法直接取K个近邻的数值平均值;加权均值法则给距离越近的近邻分配越高的权重(权重与距离成反比),再计算加权平均值,结果更精准。
-
3. 核心参数说明¶
KNN的参数较少,核心仅两个,但其取值对模型效果影响极大:
-
K值(近邻数量):
-
K值过小:模型复杂度高,容易过拟合(仅依赖少数近邻,对噪声敏感,比如异常值可能被当作近邻影响结果)。
-
K值过大:模型复杂度低,容易欠拟合(近邻中混入过多无关样本,模糊类别边界)。
-
最优取值:通常取1-20之间的奇数,需通过交叉验证(如10折交叉验证)确定,即尝试不同K值,选择验证集准确率最高的K。
-
-
距离度量方式:根据数据特征类型选择:
-
连续型特征、无异常值:优先选欧式距离。
-
离散型特征、存在异常值、高维数据:可选择曼哈顿距离,或更复杂的闵可夫斯基距离(欧式距离和曼哈顿距离的通用形式)。
-
三、算法优缺点分析¶
1. 优点¶
-
简单易用:原理直观,无需复杂的数学推导,代码实现难度低,适合机器学习入门。
-
训练效率高:训练阶段仅存储数据,无需拟合模型,训练时间几乎可以忽略,适合快速落地验证思路。
-
适应性强:非参数模型,无需假设数据分布,可处理分类、回归等多种任务,对非线性数据的拟合效果较好。
-
对异常值相对不敏感:当K值适中时,少数异常值被多数正常近邻“稀释”,对预测结果影响较小。
2. 缺点¶
-
内存占用高:需存储全部训练集数据,当训练集规模大、特征维度高时,对内存要求极高。
-
预测速度慢:预测时需与所有训练样本计算距离,数据量越大,预测耗时越长,不适合实时预测场景。
-
对无关特征敏感:无关特征会干扰距离计算,导致近邻筛选不准确,降低预测精度,需提前做特征筛选。
-
高维数据效果差:随着特征维度增加,样本间的距离差异会逐渐缩小(“维度灾难”),难以有效区分近邻。
四、适用场景与优化技巧¶
1. 适用场景¶
KNN适合小样本、低维度、非线性的数据场景,典型应用包括:
-
推荐系统:基于用户行为相似度推荐商品(如“你可能喜欢”功能)。
-
图像识别:简单的图像分类、人脸识别(基于像素特征相似度)。
-
文本分类:基于文本特征(如词频)的类别判断(如垃圾邮件识别)。
-
回归预测:小样本的数值预测(如房价预测、气温预测)。
2. 优化技巧¶
针对KNN的缺点,可通过以下方式优化性能:
-
特征预处理:对特征做标准化/归一化,删除无关特征、降维(如PCA),缓解维度灾难和无关特征干扰。
-
优化近邻搜索:采用空间索引(如KD树、Ball树)替代暴力搜索,减少距离计算次数,提升预测速度。
-
自适应K值:不同样本采用不同的K值(如距离近的样本取小K,距离远的样本取大K),提升模型精度。
-
加权投票/加权均值:给近邻分配与距离成反比的权重,让更近的样本对预测结果影响更大,优化精度。
五、实例¶
import numpy as np
import matplotlib.pyplot as plt
from sklearn.datasets import load_iris
from sklearn.preprocessing import StandardScaler
from sklearn.neighbors import KNeighborsClassifier
# 加载鸢尾花数据集
iris = load_iris()
# 选择2个区分度最高的特征:花瓣长度、花瓣宽度
X = iris.data[:, [2, 3]] # 花瓣长度(cm)、花瓣宽度(cm)
y = iris.target
# 特征标准化
scaler = StandardScaler()
X_scaled = scaler.fit_transform(X)
# 3. 训练KNN模型(经典K=3)
knn = KNeighborsClassifier(n_neighbors=3, metric='euclidean')
knn.fit(X_scaled, y)
# 生成网格点
h = 0.01
x_min, x_max = X_scaled[:, 0].min() - 1, X_scaled[:, 0].max() + 1
y_min, y_max = X_scaled[:, 1].min() - 1, X_scaled[:, 1].max() + 1
xx, yy = np.meshgrid(np.arange(x_min, x_max, h),
np.arange(y_min, y_max, h))
# 预测网格点类别,生成背景
Z = knn.predict(np.c_[xx.ravel(), yy.ravel()])
Z = Z.reshape(xx.shape)
# 绘制决策边界图
plt.figure(figsize=(10, 7))
plt.contourf(xx, yy, Z, alpha=0.3, cmap='Set3')
# 绘制决策边界线
plt.contour(xx, yy, Z, colors='k', linewidths=1.5, alpha=0.8)
# 叠加真实样本散点
markers = ['o', 's', '^'] # 圆形、正方形、三角形
colors = ['#1f77b4', '#ff7f0e', '#2ca02c'] # 蓝、橙、绿
# 中文类别名
iris_names = ['山鸢尾', '变色鸢尾', '维吉尼亚鸢尾']
for idx, cl in enumerate(np.unique(y)):
plt.scatter(
X_scaled[y == cl, 0], X_scaled[y == cl, 1],
alpha=0.8, c=colors[idx], marker=markers[idx],
label=iris_names[idx],
edgecolor='k', s=80
)
plt.xlabel('花瓣长度(标准化)', fontsize=12, fontweight='medium')
plt.ylabel('花瓣宽度(标准化)', fontsize=12, fontweight='medium')
plt.title('KNN决策边界可视化(K=3)- 鸢尾花数据集', fontsize=14, fontweight='bold')
plt.legend(loc='upper left', fontsize=10, framealpha=0.9)
plt.grid(True, alpha=0.2, linestyle='--')
plt.tight_layout()
plt.show()
KNN决策边界可视化(鸢尾花数据集) 从图中可直观看到KNN算法的分类逻辑:
1. 山鸢尾(蓝色圆形)被完全分隔在左侧区域,与另外两类无任何重叠,体现其特征辨识度极高;
2. 变色鸢尾(橙色正方形)和维吉尼亚鸢尾(绿色三角形)仅有少量边界重叠,K=3时的决策边界能精准区分这两类;
3. 平滑的决策边界说明K=3既避免了过拟合,也没有欠拟合。 六、总结¶
KNN核心靠“相似性”决策,优点是易用、适应性强,缺点是内存和速度瓶颈明显。它适合作为机器学习入门的实践案例,也能在小样本场景中实现快速落地。在实际应用中,需重点关注K值选择、特征预处理和近邻搜索优化,才能充分发挥其效果。