如何使用计算得到的数组对列进行索引?附Python代码示例
解决用索引数组提取每行最小距离的问题
嘿,我懂你在做的事——这明显是K-means聚类流程里,计算每个样本到最近质心距离的环节对吧?你已经算出了所有点到质心的距离矩阵dists,也找到了每个点对应的最近质心索引dmini,现在卡在怎么用这个索引数组准确提取每行的最小距离,我来给你捋清楚。
先说说你现有代码的问题
你写的mindists = dists[:, dmini]其实达不到预期效果:因为dmini是一维数组(形状(6,)),这样索引会把dmini里的每个值当作列索引,重复提取对应列,最后得到一个(6,6)的数组,而不是每个点对应最小距离的一维数组。
两种正确的实现方法
方法1:用np.take_along_axis(清晰直观)
这个函数专门用来沿着指定轴,根据索引数组提取元素:
# 把dmini转换成二维数组,匹配dists的维度 dmini_2d = dmini[:, np.newaxis] # 沿着列轴提取每个行对应的最小距离元素 mindists = np.take_along_axis(dists, dmini_2d, axis=1).flatten()
解释:dmini[:, np.newaxis]把一维的dmini变成(6,1)的二维数组,这样take_along_axis就能精准定位到每行中对应最小距离的那个元素,最后flatten()把结果转回一维数组,方便后续使用。
方法2:用行索引+列索引配对(更简洁高效)
利用NumPy的广播特性,直接生成行索引数组和dmini的列索引数组配对:
# 生成行索引:0到5,和每个样本一一对应 row_indices = np.arange(len(dists)) # 提取每行中dmini指定列的元素 mindists = dists[row_indices, dmini]
这是最常用的写法,一行代码搞定,结果直接是你想要的一维数组(形状(6,))。
验证结果
用你的示例数据,运行上述代码后,mindists会输出:
array([0., 0., 1., 0., 1., 2.])
完全符合每个点到最近质心的距离:
- 第1个点
[0,0]到质心[0,0]距离0 - 第2个点
[0,1]到质心[0,1]距离0 - 第3个点
[0,2]到质心[0,3]距离1 - 以此类推...
内容的提问来源于stack exchange,提问作者Antonio Sesto
相关产品推荐
相关产品推荐

