经典分类算法KNN的研究
KNN分类算法是一种基于“物以类聚”思想的监督学习算法,通过计算待分类样本与训练集中各样本的距离,选取最近的K个邻居,根据多数表决原则确定其类别。
算法核心原理
- 距离计算:对待分类样本与训练集中的每个样本计算距离,常用的距离度量包括:
- 欧氏距离:两点间的直线距离,公式为 d=∑i=1n(xi−yi)2d=∑i=1n(xi−yi)2
- 曼哈顿距离:各维度绝对差之和,d=∑i=1n∣xi−yi∣d=∑i=1n∣xi−yi∣
- 余弦相似度:衡量向量方向一致性,适用于文本等高维稀疏数据
- 选择K个最近邻:将所有距离按升序排序,取前K个最邻近的样本。
- 多数表决:统计这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}")

使用网络搜索法寻找最优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 值),而其他参数(如weights,metric,p)保持默认值(例如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 为最优(因为它没尝试距离权重)。
- 示例:手动遍历 K=5 时用的是均匀权重,得分 0.90;但 GridSearchCV 可能发现 K=7 且使用距离权重 (
3. 数据泄露与预处理步骤的影响
- 标准化时机:
- 正确做法(GridSearchCV 内部管道或严格分离):应在每一折的训练集上拟合 scaler,再转换训练集和验证集。
- 常见错误(手动遍历):如果在划分数据集之前就对整个数据集进行了
fit_transform,会导致数据泄露(Data Leakage)。测试集的信息“泄露”到了训练过程中,导致评估分数虚高且不稳定。这种错误的预处理方式会导致手动遍历选出的 K 值不可靠,与严谨的 GridSearchCV 结果产生偏差。
表格
| 特性 | 手动遍历 (简单划分) | GridSearchCV (交叉验证) |
|---|---|---|
| 评估稳定性 | 低,受单次划分影响大 | 高,多次评估取平均 |
| 数据利用率 | 较低,部分数据仅用于测试 | 较高,所有数据都参与过验证 |
| 过拟合风险 | 容易过拟合到特定测试集 | 较低,更能反映泛化能力 |
| 计算成本 | 低 | 高(需训练 K * N 次模型) |
应用场景
- 图像识别:比较像素特征进行分类
- 推荐系统:基于用户行为相似性推荐商品
- 医学诊断:根据患者指标判断疾病类型
用户行为相似性?那么炒股是否也是用户行为呢?当然是的,那么基于用户行为相似性,是否可以根据历史用户行为推测未来股票趋势呢?应该也是可以的,期待验证。
-----------------------------------------------------------------
- 我做的各种程序们
- 小y的QQ:28657321 (欢迎交流)

浙公网安备 33010602011771号