numpy的argmax()仅返回首个最大值索引,如何获取所有最大值索引?
获取Numpy数组中所有最大值的索引
np.argmax() 的设计逻辑就是仅返回第一个最大值所在的索引,所以你会得到0而不是所有符合条件的索引。要获取所有最大值的位置,可以按以下步骤操作:
- 先获取数组的最大值
- 通过布尔索引筛选出所有等于最大值的元素位置
方法一:使用 np.where()
import numpy as np arr = np.array([5,5,4]) max_value = arr.max() all_max_indices = np.where(arr == max_value)[0] print(all_max_indices) # 输出: [0 1]
np.where() 返回一个包含索引的元组,取第一个元素即可得到一维的索引数组。
方法二:使用 np.flatnonzero()
这是更简洁的写法,直接返回布尔数组中True对应的索引:
import numpy as np arr = np.array([5,5,4]) all_max_indices = np.flatnonzero(arr == arr.max()) print(all_max_indices) # 输出: [0 1]
两种方法都能精准获取所有最大值的索引,根据个人习惯选择即可。
内容的提问来源于stack exchange,提问作者Brian Smith
相关产品推荐
相关产品推荐

