【机器学习入门】KNN算法全解析:原理+代码+实战(鸢尾花案例)

(终于忙完考试,闲下来一会了,开始机器学习的篇章)
大家好!今天终于啃完了机器学习入门经典算法——KNN(K近邻),从原理到代码踩了不少坑,也总结了很多实用技巧,特意整理成笔记分享给和我一样的入门小伙伴,全程干货,建议收藏~

一、什么是KNN算法?

KNN(K-Nearest Neighbor,K近邻算法)是最简单的监督学习算法之一,核心思想可以用一句话概括:“物以类聚,人以群分”

它既可以做分类,也可以做回归,核心逻辑无差别:

  • 分类任务:给定一个待预测样本,计算它与训练集中所有样本的距离,取距离最近的K个样本,这K个样本中出现次数最多的类别,就是该待预测样本的类别;
  • 回归任务:同理,取K个最近邻样本的数值均值(或中位数)作为预测结果。

KNN的特点:

  • 非参数化:无需假设数据分布,对异常值敏感;
  • 惰性学习:训练阶段不计算任何参数,仅存储数据集,预测阶段才计算距离(训练快、预测慢);
  • 核心依赖:距离计算方式 + K值选择 + 特征预处理(重中之重)。

二、KNN核心思路解析

2.1 核心步骤(通用)

  1. 数据准备:整理带标签的训练集,确保特征为数值型(类别特征需编码);
  2. 距离计算:选择合适的距离公式,计算待预测样本与训练集所有样本的距离;
  3. 选K值:选取距离最近的K个训练集样本(K为超参数,需手动指定);
  4. 预测决策:分类取“多数投票”,回归取“均值/中位数”。

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=(x1y1)2+(x2y2)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=xmaxxminxxmin

特点
  • 优点:缩放后范围固定,距离解释性强;
  • 缺点:对异常值极其敏感(一个极端值会压缩其他特征);
  • 适用场景:特征无异常值、分布已知。

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 核心知识点

  1. KNN的核心是“距离”和“投票/均值”,无训练过程,属于惰性学习;
  2. 特征缩放是KNN的必选项:无异常值用归一化,有异常值用标准化;
  3. 数据划分必须在预处理前,避免测试集信息泄露;
  4. K值选择需调参(K太小易过拟合,K太大易欠拟合)。

5.2 常见坑点

  • ❌ 先归一化再划分数据集:导致测试集“污染”,评估结果失真;
  • ❌ 忽略异常值:归一化对异常值敏感,需先通过IQR/3σ法则处理;
  • ❌ 只看准确率:数据不平衡时,准确率无意义,需结合混淆矩阵、F1值;
  • ❌ 高维数据直接用KNN:高维空间中“距离”会失效(维度灾难),需先降维(如PCA)。

六、最后

KNN作为入门算法,虽然简单,但能帮我们理解机器学习的核心逻辑——“基于相似性做预测”。今天从原理到实战走了一遍,最大的感悟是:算法本身不难,难的是数据预处理和模型评估,细节决定成败!

如果有小伙伴有疑问,欢迎评论区交流~一起学习,一起进步!


Logo

DAMO开发者矩阵,由阿里巴巴达摩院和中国互联网协会联合发起,致力于探讨最前沿的技术趋势与应用成果,搭建高质量的交流与分享平台,推动技术创新与产业应用链接,围绕“人工智能与新型计算”构建开放共享的开发者生态。

更多推荐