零基础入门机器学习 -- 第七章K近邻算法KNN
·
学习目标:
- 通过直观案例理解 KNN 算法的原理
- 学会如何使用 Python 实现 KNN 分类
- 通过实际案例(预测水果种类)学会 KNN 的应用
- 了解 KNN 的优缺点,以及如何优化它
7.1 K 近邻(KNN)的直观理解
从现实生活看 KNN
想象一下,你走进一个陌生的水果市场,看到一个从未见过的水果,你想知道它是什么。
你有两种方法:
- 自己瞎猜,可能猜错
- 观察周围的水果,看看它更像苹果、橘子还是香蕉
你选择了方法 2,并发现:
- 这个水果的形状很圆,颜色是红色
- 周围几个红色的、圆形的水果都是苹果
- 于是你推测,这个水果大概率也是苹果
✅ 这个过程就像 K 近邻(KNN)算法的思路——“看看你周围的邻居,来判断自己属于哪一类”。
7.2 什么是 K 近邻(KNN)?
KNN(K-Nearest Neighbors,K 近邻)是一种基于距离的分类算法:
- K 代表“最近的 K 个邻居”
- 近邻 意味着谁和你最接近,谁就更可能影响你的类别
KNN 的工作流程:
- 计算距离:找到离当前样本最近的 K 个已知类别的样本
- 投票表决:看看 K 个邻居中哪个类别最多
- 最终分类:将当前样本归类到得票最多的类别
✅ 适用于:
- 分类问题(如邮件垃圾分类)
- 回归问题(如预测房价)
❌ 不适用于:
- 大规模数据(计算所有样本的距离会很慢)
- 高维数据(容易受噪声影响)
7.3 KNN 是如何工作的?
示例:判断水果种类
假设你在水果市场发现了一个未知水果,你想知道它是苹果、橘子还是香蕉。
你有一张参考表,里面是一些已知水果的甜度(1-10)和重量(50g-300g):
| 水果 | 甜度(1-10) | 重量(克) | 类别 |
|---|---|---|---|
| 苹果A | 8 | 150 | 🍎 |
| 苹果B | 7 | 140 | 🍎 |
| 苹果C | 9 | 160 | 🍎 |
| 橘子A | 6 | 120 | 🍊 |
| 橘子B | 5 | 110 | 🍊 |
| 橘子C | 6 | 125 | 🍊 |
| 香蕉A | 9 | 200 | 🍌 |
| 香蕉B | 10 | 220 | 🍌 |
| 香蕉C | 9 | 210 | 🍌 |
| 香蕉D | 8 | 205 | 🍌 |
📌 现在有一个未知水果(甜度 7.5,重量 160g),它属于哪一类?
步骤
- 计算它与所有水果的距离
- 使用欧几里得距离公式:

- 使用欧几里得距离公式:
- 找到最近的 K 个邻居
- 让邻居“投票”决定分类
7.4 代码实现:KNN 预测水果种类
1️⃣ 导入必要的库
import numpy as np
import pandas as pd
import matplotlib.pyplot as plt
from sklearn.model_selection import train_test_split
from sklearn.preprocessing import StandardScaler
from sklearn.neighbors import KNeighborsClassifier
from sklearn.metrics import accuracy_score
2️⃣ 创建水果数据集
# 创建水果数据
data = {
"水果": ["苹果A", "苹果B", "苹果C", "苹果D", "橘子A", "橘子B", "橘子C", "橘子D", "香蕉A", "香蕉B", "香蕉C", "香蕉D", "香蕉E"],
"甜度": [8, 7, 9, 6, 6, 5, 6, 7, 9, 10, 9, 8, 9],
"重量": [150, 140, 160, 155, 120, 110, 125, 130, 200, 220, 210, 205, 215],
"类别": ["苹果", "苹果", "苹果", "苹果", "橘子", "橘子", "橘子", "橘子", "香蕉", "香蕉", "香蕉", "香蕉", "香蕉"]
}
# 转换为 DataFrame
df = pd.DataFrame(data)
# 映射类别为数值(苹果=0, 橘子=1, 香蕉=2)
df["类别"] = df["类别"].map({"苹果": 0, "橘子": 1, "香蕉": 2})
# 查看数据
print(df)
3️⃣ 训练 KNN 模型
# 选择特征(甜度、重量)和目标变量(类别)
X = df[["甜度", "重量"]]
y = df["类别"]
# stratify 确保类别均匀拆分
X_train, X_test, y_train, y_test = train_test_split(X, y, test_size=0.3, stratify=y, random_state=42)
# 归一化数据
scaler = StandardScaler()
X_train = scaler.fit_transform(X_train)
X_test = scaler.transform(X_test)
# 训练 KNN 模型
knn = KNeighborsClassifier(n_neighbors=3)
knn.fit(X_train, y_train)
# 计算准确率
y_pred = knn.predict(X_test)
print("优化后模型准确率:", accuracy_score(y_test, y_pred))
4️⃣ 预测未知水果
# 假设一个新水果(甜度 7.5,重量 160 克)
new_fruit = np.array([[7.5, 160]])
# 归一化数据
new_fruit = scaler.transform(new_fruit)
# 预测水果类别
prediction = knn.predict(new_fruit)
# 映射回类别
fruit_labels = {0: "苹果", 1: "橘子", 2: "香蕉"}
print("预测的新水果类型:", fruit_labels[prediction[0]])
示例输出:
优化后模型准确率: 0.75
预测的新水果类型: 苹果
7.5 结果分析
如果 K=3,模型会选择 3 个最近的水果进行投票:
- 最近的 3 个水果可能是 苹果、苹果、橘子
- 因为苹果 2 票 > 橘子 1 票
- KNN 预测这个新水果是 苹果 🍎
7.6 总结
- KNN 通过计算“邻居”距离进行分类
- 适用于小数据集
- 适合电影分类、信用评分、手写识别等应用
- 选择合适的 K 值很重要
🔥 试试看上面的代码,并尝试不同的 K 值,看看结果如何变化! 🚀
DAMO开发者矩阵,由阿里巴巴达摩院和中国互联网协会联合发起,致力于探讨最前沿的技术趋势与应用成果,搭建高质量的交流与分享平台,推动技术创新与产业应用链接,围绕“人工智能与新型计算”构建开放共享的开发者生态。
更多推荐

所有评论(0)