K-means聚类标签正确但质心偏离聚类中心的原因分析
K-means聚类质心可视化异常问题
疑问:为何K-means聚类标签显示正确,但质心却未靠近聚类中心?具体表现为可视化图中质心挤在左下角,但图中存在三类聚类标签。
相关代码、数据输出
print(df.info()) print(df) preprocessor = ColumnTransformer( transformers=[ ('cat', OneHotEncoder(), ['State']) ], remainder='passthrough' ) kmeans = Pipeline([ ('preprocessor', preprocessor), ('kmeans', KMeans(n_clusters = 3, random_state=0, n_init = "auto")) ]).fit(df) labels = kmeans['kmeans'].labels_ print("Cluster Labels:", labels) centroids = kmeans['kmeans'].cluster_centers_ print("Centroids:", centroids) labels = kmeans['kmeans'].labels_ centroids = kmeans['kmeans'].cluster_centers_ plt.scatter(df['SumOfTotalPrice'], df['State'], c = labels) plt.scatter(centroids[:, 0], centroids[:, 1], marker='*', s=200, c='#050505') plt.xlabel('SumOfTotalPrice') plt.ylabel('State') plt.show()
数据输出:
<class 'pandas.core.frame.DataFrame'> RangeIndex: 11 entries, 0 to 10 Data columns (total 2 columns): # Column Non-Null Count Dtype --- ------ -------------- ----- 0 State 11 non-null object 1 SumOfTotalPrice 11 non-null float64 dtypes: float64(1), object(1) memory usage: 304.0+ bytes None State SumOfTotalPrice 0 AK 1.063432e+07 1 CA 4.172891e+07 2 IL 2.103149e+07 3 IN 2.270681e+08 4 KY 4.144238e+07 5 ME 2.057557e+07 6 MI 4.216375e+07 7 OH 7.970354e+08 8 PA 2.158148e+07 9 SD 1.025623e+07 10 TX 2.061534e+07 Cluster Labels: [0 0 0 2 0 0 0 1 0 0 0] Centroids: [[1.11111111e-01 1.11111111e-01 1.11111111e-01 0.00000000e+00 1.11111111e-01 1.11111111e-01 1.11111111e-01 0.00000000e+00 1.11111111e-01 1.11111111e-01 1.11111111e-01 2.55588301e+07] [0.00000000e+00 0.00000000e+00 0.00000000e+00 0.00000000e+00 0.00000000e+00 0.00000000e+00 0.00000000e+00 1.00000000e+00 0.00000000e+00 0.00000000e+00 0.00000000e+00 7.97035399e+08] [0.00000000e+00 0.00000000e+00 0.00000000e+00 1.00000000e+00 0.00000000e+00 0.00000000e+00 0.00000000e+00 0.00000000e+00 0.00000000e+00 0.00000000e+00 0.00000000e+00 2.27068150e+08]]
可视化结果

问题原因
- 可视化维度错误:独热编码后,数据集变成12维(11个State的编码列 + 1个SumOfTotalPrice列),但你可视化时取了质心的前两列(
centroids[:,0], centroids[:,1]),这两列是State独热编码的均值(比如第一类质心的前11列都是1/9,对应9个不同State的平均),这些小数值点自然挤在左下角,和你可视化用的SumOfTotalPrice完全无关。 - 类别特征不适合直接参与K-means:State是类别型特征,独热编码后用欧氏距离做聚类,会让聚类结果偏向类别分布,而且质心的类别特征列是均值,无法对应到原始的State类别,导致可视化时无法匹配到正确的y轴位置。
修复方案
方案1:仅用数值特征做聚类(更合理)
如果State只是用来标注样本,不是聚类的特征,只基于SumOfTotalPrice做聚类更符合业务逻辑:
import pandas as pd from sklearn.cluster import KMeans import matplotlib.pyplot as plt # 基于数值特征训练K-means kmeans = KMeans(n_clusters=3, random_state=0, n_init="auto").fit(df[['SumOfTotalPrice']]) labels = kmeans.labels_ centroids = kmeans.cluster_centers_ # 可视化样本 plt.scatter(df['SumOfTotalPrice'], df['State'], c=labels) # 绘制质心:将质心的x值对应到SumOfTotalPrice的均值,y值对应聚类中任意一个State的位置 state_list = df['State'].unique().tolist() for i, centroid in enumerate(centroids): # 取该聚类第一个样本的State位置作为质心的y坐标 cluster_state = df[labels == i]['State'].iloc[0] y_pos = state_list.index(cluster_state) plt.scatter(centroid[0], y_pos, marker='*', s=200, c='#050505') plt.xlabel('SumOfTotalPrice') plt.ylabel('State') plt.show()
方案2:保留类别特征并正确映射质心
如果必须将State作为聚类特征,需要从质心的12维数据中提取SumOfTotalPrice列(最后一列),并找到质心对应的State类别:
import pandas as pd from sklearn.compose import ColumnTransformer from sklearn.preprocessing import OneHotEncoder from sklearn.pipeline import Pipeline from sklearn.cluster import KMeans import matplotlib.pyplot as plt # 原始预处理和聚类流程 preprocessor = ColumnTransformer( transformers=[ ('cat', OneHotEncoder(), ['State']) ], remainder='passthrough' ) kmeans = Pipeline([ ('preprocessor', preprocessor), ('kmeans', KMeans(n_clusters=3, random_state=0, n_init="auto")) ]).fit(df) labels = kmeans['kmeans'].labels_ centroids = kmeans['kmeans'].cluster_centers_ # 可视化样本 plt.scatter(df['SumOfTotalPrice'], df['State'], c=labels) # 获取独热编码的State列名,匹配质心对应的State encoder = preprocessor.named_transformers_['cat'] state_feature_names = encoder.get_feature_names_out(['State']) state_list = df['State'].unique().tolist() for centroid in centroids: # 找到独热编码列中值最大的索引,对应最具代表性的State state_idx = centroid[:-1].argmax() state = state_feature_names[state_idx].replace('State_', '') y_pos = state_list.index(state) # 质心的x值是SumOfTotalPrice的均值(最后一列) plt.scatter(centroid[-1], y_pos, marker='*', s=200, c='#050505') plt.xlabel('SumOfTotalPrice') plt.ylabel('State') plt.show()
内容的提问来源于stack exchange,提问作者nicomp
相关产品推荐
相关产品推荐

