惯性聚合 高效追踪和阅读你感兴趣的博客、新闻、科技资讯
阅读原文 在惯性聚合中打开

推荐订阅源

aimingoo的专栏
aimingoo的专栏
Jina AI
Jina AI
WordPress大学
WordPress大学
Recent Announcements
Recent Announcements
G
Google Developers Blog
I
InfoQ
H
Hackread – Cybersecurity News, Data Breaches, AI and More
Google DeepMind News
Google DeepMind News
P
Proofpoint News Feed
MyScale Blog
MyScale Blog
M
MIT News - Artificial intelligence
让小产品的独立变现更简单 - ezindie.com
让小产品的独立变现更简单 - ezindie.com
C
Check Point Blog
J
Java Code Geeks
T
Tailwind CSS Blog
OSCHINA 社区最新新闻
OSCHINA 社区最新新闻
Microsoft Security Blog
Microsoft Security Blog
MongoDB | Blog
MongoDB | Blog
V
Visual Studio Blog
人人都是产品经理
人人都是产品经理
量子位
A
About on SuperTechFans
D
DataBreaches.Net
钛媒体:引领未来商业与生活新知
钛媒体:引领未来商业与生活新知

博客园 - 小y

HP-Socket压力测试例子,用恐怖如斯来形容。 用python自己回测股票策略 怎么让大模型查询本地数据库输出结果(Dify+豆包) OpenClaw工作原理解析【建议背诵】 AiAgent工具之网页搜索 怎么在百度搜索中屏蔽csdn 使用C#调用Yolo26模型的ONNX ESP32 使用MicroPython编写代码实现远程控制LED灯 使用Razor模板引擎实现自动生成代码 ESP32开发环境下载、USB驱动安装 2026年最新年会抽奖小程序 grafana+prometheus快速实现可视化大屏 PyAutoGUI库自动化测试脚本工具模拟键盘鼠标操作 性能卓越的开源时序数据库——QuestDB GitHub超 30000+ star , 超强大的开源项目Supervision DevExpress系列:dxValidationProvider和dxErrorProvider两种控件的用法和区别 实测:MySQL跑在Docker里会损失多少性能?寻找Docker优化点 用C# GDI编写粒子效果 动画图解嵌入式常见的通讯协议:SPI、I²C、UART、红外 C#中的MVVM框架 .net 9 中的WinUI3和MAUI有什么区别 Netty的高性能之道 一文搞定理解RPC CSS flex布局(弹性布局/弹性盒子) 为什么选择使用TypeScript?这篇文章讲透了
经典分类算法KNN的研究
小y · 2026-05-19 · via 博客园 - 小y

KNN分类算法‌是一种基于“物以类聚”思想的监督学习算法,通过计算待分类样本与训练集中各样本的距离,选取最近的K个邻居,根据多数表决原则确定其类别。

算法核心原理

  1. ‌距离计算‌:对待分类样本与训练集中的每个样本计算距离,常用的距离度量包括:
    • ‌欧氏距离‌:两点间的直线距离,公式为 d=∑i=1n(xi−yi)2d=i=1n(xiyi)2
    • ‌曼哈顿距离‌:各维度绝对差之和,d=∑i=1n∣xi−yi∣d=i=1nxiyi
    • ‌余弦相似度‌:衡量向量方向一致性,适用于文本等高维稀疏数据
  2. ‌选择K个最近邻‌:将所有距离按升序排序,取前K个最邻近的样本。
  3. ‌多数表决‌:统计这K个邻居中各类别出现的频率,将频率最高的类别作为预测结果。

算法实现

接下来我们将通过 ‌Python‌ 和 ‌scikit-learn‌ 库,以经典的 ‌鸢尾花数据集(Iris Dataset)‌ 为例,完整演示 KNN 分类算法的实现流程。

这个流程涵盖了从数据加载、预处理、模型训练、超参数调优到最终评估的全过程。

from sklearn.preprocessing import StandardScaler
from sklearn.neighbors import KNeighborsClassifier
from sklearn.metrics import classification_report, confusion_matrix, accuracy_score
from sklearn.datasets import load_iris
from sklearn.model_selection import train_test_split
import seaborn as sns
import matplotlib.pyplot as plt
import math

# 1.加载鸢尾花数据集
iris=load_iris()
X,y=iris.data,iris.target
# 2. 数据标准化
scaler = StandardScaler()
X_scaled = scaler.fit_transform(X)
# 3. 划分训练集和测试集
# test_size=0.3 表示 30% 用于测试,random_state 确保结果可复现
X_train, X_test, y_train, y_test = train_test_split(
    X_scaled, y, test_size=0.3, random_state=42, stratify=y
)

# 4. 寻找最佳 K 值
# 假设 X_train 是你的训练数据特征矩阵
N = X_train.shape[0]  # 获取训练集样本数量 N
max_k = int(math.sqrt(N))  # 计算根号 N 并取整
# 生成从 1 到 max_k 之间的所有奇数
k_values = range(1, max_k + 1, 2)
#缓存预测结果及评分
accuracies = []

#遍历K值,找到最佳的K值
for k in k_values:
    knn = KNeighborsClassifier(n_neighbors=k)
    knn.fit(X_train, y_train)
    y_pred = knn.predict(X_test)
    accuracies.append(accuracy_score(y_test, y_pred))

# 绘制 K 值与准确率的关系图
plt.figure(figsize=(10, 6))
plt.plot(k_values, accuracies, marker='o', linestyle='-')
plt.title('Accuracy vs K')
plt.xlabel('K')
plt.ylabel('Accuracy')
plt.xticks(k_values)
plt.grid(True)
plt.show()

