如何正确计算matplotlib绘制的PCA散点图中各类别质心
问题场景
我通过绘制特征提取得到的前两个PCA主成分(PCA1、PCA2)生成了对应的散点图。
上图为3个类别的PCA1(横轴)与PCA2(纵轴)散点图,原有绘图代码如下:
target_names = ['class_1', 'class_2', 'class_3'] plt.figure(figsize=(11, 8)) Xt = pca.fit_transform(X) plot = plt.scatter(Xt[:,0], Xt[:,1], c=y, cmap=plt.cm.jet, s=30, linewidths=0, alpha=0.7) #centers = kmeans.cluster_centers_ #plt.scatter(centers[:, 0], centers[:, 1], c=['black', 'green', 'red'], marker='^', s=100, #alpha=0.5); plt.legend(handles=plot.legend_elements()[0], labels=list(target_names)) plt.show()
需求为正确获取该图中每个类别对应的质心,所用数据集前几行如下:
Xt1 Xt2 y -107.988187 -23.70121 1 -128.578852 -20.222378 1 -124.522967 -25.298283 1 -96.222918 -25.028239 1 -95.152954 -23.94496 1 -113.275804 -26.563129 1 -101.803 -24.22359 1 -94.662469 -22.94211 1 -104.118882 -24.037226 1 439.765098 -101.532469 2 50.100362 -34.278841 2 -69.229603 62.178599 2 -60.915475 53.296491 2 64.797364 91.991527 2 -112.815192 0.263505 0 -91.287067 -25.207217 0 -74.181941 -2.457892 0 -83.273718 -0.608004 0 -100.881393 -22.387571 0 -107.861711 -15.848869 0 -85.866992 -18.79126 0 -53.96314 -28.885316 0 -59.195432 -3.373361 0
实现方法
核心逻辑
有监督标签下的类别质心,就是同一类别所有样本在两个PCA维度上的坐标平均值,不需要调用聚类算法计算,直接按标签分组求均值即可。
完整代码
import numpy as np # 原有PCA计算和绘图代码 target_names = ['class_0', 'class_1', 'class_2'] # 注意标签取值为0/1/2,名称顺序要和标签升序对应 plt.figure(figsize=(11, 8)) Xt = pca.fit_transform(X) plot = plt.scatter(Xt[:,0], Xt[:,1], c=y, cmap=plt.cm.jet, s=30, linewidths=0, alpha=0.7) # 计算各标签对应质心 class_centers = [] for label in np.unique(y): # 筛选当前标签的所有样本 label_mask = (y == label) samples = Xt[label_mask] # 沿样本轴求两个PCA维度的均值 center = np.mean(samples, axis=0) class_centers.append(center) class_centers = np.array(class_centers) # 可选:将质心绘制在原图上验证效果 center_plot = plt.scatter(class_centers[:,0], class_centers[:,1], c=['black','green','red'], marker='^', s=100, alpha=0.8) plt.legend(handles=[*plot.legend_elements()[0], center_plot], labels=[*target_names, 'class_center']) plt.show()
注意点
- 之前注释的
kmeans.cluster_centers_是无监督聚类生成的中心,和已有真实标签的类别质心不是同一个概念,二者结果大概率不重合,不要混用。 - 最终得到的
class_centers数组按标签升序存储质心坐标,第0行是标签0的(PCA1, PCA2)坐标,第1行是标签1的坐标,以此类推,可直接打印查看具体数值。 - 要保证
target_names的顺序和标签升序顺序一致,否则图例和实际类别、质心颜色会匹配错误。
内容的提问来源于stack exchange,提问作者Joe
相关产品推荐
相关产品推荐

