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

求助:如何在K-means聚类中根据最近质心为数据点分配颜色?

嘿,我来帮你解决这个K-means聚类点颜色匹配的问题!其实你的代码里已经有了关键的一步——calc_dist_euclidean函数返回的就是每个数据点对应的最近质心索引(也就是簇标签),我们只需要把这个标签和颜色映射起来就行。

我给你补全完整的K-means流程,并且加上颜色可视化的部分,你可以直接参考:

完整实现代码

import numpy as np
import pandas as pd
from matplotlib import pyplot as plt

def calc_dist_euclidean(vec_1, vec_2):
    # vec_1:所有数据点(形状:n_samples × n_features)
    # vec_2:当前质心(形状:k × n_features)
    distances = np.sqrt(((vec_1 - vec_2[:, np.newaxis]) ** 2).sum(axis=2))
    # 返回每个数据点对应的最近质心索引(簇标签)
    return np.argmin(distances, axis=0)

def kmeans(data, k, max_iter=100):
    # 随机初始化质心:从数据集中选k个不重复的点
    centroids = data[np.random.choice(data.shape[0], k, replace=False)]
    
    for _ in range(max_iter):
        # 给每个数据点分配对应的簇
        cluster_labels = calc_dist_euclidean(data, centroids)
        
        # 更新质心:计算每个簇的均值作为新质心
        new_centroids = np.array([data[cluster_labels == i].mean(axis=0) for i in range(k)])
        
        # 如果质心不再变化,提前终止迭代
        if np.all(centroids == new_centroids):
            break
        
        centroids = new_centroids
    
    return cluster_labels, centroids

# ---------------------- 测试示例 ----------------------
# 生成3组高斯分布的测试数据(你可以替换成自己的真实数据)
np.random.seed(42)
data = np.vstack([
    np.random.normal(0, 1, (100, 2)), 
    np.random.normal(5, 1, (100, 2)),
    np.random.normal(10, 1, (100, 2))
])

# 运行K-means聚类,设置簇数k=3
k = 3
cluster_labels, centroids = kmeans(data, k)

# ---------------------- 可视化部分 ----------------------
plt.figure(figsize=(8, 6))
# 绘制数据点:用cluster_labels作为颜色依据,自动分配不同颜色
scatter = plt.scatter(
    data[:, 0], data[:, 1], 
    c=cluster_labels,  # 核心:用簇标签映射颜色
    cmap='viridis',    # 可选:指定颜色映射方案,比如rainbow、plasma等
    alpha=0.6          # 让点半透明,避免重叠遮挡
)
# 绘制质心:用黑色大X标记,和数据点区分开
plt.scatter(
    centroids[:, 0], centroids[:, 1], 
    c='black', s=200, marker='X', label='Centroids'
)
# 添加图例、颜色条和标题
plt.legend()
plt.colorbar(scatter, label='Cluster Label')
plt.title('K-means Clustering: Points Colored by Nearest Centroid')
plt.xlabel('Feature 1')
plt.ylabel('Feature 2')
plt.show()

关键说明

  • 簇标签的作用:calc_dist_euclidean返回的cluster_labels是一个和数据点数量等长的数组,每个元素是0到k-1的整数,代表该点属于哪个簇(对应最近的质心)。
  • 颜色映射:在plt.scatter中设置c=cluster_labels,Matplotlib会自动把每个整数标签映射到不同的颜色,你可以通过cmap参数切换喜欢的配色方案。
  • 质心标记:把质心用不同形状/颜色突出显示,能更直观地看到每个簇的中心位置。

如果你用的是Pandas DataFrame存储数据,只需要把cluster_labels添加为DataFrame的一列,然后在绘图时指定对应列即可,比如:

df = pd.DataFrame(data, columns=['x', 'y'])
df['cluster'] = cluster_labels
plt.scatter(df['x'], df['y'], c=df['cluster'], cmap='viridis')

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.26 10:28:06