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

如何将K-means散点图与ground truth散点图进行对比?

K-Means聚类结果与真实分组的对比可视化

需求说明

我有一个包含height_mean(x)、weight_mean(y)变量的数据集,另有一列指定样本所属的真实分组(共11组)。已通过K-means完成聚类并绘制了聚类结果图,现在需要生成一张对比散点图:将聚类结果与真实分组匹配的点标记为绿色,不匹配的标记为红色。

原K-Means实现代码

import matplotlib.pyplot as plt
import numpy as np
from sklearn.cluster import KMeans  
import pandas as pd

df = pd.read_csv("FILENAME")
print(df)

x = df['height_mean']
y = df['weight_mean']

points = df[['height_mean', 'weight_mean']].values

n_clusters = 11

kmeans = KMeans(n_clusters=n_clusters)
kmeans.fit(points)

labels = kmeans.labels_
centers = kmeans.cluster_centers_

plt.scatter(x, y, c=labels, cmap='viridis')
plt.scatter(centers[:, 0], centers[:, 1], c='red', marker='x', s=100)
plt.xlabel('height')
plt.ylabel('weight')
plt.title("K-Means Clustering")

plt.show()
print(df)

修改后的完整代码(含对比可视化)

import matplotlib.pyplot as plt
import numpy as np
from sklearn.cluster import KMeans  
import pandas as pd

df = pd.read_csv("FILENAME")
# 替换为你数据集中真实分组的列名,比如df['true_group']
true_labels = df['z'].values  

x = df['height_mean']
y = df['weight_mean']

points = df[['height_mean', 'weight_mean']].values

n_clusters = 11

# 拟合KMeans模型
kmeans = KMeans(n_clusters=n_clusters)
kmeans.fit(points)
labels = kmeans.labels_
centers = kmeans.cluster_centers_

# 创建画布,同时展示三张对比图
plt.figure(figsize=(15, 5))

# 1. K-Means聚类结果图
plt.subplot(1, 3, 1)
plt.scatter(x, y, c=labels, cmap='viridis')
plt.scatter(centers[:, 0], centers[:, 1], c='red', marker='x', s=100)
plt.xlabel('height')
plt.ylabel('weight')
plt.title("K-Means Clustering")

# 2. 真实分组分布图
plt.subplot(1, 3, 2)
plt.scatter(x, y, c=true_labels, cmap='viridis')
plt.xlabel('height')
plt.ylabel('weight')
plt.title("Ground Truth Groups")

# 3. 聚类结果与真实分组对比图
plt.subplot(1, 3, 3)
# 生成颜色数组:匹配为绿色,不匹配为红色
color_list = ['green' if l == t else 'red' for l, t in zip(labels, true_labels)]
plt.scatter(x, y, c=color_list)
plt.xlabel('height')
plt.ylabel('weight')
plt.title("K-Means vs Ground Truth")

plt.tight_layout()
plt.show()

核心修改说明

  • 提取真实分组标签:必须将代码中的'z'替换为你数据集中真实分组列的实际名称
  • 颜色映射逻辑:通过列表推导式遍历聚类标签与真实标签,逐一判断匹配情况并赋值对应颜色
  • 子图布局:用plt.subplot将三张图放在同一画布中,便于直观对比聚类效果与真实分组的差异

内容的提问来源于stack exchange,提问作者dburbank

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.11 23:09:56