如何在K-Means代码中加入train_test_split计算accuracy_score并绘制质心
实现方案
首先明确注意点:K-Means输出的簇编号和真实标签编号没有天然对应关系,直接计算准确率会得到错误结果,我们需要先做标签对齐再计算指标。
完整修改后代码
from sklearn.cluster import KMeans from sklearn import metrics from sklearn.model_selection import train_test_split import numpy as np import matplotlib.pyplot as plt from scipy.optimize import linear_sum_assignment # 原始坐标数据 X = np.array([3, 1, 1, 2, 1, 6, 6, 6, 5, 6, 7, 8, 9, 8, 9, 9, 8]) Y = np.array([5, 4, 6, 6, 5, 8, 6, 7, 6, 7, 1, 2, 1, 2, 3, 2, 3]) data = np.array(list(zip(X, Y))).reshape(len(X), 2) # 构造真实标签:对应数据集天然划分的3个簇 true_labels = np.array([0]*5 + [1]*5 + [2]*7) # 拆分训练集、测试集,拆分比例可自行调整 X_train, X_test, y_train, y_test = train_test_split(data, true_labels, test_size=0.3, random_state=42) colors = ['b', 'g', 'c'] markers = ['o', 'v', 's'] # 用训练集拟合KMeans模型 model = KMeans(n_clusters=3, random_state=42).fit(X_train) centers = np.array(model.cluster_centers_) # 预测测试集聚类标签 y_pred = model.predict(X_test) # 标签对齐:用匈牙利算法匹配聚类标签和真实标签,保证准确率计算合理 def align_labels(y_true, y_pred): cm = metrics.confusion_matrix(y_true, y_pred) row_ind, col_ind = linear_sum_assignment(-cm) aligned_pred = np.zeros_like(y_pred) for i, j in zip(row_ind, col_ind): aligned_pred[y_pred == j] = i return aligned_pred aligned_y_pred = align_labels(y_test, y_pred) accuracy = metrics.accuracy_score(y_test, aligned_y_pred) print(f"测试集准确率:{accuracy:.2f}") # 绘图逻辑 plt.title(f'K-Means Centroids (Test Accuracy: {accuracy:.2f})') # 绘制训练集点(半透明展示) train_pred = model.predict(X_train) for i, l in enumerate(train_pred): plt.plot(X_train[i,0], X_train[i,1], color=colors[l], marker=markers[l], ls='None', alpha=0.6) # 绘制测试集点(带黑边区分) for i, l in enumerate(aligned_y_pred): plt.plot(X_test[i,0], X_test[i,1], color=colors[l], marker=markers[l], ls='None', edgecolor='k') # 绘制聚类质心 plt.scatter(centers[:,0], centers[:,1], marker="x", color='r', s=200, linewidths=3) plt.xlim([0, 10]) plt.ylim([0, 10]) plt.show()
关键改动说明
- 新增数据集拆分逻辑,默认按照7:3比例拆分训练集和测试集
- 新增标签对齐函数,解决聚类标签和真实标签序号不匹配的问题,保证准确率计算结果有效
- 绘图时区分了训练集(半透明)和测试集(带黑边),同时保留了红色质心标记,标题中直接展示测试集准确率
- 新增
random_state参数保证运行结果可复现
内容的提问来源于stack exchange,提问作者nicolaser55
相关产品推荐
相关产品推荐