# 获取最佳 K 值
best_k = k_values[np.argmax(accuracies)]
print(f"最佳 K 值: {best_k}, 对应准确率: {max(accuracies):.4f}")

  16

使用网络搜索法寻找最优K值

前面的代码是用手动遍历法选择最优K值,也可用网络搜索法选择最优K值

from sklearn.preprocessing import StandardScaler
from sklearn.neighbors import KNeighborsClassifier
from sklearn.metrics import classification_report, confusion_matrix, accuracy_score
from sklearn.datasets import load_iris
from sklearn.model_selection import train_test_split,GridSearchCV
import seaborn as sns
import matplotlib.pyplot as plt
import math

# 1.加载鸢尾花数据集
iris=load_iris()
X,y=iris.data,iris.target
# 2. 数据标准化
scaler = StandardScaler()
X_scaled = scaler.fit_transform(X)
# 3. 划分训练集和测试集
# test_size=0.3 表示 30% 用于测试,random_state 确保结果可复现
X_train, X_test, y_train, y_test = train_test_split(
    X_scaled, y, test_size=0.3, random_state=42, stratify=y
)

# 4. 寻找最佳 K 值
n_neighbors=tuple(range(1,31,1))
#创建网络搜索实例
cv=GridSearchCV(estimator=KNeighborsClassifier(),param_grid={'n_neighbors':n_neighbors},cv=5)
cv.fit(X_scaled,y)
# 获取最佳 K 值
best_k = cv.best_params_["n_neighbors"]
print(best_k)
knn = KNeighborsClassifier(n_neighbors= best_k)
knn.fit(X_train, y_train)
print(f"最佳 K 值: {best_k}, 对应准确率: {knn.score(X_train, y_train):.4f}")

  结果:6

最佳 K 值: 6, 对应准确率: 0.9524

为何结果与前面不一样?

1. 评估机制不同:交叉验证 vs. 单次划分

这是造成结果差异的最主要原因。

  • ‌手动遍历法(简单 Hold-out)‌:
    通常将数据一次性划分为训练集和测试集(例如 70% 训练,30% 测试)。模型只在‌这一种固定的数据划分‌上进行训练和评估。如果这次划分中,测试集恰好包含了一些容易分类的样本,或者训练集缺失了某些关键特征分布,得到的准确率就会带有偶然性(高方差)。
  • ‌GridSearchCV(交叉验证)‌:
    默认使用 K 折交叉验证(如 5 折或 10 折)。它将训练数据分成 K 份,轮流用其中 K-1 份训练,1 份验证,重复 K 次并取‌平均得分‌。
    • ‌结果更稳定‌:交叉验证利用了更多数据进行评估,减少了因数据划分随机性带来的偏差。
    • ‌差异来源‌:手动遍历可能因为某次特定的“幸运”划分,使得某个非最优的 K 值在测试集上表现极好;而 GridSearchCV 通过平均化,更能反映模型在整体数据分布上的真实性能,因此选出的 K 值往往更具泛化能力,但也可能与单次划分的最佳 K 不同。

2. 搜索空间与参数组合的差异

  • ‌手动遍历‌:
    初学者在手动写循环时,往往只调整 n_neighbors (K 值),而其他参数(如 weightsmetricp)保持默认值(例如 weights='uniform'metric='minkowski'p=2)。
  • ‌GridSearchCV‌:
    通常用于同时搜索多个超参数的组合。如果你在 param_grid 中不仅定义了 K 值,还定义了其他参数(如 weights=['uniform', 'distance']),GridSearchCV 会寻找‌全局最优组合‌。
    • ‌示例‌:手动遍历 K=5 时用的是均匀权重,得分 0.90;但 GridSearchCV 可能发现 K=7 且使用距离权重 (weights='distance') 时,得分高达 0.92。此时 GridSearchCV 返回的最优 K 是 7,而手动遍历若只看 K 值可能会误判 5 为最优(因为它没尝试距离权重)。

3. 数据泄露与预处理步骤的影响

  • ‌标准化时机‌:
    • ‌正确做法(GridSearchCV 内部管道或严格分离)‌:应在每一折的训练集上拟合 scaler,再转换训练集和验证集。
    • ‌常见错误(手动遍历)‌:如果在划分数据集之前就对‌整个数据集‌进行了 fit_transform,会导致数据泄露(Data Leakage)。测试集的信息“泄露”到了训练过程中,导致评估分数虚高且不稳定。这种错误的预处理方式会导致手动遍历选出的 K 值不可靠,与严谨的 GridSearchCV 结果产生偏差。
特性手动遍历 (简单划分)GridSearchCV (交叉验证)
‌评估稳定性‌ 低,受单次划分影响大 高,多次评估取平均
‌数据利用率‌ 较低,部分数据仅用于测试 较高,所有数据都参与过验证
‌过拟合风险‌ 容易过拟合到特定测试集 较低,更能反映泛化能力
‌计算成本‌ 高(需训练 K * N 次模型)
 

应用场景

  • ‌图像识别‌:比较像素特征进行分类
  • ‌推荐系统‌:基于用户行为相似性推荐商品
  • ‌医学诊断‌:根据患者指标判断疾病类型

用户行为相似性?那么炒股是否也是用户行为呢?当然是的,那么基于用户行为相似性,是否可以根据历史用户行为推测未来股票趋势呢?应该也是可以的,期待验证。