基于KMeans聚类修改3D散点图数据点的轮廓颜色
KMeans聚类3D散点图颜色匹配方案
问题描述
我有一个包含RGB值的数据集,已将其绘制成3D散点图。运行KMeans聚类后能正常显示数据点,但需要实现两个效果:
- 每个聚类中心使用不同颜色
- 每个数据点的轮廓颜色与其所属聚类中心的颜色匹配
当前运行代码后所有聚类中心均为红色,无法实现需求。
原代码
# Import scikit-learn, a machine learning library. from sklearn.cluster import KMeans # Load our classifier. num_clusters = 5 # You can change this if you want more/less than 5 bins! # Fit to our data. clusters_by_set = {} for dataset_name, points in datasets.items(): kmeans_cluster = KMeans(n_clusters=num_clusters, random_state=0) clusters_by_set[dataset_name] = kmeans_cluster.fit(points[['red','green','blue']]) # Add cluster centers to scatter. for dataset_name, points in datasets.items(): clusters = pd.DataFrame(clusters_by_set[dataset_name].cluster_centers_, columns=['red','green','blue']) fig = create_3d_scatter(points, dataset_name) # # Maybe color all points to match their cluster color? fig.add_trace(dict(type='scatter3d', x=clusters['red'], y=clusters['green'], z=clusters['blue']) ) pio.show(fig)
解决方案
核心改动是获取聚类标签、生成对应颜色映射,分别处理数据点轮廓和聚类中心的颜色:
修改后的完整代码
from sklearn.cluster import KMeans import pandas as pd import plotly.io as pio import plotly.graph_objects as go from matplotlib.cm import get_cmap num_clusters = 5 clusters_by_set = {} # 完成聚类并保存模型、标签及原始数据 for dataset_name, points in datasets.items(): kmeans_cluster = KMeans(n_clusters=num_clusters, random_state=0) kmeans_cluster.fit(points[['red','green','blue']]) clusters_by_set[dataset_name] = { 'model': kmeans_cluster, 'labels': kmeans_cluster.labels_, 'points': points } # 生成聚类对应颜色(用tab10色板,区分度高) cmap = get_cmap('tab10', num_clusters) cluster_colors = [f'rgb({int(cmap(i)[0]*255)}, {int(cmap(i)[1]*255)}, {int(cmap(i)[2]*255)})' for i in range(num_clusters)] # 绘制带颜色匹配的3D散点图 for dataset_name, cluster_data in clusters_by_set.items(): kmeans_model = cluster_data['model'] labels = cluster_data['labels'] points = cluster_data['points'] # 初始化图表 fig = go.Figure() # 添加数据点:填充色和轮廓色均匹配所属聚类 fig.add_trace(go.Scatter3d( x=points['red'], y=points['green'], z=points['blue'], mode='markers', marker=dict( size=5, color=[cluster_colors[label] for label in labels], line=dict( color=[cluster_colors[label] for label in labels], width=2 ) ), name='数据点' )) # 添加聚类中心:每个中心用对应聚类颜色,大尺寸叉号突出 centers = kmeans_model.cluster_centers_ for i in range(num_clusters): fig.add_trace(go.Scatter3d( x=[centers[i][0]], y=[centers[i][1]], z=[centers[i][2]], mode='markers', marker=dict( size=12, color=cluster_colors[i], symbol='x' ), name=f'聚类中心{i+1}' )) # 设置图表布局 fig.update_layout( title=f'{dataset_name} - KMeans聚类3D散点图', scene=dict( xaxis_title='红色通道', yaxis_title='绿色通道', zaxis_title='蓝色通道' ), legend=dict(title='图例') ) pio.show(fig)
关键说明
- 通过
kmeans_model.labels_获取每个数据点的聚类归属 - 用
matplotlib的tab10色板生成区分度高的颜色,也可替换为自定义RGB列表(如['#FF5733', '#33FF57', '#3357FF']) - 数据点的
line.color参数设置为所属聚类颜色,实现轮廓匹配 - 聚类中心单独循环绘制,用大尺寸叉号marker与数据点区分
内容的提问来源于stack exchange,提问作者Anthony Albert
相关产品推荐
相关产品推荐

