np.argpartition返回前3个索引与np.argsort不匹配问题咨询
np.argpartition返回结果与np.argsort不一致的原因及解决方法
np.argpartition执行的是分块划分操作而非全排序,它的核心规则是:
- 当你指定
kth参数为N时,仅保证返回的索引数组中,下标为N的位置对应原数组排序后的第N+1小的元素 - 下标小于N的位置对应的元素,全部小于下标为N的元素;下标大于N的位置对应的元素,全部大于下标为N的元素
kth前后两个区间内的元素顺序不做任何排序保证,这就是你看到前3个索引和np.argsort结果不一致的核心原因。
你当前使用kth=3的调用方式,返回结果的前3个索引确实对应每行最小的3个元素,只是这3个索引没有按对应元素的大小排序,如果你仅需要提取最小的3个元素、不要求这3个元素本身有序,那原有代码的前3个索引已经可以满足需求。
如果需要前3个索引也按对应元素从小到大有序,同时保留比全排序更高的效率,可以采用「先分块拿前k个,再局部排序」的方案,代码实现如下:
import numpy as np example = np.array([[5,6,7,3,4],[1,2,3,7,5],[6,7,4,2,3],[1,2,3,5,9],[2,3,6,1,2,]]) k = 3 # 第一步:用argpartition获取每行最小的k个元素的索引(区间内无序) part_idx = np.argpartition(example, kth=k-1, axis=1)[:, :k] # 第二步:对每行的k个索引做局部排序,得到有序的前k个索引 row_indices = np.arange(example.shape[0])[:, np.newaxis] sorted_topk_idx = part_idx[row_indices, np.argsort(example[row_indices, part_idx])] print(sorted_topk_idx)
运行输出结果和np.argsort返回的前3个索引完全一致:
[[3 4 0] [0 1 2] [3 4 2] [0 1 2] [3 0 4]]
该方案的时间复杂度为每行O(n + k logk),当k远小于每行的元素长度时,效率远高于全排序的O(n logn),符合性能优化的需求。
内容的提问来源于stack exchange,提问作者Дима Лимонов
相关产品推荐
相关产品推荐

