如何用Numpy筛选第一列重复行并保留第二列最大值行
解决Numpy数组按第一列分组保留第二列最大值行的问题
嗨,这是个很常见的分组筛选需求,我给你两种实用的Numpy实现方法,都能完美解决你的问题:
方法一:利用排序+唯一值索引(高效简洁)
这种方法通过先排序让每组的目标行处于固定位置,再提取唯一值对应的行,代码更简洁且性能更好:
import numpy as np # 原始数组 arr = np.array([[1,2], [2,9], [2,8], [2,7], [1,11], [3,6]]) # 按第一列升序、第二列降序排序,这样每个第一列分组的第一个元素就是第二列最大的行 sorted_arr = arr[np.lexsort((-arr[:, 1], arr[:, 0]))] # 获取第一列每个唯一值第一次出现的索引(因为排序后最大值在最前面) _, unique_indices = np.unique(sorted_arr[:, 0], return_index=True) # 提取结果 result = sorted_arr[unique_indices] print(result)
运行后输出就是你想要的:
[[ 2 9] [ 1 11] [ 3 6]]
方法二:分组查找最大值索引(逻辑更直观)
如果想更清晰地看到分组筛选的过程,可以用这种方法,先对第一列分组,再逐个找到每组第二列的最大值对应的行:
import numpy as np arr = np.array([[1,2], [2,9], [2,8], [2,7], [1,11], [3,6]]) # 获取第一列的唯一值,以及每个元素对应的分组索引 unique_first_col, group_indices = np.unique(arr[:, 0], return_inverse=True) # 遍历每个分组,找到该组内第二列最大值的位置 max_positions = [] for i in range(len(unique_first_col)): # 筛选出当前分组的所有行 group_rows = arr[group_indices == i] # 找到第二列最大值的索引 max_idx = np.argmax(group_rows[:, 1]) # 记录原始数组中对应的行索引 max_positions.append(np.where(group_indices == i)[0][max_idx]) # 提取最终结果 result = arr[max_positions] print(result)
两种方法都能得到正确结果,第一种更适合处理大数据量的数组,因为排序和唯一值操作都是Numpy的底层优化实现,速度更快。
内容的提问来源于stack exchange,提问作者TomK
相关产品推荐
相关产品推荐

