查找2D ndarray中最后一列值最大的n个子数组的索引
如何找到2D NumPy数组中最后一列数值最大的n个子数组的索引
假设你已经有了如下的2D NumPy数组:
import numpy as np arr = np.array([ [ 0, 0, 1024, 1024, 572911], [ 0, 0, 822, 1024, 316522], [ 700, 0, 95, 40, 2189], [ 986, 72, 38, 37, 648], [ 933, 158, 50, 51, 823], [ 720, 172, 39, 85, 1806], [ 863, 204, 161, 255, 12329], [ 407, 275, 55, 34, 838], [ 522, 414, 234, 161, 12692], [ 861, 547, 49, 73, 1373], [ 972, 564, 52, 49, 1252], [ 929, 577, 25, 42, 376], [ 703, 608, 37, 64, 1082], [ 565, 612, 106, 152, 6278], [ 556, 615, 468, 409, 113116], [ 114, 674, 43, 44, 478], [ 155, 733, 150, 53, 3505], [ 57, 991, 23, 30, 358] ])
方法一:使用argsort(简单直观)
np.argsort()会返回数组元素从小到大排序后的索引,取最后n个索引并反转,就能得到最大的n个值对应的原索引:
n = 2 # 提取数组的最后一列 last_col = arr[:, -1] # 获取排序后的索引,倒序取前n个 top_n_indices = np.argsort(last_col)[-n:][::-1] print(top_n_indices) # 输出: [0 1]
方法二:使用argpartition(高效处理大数组)
如果数组规模很大,完全排序会浪费性能,np.argpartition()可以只将最大的n个元素移到数组末尾,不需要完全排序,效率更高:
n = 2 last_col = arr[:, -1] # 定位最大的n个元素的索引位置 top_n_indices = np.argpartition(last_col, -n)[-n:] # 对这n个索引对应的最后一列值排序,得到从大到小的索引顺序 top_n_indices = top_n_indices[np.argsort(-last_col[top_n_indices])] print(top_n_indices) # 输出: [0 1]
验证结果
拿到索引后,可直接提取对应的子数组:
top_n_subarrays = arr[top_n_indices] print(top_n_subarrays) # 输出: # [[ 0 0 1024 1024 572911] # [ 0 0 822 1024 316522]]
内容的提问来源于stack exchange,提问作者Apollo
相关产品推荐
相关产品推荐

