【机器学习入门】KNN算法全解析:原理+代码+实战(鸢尾花案例)
【机器学习入门】KNN算法全解析:原理+代码+实战(鸢尾花案例)
(终于忙完考试,闲下来一会了,开始机器学习的篇章)
大家好!今天终于啃完了机器学习入门经典算法——KNN(K近邻),从原理到代码踩了不少坑,也总结了很多实用技巧,特意整理成笔记分享给和我一样的入门小伙伴,全程干货,建议收藏~
一、什么是KNN算法?
KNN(K-Nearest Neighbor,K近邻算法)是最简单的监督学习算法之一,核心思想可以用一句话概括:“物以类聚,人以群分”。
它既可以做分类,也可以做回归,核心逻辑无差别:
- 分类任务:给定一个待预测样本,计算它与训练集中所有样本的距离,取距离最近的K个样本,这K个样本中出现次数最多的类别,就是该待预测样本的类别;
- 回归任务:同理,取K个最近邻样本的数值均值(或中位数)作为预测结果。
KNN的特点:
- 非参数化:无需假设数据分布,对异常值敏感;
- 惰性学习:训练阶段不计算任何参数,仅存储数据集,预测阶段才计算距离(训练快、预测慢);
- 核心依赖:距离计算方式 + K值选择 + 特征预处理(重中之重)。
二、KNN核心思路解析
2.1 核心步骤(通用)
- 数据准备:整理带标签的训练集,确保特征为数值型(类别特征需编码);
- 距离计算:选择合适的距离公式,计算待预测样本与训练集所有样本的距离;
- 选K值:选取距离最近的K个训练集样本(K为超参数,需手动指定);
- 预测决策:分类取“多数投票”,回归取“均值/中位数”。
2.2 常见距离表示方式
距离是KNN的“灵魂”,不同距离适用于不同场景,最常用的有3种:
2.2 常见距离表示方式
距离是KNN的“灵魂”,不同距离适用于不同场景,最常用的有3种:
| 距离类型 | 公式(以二维特征为例) | 适用场景 |
|---|---|---|
| 欧氏距离(最常用) | d=(x1−y1)2+(x2−y2)2d = \sqrt{(x_1-y_1)^2 + (x_2-y_2)^2}d=(x1−y1)2+(x2−y2)2 | 连续特征、各特征量纲一致(或已归一化) |
| 曼哈顿距离 | $d = | x_1-y_1 |
| 余弦相似度 | $\cos\theta = \frac{x·y}{ |
三、特征预处理:归一化 vs 标准化
KNN对特征量纲极度敏感(比如“收入(0-10万)”和“年龄(0-100)”,距离会被收入主导),因此特征缩放是KNN必做的预处理步骤,核心是让所有特征对距离的贡献权重一致。
3.1 归一化(Min-Max Scaling)
公式
将特征缩放到 [0,1] 区间(可自定义区间如[a,b]):
xscaled=x−xminxmax−xminx_{scaled} = \frac{x - x_{min}}{x_{max} - x_{min}}xscaled=xmax−xminx−xmin
特点
- 优点:缩放后范围固定,距离解释性强;
- 缺点:对异常值极其敏感(一个极端值会压缩其他特征);
- 适用场景:特征无异常值、分布已知。
3.2 标准化(Z-Score Normalization)
公式
将特征转换为 均值=0,方差=1 的标准正态分布:
xscaled=x−μσx_{scaled} = \frac{x - \mu}{\sigma}xscaled=σx−μ
(μ\muμ为特征均值,σ\sigmaσ为特征标准差)
特点
- 优点:对异常值鲁棒性更强,不改变特征分布趋势;
- 缺点:缩放后范围不固定(通常在[-3,3]);
- 适用场景:特征有极端异常值、高维数据、后续需结合其他对分布敏感的算法。
3.3 核心区别总结
| 维度 | 归一化(Min-Max) | 标准化(Z-Score) |
|---|---|---|
| 缩放范围 | 固定区间(如[0,1]) | 无固定范围(均值0,方差1) |
| 异常值敏感 | 极敏感 | 相对鲁棒 |
| 适用场景 | 无异常值、需固定范围 | 有异常值、高维数据 |
四、实战:KNN完整流程(鸢尾花案例)
以经典的鸢尾花分类任务为例,完整实现KNN的全流程:加载数据→数据切分→特征工程→模型训练→预测→评估。
4.1 环境准备
# 导入必备库
import numpy as np
import pandas as pd
from sklearn.datasets import load_iris
from sklearn.model_selection import train_test_split
from sklearn.neighbors import KNeighborsClassifier, KNeighborsRegressor
from sklearn.preprocessing import MinMaxScaler, StandardScaler
from sklearn.metrics import accuracy_score, mean_squared_error, classification_report, confusion_matrix
import matplotlib.pyplot as plt
import seaborn as sns
# 中文显示设置
plt.rcParams['font.sans-serif'] = ['SimHei']
plt.rcParams['axes.unicode_minus'] = False
4.2 完整流程实现
步骤1:加载数据集
鸢尾花数据集包含3类鸢尾花(setosa、versicolor、virginica),4个数值特征:萼片长度、萼片宽度、花瓣长度、花瓣宽度。
# 加载数据
iris = load_iris()
X = iris.data # 特征矩阵
y = iris.target # 分类标签(0/1/2)
feature_names = iris.feature_names # 特征名
target_names = iris.target_names # 类别名
# 查看数据基本信息
print("特征维度:", X.shape)
print("特征名:", feature_names)
print("类别:", target_names)
步骤2:划分训练集/测试集(关键!避免数据泄露)
核心原则:先划分数据集,再做特征预处理!
# 划分训练集(80%)和测试集(20%)
X_train, X_test, y_train, y_test = train_test_split(
X, y, test_size=0.2, random_state=42 # random_state固定随机种子,保证结果可复现
)
print("训练集维度:", X_train.shape)
print("测试集维度:", X_test.shape)
步骤3:特征工程(归一化/标准化)
对比两种缩放方式的效果,这里以分类任务为例:
# -------------- 方式1:归一化(Min-Max) --------------
scaler_minmax = MinMaxScaler()
X_train_minmax = scaler_minmax.fit_transform(X_train) # 仅用训练集拟合scaler
X_test_minmax = scaler_minmax.transform(X_test) # 测试集复用训练集的scaler
# -------------- 方式2:标准化(Z-Score) --------------
scaler_standard = StandardScaler()
X_train_standard = scaler_standard.fit_transform(X_train)
X_test_standard = scaler_standard.transform(X_test)
# 查看缩放后的特征范围(验证效果)
print("归一化后训练集特征最小值:", X_train_minmax.min(axis=0))
print("归一化后训练集特征最大值:", X_train_minmax.max(axis=0))
print("标准化后训练集特征均值:", np.round(X_train_standard.mean(axis=0), 2))
print("标准化后训练集特征方差:", np.round(X_train_standard.var(axis=0), 2))
步骤4:模型训练(分类+回归)
4.4.1 KNN分类(核心)
# 初始化KNN分类器(K=3,欧氏距离)
knn_clf = KNeighborsClassifier(n_neighbors=3, metric='euclidean')
# 用归一化后的训练集训练模型
knn_clf.fit(X_train_minmax, y_train)
# 模型预测
y_pred = knn_clf.predict(X_test_minmax)
# 预测概率(可选,查看每个样本属于各类别的概率)
y_pred_proba = knn_clf.predict_proba(X_test_minmax)
print("测试集预测结果:", y_pred)
print("前5个样本预测概率:\n", np.round(y_pred_proba[:5], 2))
4.4.2 KNN回归(拓展)
若将鸢尾花的“花瓣长度”作为预测目标,其他特征作为输入,可实现回归任务:
# 构造回归数据集:预测花瓣长度(第2列特征)
X_reg = np.delete(X, 2, axis=1) # 输入:除花瓣长度外的3个特征
y_reg = X[:, 2] # 目标:花瓣长度
# 划分回归数据集
X_reg_train, X_reg_test, y_reg_train, y_reg_test = train_test_split(
X_reg, y_reg, test_size=0.2, random_state=42
)
# 归一化+回归训练
scaler_reg = MinMaxScaler()
X_reg_train_scaled = scaler_reg.fit_transform(X_reg_train)
X_reg_test_scaled = scaler_reg.transform(X_reg_test)
knn_reg = KNeighborsRegressor(n_neighbors=3)
knn_reg.fit(X_reg_train_scaled, y_reg_train)
y_reg_pred = knn_reg.predict(X_reg_test_scaled)
print("回归预测前5个结果:", np.round(y_reg_pred[:5], 2))
print("真实前5个值:", np.round(y_reg_test[:5], 2))
步骤5:模型评估
5.1 分类任务评估(多维度)
# 1. 基础准确率
acc = accuracy_score(y_test, y_pred)
print(f"KNN分类准确率:{acc:.4f}")
# 2. 分类报告(精确率、召回率、F1值)
print("\n分类报告:")
print(classification_report(y_test, y_pred, target_names=target_names))
# 3. 混淆矩阵(可视化)
cm = confusion_matrix(y_test, y_pred)
plt.figure(figsize=(8, 6))
sns.heatmap(cm, annot=True, fmt='d', cmap='Blues',
xticklabels=target_names, yticklabels=target_names)
plt.xlabel('预测类别')
plt.ylabel('真实类别')
plt.title('KNN分类混淆矩阵')
plt.show()
5.2 回归任务评估
# 回归评估指标:均方误差(MSE)、均方根误差(RMSE)
mse = mean_squared_error(y_reg_test, y_reg_pred)
rmse = np.sqrt(mse)
print(f"\nKNN回归MSE:{mse:.4f}")
print(f"KNN回归RMSE:{rmse:.4f}")
# 可视化回归预测结果
plt.figure(figsize=(10, 6))
plt.plot(range(len(y_reg_test)), y_reg_test, label='真实值', color='blue', marker='o')
plt.plot(range(len(y_reg_pred)), y_reg_pred, label='预测值', color='red', marker='x')
plt.xlabel('样本序号')
plt.ylabel('花瓣长度(cm)')
plt.title('KNN回归预测结果对比')
plt.legend()
plt.show()
五、学习总结 & 避坑指南
5.1 核心知识点
- KNN的核心是“距离”和“投票/均值”,无训练过程,属于惰性学习;
- 特征缩放是KNN的必选项:无异常值用归一化,有异常值用标准化;
- 数据划分必须在预处理前,避免测试集信息泄露;
- K值选择需调参(K太小易过拟合,K太大易欠拟合)。
5.2 常见坑点
- ❌ 先归一化再划分数据集:导致测试集“污染”,评估结果失真;
- ❌ 忽略异常值:归一化对异常值敏感,需先通过IQR/3σ法则处理;
- ❌ 只看准确率:数据不平衡时,准确率无意义,需结合混淆矩阵、F1值;
- ❌ 高维数据直接用KNN:高维空间中“距离”会失效(维度灾难),需先降维(如PCA)。
六、最后
KNN作为入门算法,虽然简单,但能帮我们理解机器学习的核心逻辑——“基于相似性做预测”。今天从原理到实战走了一遍,最大的感悟是:算法本身不难,难的是数据预处理和模型评估,细节决定成败!
如果有小伙伴有疑问,欢迎评论区交流~一起学习,一起进步!
DAMO开发者矩阵,由阿里巴巴达摩院和中国互联网协会联合发起,致力于探讨最前沿的技术趋势与应用成果,搭建高质量的交流与分享平台,推动技术创新与产业应用链接,围绕“人工智能与新型计算”构建开放共享的开发者生态。
更多推荐



所有评论(0)