如何用Scikit-learn K-Means计算并存储样本到最近聚类中心的距离
嘿,我来帮你搞定这个需求!要计算并存储Scikit-learn K-Means中每个样本到最近聚类中心的距离,其实有两种非常直接的方式,我结合你的代码结构给你详细说明:
计算每个样本到最近聚类中心的距离
当你训练完KMeans模型后,只需要利用模型自带的属性和方法,就能轻松获取到所需的距离值:
方法1:用KMeans的transform()方法(推荐)
KMeans的transform()方法会直接返回一个二维数组,其中每一行代表对应样本到所有聚类中心的欧氏距离。我们只需要对每一行取最小值,就能得到该样本到最近中心的距离。
方法2:手动计算距离(更灵活)
如果你偏好手动实现,可以用sklearn.metrics.pairwise.euclidean_distances函数,传入你的样本数据和模型训练好的聚类中心,同样取每行最小值即可。
完整代码示例
我调整了你的代码,补充了关键部分,同时修正了已弃用的samples_generator导入路径:
import numpy as np import matplotlib.pyplot as plt from sklearn.metrics.pairwise import euclidean_distances from sklearn.cluster import KMeans from sklearn.datasets import make_blobs # 注意:原samples_generator已被弃用,改用这个 def getDataFromTransaction(t): # 这里替换成你实际的交易数据特征提取逻辑 # 示例:返回2维随机特征,仅作演示 return np.random.rand(2) def learn(records): # 处理交易数据,转成特征数组 data = np.array([getDataFromTransaction(t) for t in records]) # 初始化并训练KMeans模型 kmeans = KMeans(n_clusters=3, random_state=42) kmeans.fit(data) # --- 方法1:用transform()快速获取距离 --- # 获取每个样本到所有聚类中心的距离 all_cluster_distances = kmeans.transform(data) # 取每行最小值,得到每个样本到最近中心的距离 nearest_distances = np.min(all_cluster_distances, axis=1) # --- 方法2:手动计算距离(验证用) --- # 用euclidean_distances计算样本到中心的距离 manual_distances = euclidean_distances(data, kmeans.cluster_centers_) manual_nearest_distances = np.min(manual_distances, axis=1) # 验证两种方法结果一致(浮点误差范围内) assert np.allclose(nearest_distances, manual_nearest_distances) # 存储距离:可以保存为numpy文件,或者转成CSV/Excel np.save('nearest_cluster_distances.npy', nearest_distances) # 也可以转成Pandas DataFrame保存(如果需要) # import pandas as pd # pd.DataFrame(nearest_distances, columns=['nearest_distance']).to_csv('distances.csv', index=False) # 可选:打印示例结果 print("前5个样本的最近聚类中心距离:", nearest_distances[:5]) print("训练得到的聚类中心:", kmeans.cluster_centers_) # 测试用例:生成100条模拟交易记录 if __name__ == "__main__": mock_records = [{} for _ in range(100)] # 模拟你的交易记录列表 learn(mock_records)
关键知识点说明
kmeans.cluster_centers_:这个属性存储了训练好的所有聚类中心的坐标,形状为(聚类数量, 特征维度)transform()方法:内部已经帮你实现了欧氏距离的计算,和手动调用euclidean_distances结果完全一致,只是更便捷np.min(..., axis=1):axis=1表示按行取最小值,也就是对每个样本,找出到所有中心里最近的那个距离- 存储方式:用
np.save()保存成numpy文件可以保留数组结构,方便后续直接加载使用;如果需要更通用的格式,转成CSV也是不错的选择
内容的提问来源于stack exchange,提问作者Chris_96
相关产品推荐
相关产品推荐

