NumPy提取n维数组子数组前k小值索引及正确取值方法
按行提取二维数组前k小值索引及对应值的实现
问题场景
给定如下numpy二维数组:
import numpy as np X = np.array([[0.65108716, 0.72213542, 0.62142414, 0.80734795, 0.79485172, 0.83946013, 0.79192978, 0.76614672, 0. , 0.84231442], [0.71353155, 0.58493483, 0.76903558, 0.77678972, 0.71837986, 0.56127471, 0.72591233, 0.75986564, 0.83495295, 0.03016315]])
需求为从每个子数组(按行)中提取前k个最小值的索引,再根据索引取出对应的元素值。
问题复现
当k=1,即提取每行最小值索引时,使用如下代码:
top_n_indices = np.argsort(X)[:, :1]
得到的索引结果符合预期:
[[8], [9]]
但直接调用np.take(X, top_n_indices)提取对应值时,返回结果错误:
[[0. ], [0.84231442]]
预期正确结果为:
[[0. ], [0.03016315]]
要求不使用列表推导式实现该取值需求。
错误原因
np.take默认会将输入数组展平为一维数组后再匹配索引取值。返回的第二个错误值0.84231442,实际是数组展平后全局索引9对应的元素(即第一行最后一个值),并非第二行索引9对应的元素,因此结果不符合预期。
无列表推导式的实现方案
- 方案1:指定
np.take的轴向参数np.take原生支持axis参数,用来指定取值操作对应的数组维度,针对按行取值的场景,指定axis=1即可:
# 提取前k个最小值的索引,k=1时切片为:1,k为其他值时修改切片范围即可 top_n_indices = np.argsort(X)[:, :1] result = np.take(X, top_n_indices, axis=1)
运行后得到的结果与预期完全一致:
[[0. ], [0.03016315]]
- 方案2:使用numpy二维高级索引直接取值
numpy支持行索引+列索引的配对高级索引,先生成与列索引形状匹配的行索引数组,即可直接取值,无需调用np.take:
top_n_indices = np.argsort(X)[:, :1] # 生成形状为(行数,1)的行索引,和列索引形状对齐 row_indices = np.arange(X.shape[0]).reshape(-1, 1) result = X[row_indices, top_n_indices]
该方法返回结果与方案1完全一致。
效率提示:如果不需要获取最小值的索引,只需要拿到前k个最小值,直接使用
np.partition效率更高,该方法不需要对数组全排序,时间复杂度低于np.argsort,例如取每行前1小值可写为np.partition(X, kth=1, axis=1)[:, :1]。
内容的提问来源于stack exchange,提问作者NineWasps
相关产品推荐
相关产品推荐

