You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何正确计算matplotlib绘制的PCA散点图中各类别质心

问题场景

我通过绘制特征提取得到的前两个PCA主成分(PCA1、PCA2)生成了对应的散点图。
PCA类别散点图
上图为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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.30 14:42:24