1. 需求分析
会员客户聚类分析。
2. 数据说明
这是 Kaggle 经典客户分群教学数据集,共 200 条商场会员脱敏样本,专门用于无监督聚类(K-Means 为主)入门练习,为模拟商业数据,无真实用户隐私。
商场客户分群数据集来源:Mall Customer Segmentation Data
| 字段名 | 含义 | 数据类型 |
|---|---|---|
| CustomerID | 客户唯一编号 | 整数,主键 |
| Gender | 性别(Male/Female) | 分类变量 |
| Age | 客户年龄 | 数值 |
| Annual Income (k$) | 年收入(单位:千美元) | 数值 |
| Spending Score (1-100) | 商场消费打分,1 最低、100 最高,由消费金额 / 频次综合计算 | 核心聚类数值 |
数据特征
- 无缺失值、无异常脏数据,清洗成本极低;
- 包含人口属性(性别、年龄)与消费经济属性(收入、消费分)两类特征;
- 核心分析维度:年收入 & 消费得分,二者常用来划分典型客群。
3. 建模
import pandas as pd from matplotlib import pyplot as plt from sklearn.cluster import KMeans from sklearn.metrics import silhouette_score, calinski_harabasz_score3.1 加载数据
# 1. 获取数据 data = pd.read_csv('./data/customers.csv') data.head()3.2 数据预处理
# 2. 数据预处理 # 2.1 提取特征 x = data.iloc[:,3:5] x.head()3.3 特征工程(无)
3.4 模型训练
确定k值;训练
# 4. 模型训练 # 4.1 确定k值 # 定义sse_list, sc_list, 记录不同k值的评估效果 sse_list = [] sc_list = [] for k in range(2,20): estimator = KMeans(n_clusters=k, max_iter=100, random_state=1234) estimator.fit(x) y_pre = estimator.predict(x) sse_list.append(estimator.inertia_) sc_list.append(silhouette_score(x, y_pre)) # 画图(sse肘部图、轮廓系数sc图) # 两张图放在1个画布中 plt.figure(figsize=(6,3)) plt.subplot(1,2,1) plt.plot(range(2,20), sse_list, marker='o', label='SSE') plt.grid() plt.title('SSE') plt.subplot(1,2,2) plt.plot(range(2,20), sc_list, marker='o', label='SC') plt.grid() plt.title('SC') plt.tight_layout() plt.show()# 4.2 模型训练 # 根据sse肘部图和sc图,确定k值=5 estimator = KMeans(n_clusters=5, max_iter=100, random_state=1234) estimator.fit(x) y_pre = estimator.predict(x)3.5 模型评估
# 5. 模型评估 print('sc指数:', silhouette_score(x, y_pre)) # [-1,1], 越接近1聚类效果越好;>0.5聚类良好;负数为大量样本错分,聚类失效 print('CH指数:', calinski_harabasz_score(x, y_pre)) # [0,inf], 越接近inf聚类效果越好 print('SSE的负数:', estimator.score(x)) # 越接近0聚类效果越好# 画图 plt.figure(figsize=(6,5)) plt.scatter(x.iloc[:,0], x.iloc[:,1], c=y_pre) plt.xlabel(x.columns.tolist()[0]) plt.ylabel(x.columns.tolist()[1]) plt.scatter(estimator.cluster_centers_[:,0], estimator.cluster_centers_[:,1], marker='*', s=200, c='r') plt.title('K-Means') plt.show()