31 客户购买意愿分析(K近邻算法)
31.1 引言K近邻算法(KNN)
KNN算法:监督学习的分类算法
- 原理: “物以类聚,人以群分”
- 决策: 根据K个最近邻居的类别投票
- 金融应用: 客户分类、信用评分
31.2 本章学习目标
先修内容:阅读本章前,建议先完成第 章节 10 章(Pandas 数据框基础)、第 章节 6 章(NumPy 数组运算,涉及欧氏距离计算)与第 章节 30 章(K-Means,便于对比监督学习与无监督学习)。
- KNN”最近邻投票”的决策规则与欧氏距离度量;
- 特征标准化对距离型算法的必要性,以及
fit_transform与transform的区别; - 用网格搜索加交叉验证选择超参数 K 的方法;
- 用准确率与混淆矩阵评估分类模型;
- 在教学平台完成社交网络广告数据的客户购买意愿预测任务。
31.3 算法原理
决策规则: 对于样本x,找到其K个最近邻居,这K个邻居中最多数样本所属的类别即为x的预测类别。
距离度量: 常用欧氏距离
\[ d(x,y) = \sqrt{\sum_{i=1}^{n}(x_i - y_i)^2} \]
31.4 社交网络广告案例
任务要求:将下方代码原样输入教学平台并运行(注释可省略)。该代码读取社交网络广告数据,以年龄与估计薪资为特征、是否购买为标签,先划分训练集/测试集并做标准化,再用网格搜索与 4 折交叉验证选择最优 K 值的 KNN 模型,最后输出测试集预测结果、准确率、交叉验证最佳得分与最佳模型。
# 注:04_Social_Network_Ads.csv数据文件本地没有,但平台已经内置
# ⚠️ 平台原始代码 - 请原样输入至教学平台(注释除外),平台才会判定答案正确
# 导入相关模块
import pandas as pd
import numpy as np # 导入NumPy数值计算库
from sklearn.model_selection import train_test_split # 导入Scikit-learn的train_test_split模块
from sklearn.preprocessing import StandardScaler # 导入Scikit-learn的StandardScaler模块
from sklearn.neighbors import KNeighborsClassifier # 导入Scikit-learn的KNeighborsClassifier模块
from sklearn.model_selection import GridSearchCV # 导入Scikit-learn的GridSearchCV模块
from sklearn.metrics import confusion_matrix # 导入Scikit-learn的confusion_matrix模块
# 导入数据集
data = pd.read_csv("04_Social_Network_Ads.csv")
# 数据集划分
x = data[["Age","EstimatedSalary"]]
y = data["Purchased"] # 提取Purchased列作为y变量
# 划分训练集和测试集
x_train, x_test, y_train, y_test = train_test_split(x, y, test_size=0.30, random_state=0)
# 数据标准化
transfer1 = StandardScaler()
x_train = transfer1.fit_transform(x_train) # 对数据进行变换
x_test = transfer1.transform(x_test) # 对数据进行变换
# 训练模型
estimator = KNeighborsClassifier(algorithm='kd_tree')
# 模型选择与调优——网格搜索和交叉验证
# 准备要调的超参数
param_dict = {"n_neighbors": [1, 3, 5, 7, 9, 11, 13]}
estimator = GridSearchCV(estimator, param_grid=param_dict, cv=4) # 创建网格搜索交叉验证对象,自动寻找最优超参数
estimator.fit(x_train,y_train) # 在数据上训练estimator模型
# 模型评估
y_pre = estimator.predict(x_test)
print("预测结果:\n", y_pre) # 输出预测结果:\n
print("准确率为:\n", estimator.score(x_test, y_test)) # 输出准确率为:\n
# print("对比真实值和预测值:\n", y_test==y_pre)
print("在交叉验证中最好的结果为:\n", estimator.best_score_)
print("使用网格搜索的最好的模型:\n", estimator.best_estimator_) # 输出使用网格搜索的最好的模型:\n预期输出:平台依次打印预测结果数组(0/1 构成)、测试集准确率、交叉验证最佳得分 best_score_ 与最佳模型 best_estimator_(其 n_neighbors 即网格搜索选中的 K)。判读要点:准确率与 best_score_ 都在 0~1 之间,后者是交叉验证在训练集各折上的平均表现,与测试集得分没有固定的大小关系,略低或略高都属正常;best_estimator_ 中还可以看到 algorithm='kd_tree' 等模型配置。04_Social_Network_Ads.csv 数据文件本地没有、平台已内置,具体数值以平台运行结果为准。
31.5 模型训练与调优
# 注:本块示例使用标准化后的训练集(记作 X_train_scaled / X_test_scaled),
# 对应平台任务中 transfer1.fit_transform 变换后的小写 x_train 与 transform 后的 x_test;
# 运行前请先执行:X_train_scaled, X_test_scaled = x_train, x_test,
# 并补充导入:from sklearn.metrics import accuracy_score
# ==================== 创建KNN分类器 ====================
# KNN通过投票机制分类:K个邻居中多数类作为预测类
knn = KNeighborsClassifier(algorithm='kd_tree') # 使用KD树算法加速最近邻搜索
# ==================== 定义参数网格 ====================
# K值是KNN的超参数,需要通过交叉验证选择最优值
param_grid = {'n_neighbors': [1, 3, 5, 7, 9, 11, 13]} # 测试不同的K值
# ==================== 网格搜索+交叉验证 ====================
# 通过交叉验证找到使准确率最高的K值
grid_search = GridSearchCV(
knn, # 基础模型
param_grid=param_grid, # 参数网格
cv=4, # 4折交叉验证(将训练集分成4份,轮流做验证集)
n_jobs=-1 # 使用所有CPU核心并行计算
)
# ==================== 训练模型 ====================
grid_search.fit(X_train_scaled, y_train) # 在标准化后的训练集上拟合
# ==================== 输出最佳参数 ====================
print(f'最佳K值: {grid_search.best_params_["n_neighbors"]}') # 打印最优K值
print(f'交叉验证最佳准确率: {grid_search.best_score_:.4f}') # 打印交叉验证的最高准确率31.6 模型评估
# 注:本块承接上方网格搜索块的 grid_search 与 X_test_scaled(见其块首的变量映射说明)
# ==================== 预测测试集 ====================
# 使用最优模型对测试集进行预测
y_pred = grid_search.predict(X_test_scaled) # 输出预测类别(0或1)
# ==================== 计算准确率 ====================
accuracy = accuracy_score(y_test, y_pred) # 计算预测准确率
print(f'测试集准确率: {accuracy:.4f}') # 打印准确率,范围[0,1]
# ==================== 生成混淆矩阵 ====================
# 混淆矩阵展示了预测结果与真实情况的对比
cm = confusion_matrix(y_test, y_pred) # 生成2x2混淆矩阵
print(f'\n混淆矩阵:')
print(cm)
# 混淆矩阵格式:
# 预测0 预测1
# 实际0 [TN [FP
# 实际1 [FN [TP
# ==================== 可视化混淆矩阵 ====================
import matplotlib.pyplot as plt # 导入绘图库
import seaborn as sns # 导入统计绘图库
plt.figure(figsize=(8, 6)) # 创建8x6英寸的画布
sns.heatmap(cm, annot=True, fmt='d', cmap='Blues') # 绘制热力图,annot显示数值,fmt整数格式
plt.title('混淆矩阵', fontsize=14) # 设置标题
plt.xlabel('预测类别', fontsize=12) # x轴标签
plt.ylabel('真实类别', fontsize=12) # y轴标签
plt.tight_layout() # 自动调整布局,避免标签重叠
plt.show() # 显示图形
# ==================== 输出最优模型 ====================
print(f'\n最佳模型: {grid_search.best_estimator_}') # 打印最优模型的完整配置31.7 分析结论
- 标准化必要: 年龄和薪资尺度差异大
- K值选择: 通过交叉验证确定
- 模型性能: 达到较高准确率
- 商业应用: 预测用户是否购买
31.8 本章小结
要点:
- KNN 是监督学习分类算法:模型不显式”训练”,预测时才寻找 K 个最近邻居,按多数票决定类别;
- 距离对量纲敏感,年龄(几十)与薪资(数万)同列计算前必须标准化,本章用
StandardScaler完成; - 超参数 K 通过网格搜索(
GridSearchCV)配合交叉验证自动选择,避免凭感觉拍定; - 模型评估看测试集准确率与混淆矩阵,二者分别回答”总体对多少”与”错在哪一类”。
易错点:
- 对测试集误用
fit_transform而不是transform,导致测试集用了自己的均值方差,信息泄露且口径错误; - K 取 1 易过拟合(过度敏感于个别近邻),K 取过大又会淹没局部结构,应交由交叉验证权衡;
- 把交叉验证得分与测试集准确率混为一谈,或因两者数值不同而怀疑代码有错;
- 忽略 KNN 的”K”与 K-Means 的”K”含义完全不同:前者是参与投票的邻居数,后者是簇的个数。
31.9 动手与思考
以下练习每题附参考答案(默认折叠)。请先独立完成并写下你的判断,再点开对照,最后上机验证。
输出预测:某测试点的 5 个近邻按距离从小到大依次为 d=(0.4, 0.9, 1.1, 1.5, 2.3),对应标签 y=(0, 1, 1, 0, 1)。请分别预测 K=1 与 K=3 时的类别,并说明 K 从 1 增至 3 时预测为什么可能改变。
参考答案(先写下你的预测再点开)
解题思路:KNN 预测前先把邻居按距离排序。本题 d 已按从小到大给出,对应标签依次为 0、1、1、0、1。K=1 时只看最近的那个邻居(d=0.4),其标签为 0,预测类别为 0。K=3 时取最近三个邻居,标签为 (0, 1, 1),多数票计数 0 类 1 票、1 类 2 票,预测类别为 1。K 从 1 增至 3 时预测发生改变,原因在于:K=1 的结果完全由单一最近邻决定,对个别样本极其敏感(这正是本章小结“易错点”第 2 条所说的“K 取 1 易过拟合”);K=3 引入投票平均,第 2、3 近邻都是 1 类,两票对一票把预测“翻”了过来——说明同一个测试点在不同 K 下的预测可能截然不同,K 的取值本身就是需要用交叉验证来选择的超参数。
# 验证脚本:手工模拟KNN的最近邻投票 import numpy as np # 导入NumPy数值计算库 neighbor_distances = np.array([0.4, 0.9, 1.1, 1.5, 2.3]) # 5个近邻的距离(从小到大) neighbor_labels = np.array([0, 1, 1, 0, 1]) # 按同一顺序对应的邻居标签 sorted_labels = neighbor_labels[np.argsort(neighbor_distances)] # 按距离升序重排标签,保证投票取的是最近邻居 print('按距离排序后的邻居标签:', sorted_labels) # 排序后依次为0,1,1,0,1 k1_prediction = sorted_labels[0] # K=1:直接取最近邻的标签 k3_votes = np.bincount(sorted_labels[:3]) # K=3:统计最近三个邻居各类的票数 print('K=1 预测类别:', k1_prediction) # 输出K=1的预测 print('K=3 票数统计(下标为类别):', k3_votes) # 0类1票、1类2票 print('K=3 预测类别:', k3_votes.argmax()) # 多数票决定K=3的预测预期输出(本机 peter 环境实际运行结果,具体以平台运行结果为准):
按距离排序后的邻居标签: [0 1 1 0 1] K=1 预测类别: 0 K=3 票数统计(下标为类别): [1 2] K=3 预测类别: 1回扣本章:对应本章小结“要点”第 1 条(K 个最近邻居按多数票决定类别)与“易错点”第 2 条(K 取 1 易过拟合,K 的大小由交叉验证权衡)。
概念辨析:KNN 与 K-Means 都依赖距离计算,却分属监督学习与无监督学习,二者在”是否有标签”“K 的含义”“输出是什么”上有何区别?
参考答案(点开前请先独立完成)
解题思路:从三个维度逐一对照。第一,是否有标签:KNN 是监督学习,训练数据必须带类别标签(如本章的“是否购买”Purchased 列),预测时按已知标签投票;K-Means 是无监督学习(见第 章节 30 章),输入只有特征矩阵,算法按距离结构自行发现分组,全程不需要标签。第二,K 的含义:KNN 的 K 是参与投票的最近邻居个数,控制的是决策的“平滑程度”;K-Means 的 K 是要把样本划分成几个簇,控制的是分组数量。第三,输出是什么:KNN 对每个新样本输出一个类别标签(离散值),回答“这个样本属于哪一类”;K-Means 输出的是全体样本的簇编号与各簇质心,回答“这批样本内部能分成哪几组”。两者的共同点是都以距离(如欧氏距离)度量相似,因而都对特征量纲敏感、都需要先做标准化。
回扣本章:对应本章小结“易错点”第 4 条——两个“K”含义完全不同,不要混淆;本章“分析结论”与第 章节 30 章的对照也正源于此。
概念辨析:
transfer1.fit_transform(x_train)与transfer1.transform(x_test)为何不能互换使用?如果对两者都用fit_transform,评估结果会怎样偏移?参考答案(点开前请先独立完成)
解题思路:
fit与transform是两个动作:fit从数据中“学参数”(StandardScaler学的是均值与标准差),transform用已学的参数做变换。fit_transform(x_train)等于先在训练集上学均值方差、再用这组参数标准化训练集;transform(x_test)则沿用训练集学到的参数去变换测试集——测试集必须站在训练集的“同一把尺子”上,模型面对的才是真正未见过的数据口径。两者不能互换:若对测试集也用fit_transform,测试集会用它自己的均值与标准差做标准化,产生两重后果。其一,口径错位:训练与预测使用了不同的标准化参数,距离计算(年龄与薪资的相对权重)随之改变,KNN 找到的“最近邻”可能不再是同一批样本;其二,信息泄露:测试集的分布信息(均值、方差)提前进入了建模流程,相当于让考试题参与了复习,测试集准确率往往被虚高估计,失去评估意义。这正是本章小结“易错点”第 1 条强调的规则:训练集用fit_transform,测试集只用transform。回扣本章:对应本章小结“易错点”第 1 条(对测试集误用
fit_transform导致口径错误与信息泄露)与“要点”第 2 条(标准化是距离型算法的必要预处理)。变式任务:把平台代码中的
param_dict改为{"n_neighbors": [1, 3, 5, 7, 9, 11, 13, 15, 17, 19]},观察最优 K 与交叉验证得分是否变化;再尝试把特征改为仅用Age一列重跑,准确率下降说明了什么?参考答案(点开前请先独立完成)
解题思路:改造思路——两处改动各自独立,应分别对照。第一处只改
param_dict一行,把候选 K 集合从 7 个奇数扩到 10 个(新增 15、17、19):网格搜索仍会在候选集内选交叉验证得分最高的 K,若原最优 K 依然胜出,则最优 K 与得分不变;只有当更大的 K 在 4 折交叉验证中得分更高时,最优 K 才会变大、得分才会更新——扩大候选集本身不会“必然”改变结果,它只是扩大了搜索范围。第二处把x = data[["Age","EstimatedSalary"]]改为x = data[["Age"]],只保留年龄一个特征:KNN 依赖特征空间中的距离判别,删去薪资维度等于丢掉一部分判别信息,测试集准确率大概率下降;若下降明显,说明估计薪资对购买意愿有增量预测力,两个特征联合的距离结构比单一年龄更有效。结构性判读——第一组对照看best_estimator_中的n_neighbors与best_score_是否变动;第二组对照看测试集准确率与混淆矩阵中错分样本是否增多,并注意仅用 Age 时标准化只对一列做,量纲问题自动消失,性能变化只能归因于信息量减少而非量纲。预期输出:数据文件由教学平台内置(04_Social_Network_Ads.csv 本地没有),本机无法复现,具体数值以平台运行结果为准。
注意:以上为变式的改造思路;列表 31.1 对应的平台原始代码块仍须原样输入教学平台,不要用本变式替换。
回扣本章:对应本章小结“要点”第 3 条(K 由网格搜索加交叉验证选择,不凭感觉拍定)与“易错点”第 2 条(K 的大小是精度与平滑之间的权衡)。