如何获取多维数组每行最小值的索引?代码问题排查
问题分析与解决方案
首先,咱们先搞清楚你代码里的核心问题:np.unravel_index在这里被误用了。
为什么你的代码会出错?
你调用np.argmin(my_array, axis=2)得到的结果是形状为(2,2)的数组,每个元素代表对应(axis0, axis1)位置上,第三个维度(axis=2)里最小值的索引——比如正确的结果应该是:
array([[1, 0], [0, 0]])
这个结果里的每个数字,都是局部索引(只对应axis=2维度),而不是整个数组扁平化后的全局一维索引。
但np.unravel_index的作用是把全局一维索引转换成多维坐标。比如你传入的1会被解析成整个数组的第1个元素(对应(0,0,1),这刚好是对的),但传入的0会被解析成整个数组的第0个元素(对应(0,0,0)),而不是你期望的(0,1,0)或(1,0,0)。这就是为什么用my_array[idx_arr]取出的值完全不符合预期。
正确的实现方式
方法1:构造完整的多维索引数组
既然np.argmin(axis=2)已经给出了axis2的索引,我们只需要生成对应axis0和axis1的索引数组,再组合起来即可:
import numpy as np my_array = np.array([[[ 0.64, 0.49, 2.56], [ 7.84, 13.69, 21.16]], [[ 33.64, 44.89, 57.76], [ 77.44, 94.09, 112.36]]]) # 生成axis0和axis1的索引网格 i, j = np.indices(my_array.shape[:2]) # 获取axis2维度的最小值索引 axis2_idx = np.argmin(my_array, axis=2) # 组合成完整的多维索引 idx_arr = (i, j, axis2_idx) # 取出对应的值 result = my_array[idx_arr] print(result) # 输出:[[ 0.49 7.84] # [33.64 77.44]]
方法2:用np.take_along_axis直接取值(更简洁)
如果你只需要获取最小值,不需要索引的话,可以用np.take_along_axis一步到位:
min_vals = np.take_along_axis( my_array, np.argmin(my_array, axis=2)[..., np.newaxis], # 增加一个维度匹配axis=2 axis=2 ).squeeze() # 去掉多余的维度 print(min_vals) # 输出和上面一致
这样就能得到你预期的0.49、7.84、33.64、77.44这四个值啦~
内容的提问来源于stack exchange,提问作者Wang Lee
相关产品推荐
相关产品推荐

