如何用numpy.argpartition获取数组中N个最小值(数组长度可能小于N)
解决方案
可以通过动态计算kth参数的方式,彻底避免分支判断,完全用numpy原生方法实现需求:
实现代码
def get_n_smallest(arr, N): k = np.minimum(N, arr.size) return arr[np.argpartition(arr, k)][:k]
验证示例
import numpy as np small_arr = np.array([3,1,4]) large_arr = np.array([3,1,4,5,0,2]) # 数组长度小于N的情况,返回整个数组 print(get_n_smallest(small_arr, 4)) # 输出结果为数组所有元素(顺序由argpartition特性决定) # 数组长度大于等于N的情况,返回前N个最小值 result = get_n_smallest(large_arr, 4) print(np.sort(result)) # 输出 [0 1 2 3],确认是4个最小值
原理说明
np.minimum(N, arr.size)会自动取N和数组长度中的较小值,确保传入argpartition的kth参数始终在数组索引的有效范围内,不会触发越界错误。- 最后通过切片
[:k],保证返回元素数量不超过N:当数组长度不足N时,k等于数组长度,切片后返回整个数组;当数组长度足够时,返回前N个经分区后的最小值。
内容的提问来源于stack exchange,提问作者blackbrandt
相关产品推荐
相关产品推荐

