如何查找NumPy二维数组首列指定值且次列最大的行索引
问题场景
存在一个形状为(n,2)的NumPy数组a,示例如下:
import numpy as np a = np.array([ [6, 185.153], [6, 9.50864], [1, 9.31425], [1, 16.4629], [6, 19.6971], [1, 2.02113], [1, 14.0193], [5, 2.92495], [3, 56.0731], [3, 77.6965], ])
需求为查找符合以下条件的行索引:
- 行的第一列取值等于指定值
M(示例中M=3) - 在所有第一列等于
M的行中,该行第二列的取值为该组最大值
注:原问题描述中给出的预期输出为索引8,但按示例数组实际计算,M=3分组下第二列最大值为77.6965,对应索引9;如果预期为8则是要找分组下第二列最小值。
原实现代码如下,运行无法得到正确结果:
indx_nonremoved=np.where([minimum_merge.max(axis=1) ==3 ])[1]
原有代码问题
- 变量引用错误:代码中使用了未定义的变量
minimum_merge,实际需要操作的数组是a - 逻辑完全偏离需求:
max(axis=1)是对每一行取行内最大值,将其和3比较的逻辑,是在筛选“行内最大值等于3”的行,和“第一列匹配M、同组第二列取极值”的需求没有关系 - 维度处理错误:给判断条件额外套了一层方括号,凭空增加了一个数组维度,后续取
[1]维度的索引也不符合逻辑。
正确实现代码
逻辑分三步即可,以取分组最大值、M=3为例:
M = 3 # 生成第一列等于M的行掩码 col1_match = a[:, 0] == M # 取出匹配行的第二列,计算最大值 group_max = a[col1_match, 1].max() # 匹配同时满足两个条件的行索引 res_index = np.where(col1_match & (a[:, 1] == group_max))[0] print(res_index) # 输出 [9]
如果同分组下不存在重复的最大值,直接取res_index[0]即可得到单个整数索引。
如果实际需求是取分组下第二列的最小值(对应示例预期输出8),只需要把代码中的.max()替换为.min()即可,此时运行输出为[8]。
内容的提问来源于stack exchange,提问作者david
相关产品推荐
相关产品推荐

