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

如何用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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.25 07:13:12