为何在此处使用argsort而非argmin?Numpy习题相关疑问
使用numpy.argsort提取每行最接近0.5的元素
问题背景
这是一道SciPy讲义中的习题:生成一个10×3的[0,1]区间随机数数组,要求用abs和argsort找到每行中最接近0.5的元素,再通过花式索引提取数值。
基于argsort的实现代码
import numpy as np # 生成10×3的[0,1]随机数组 a = np.random.random((10, 3)) # 计算每个元素与0.5的绝对差 diff = np.abs(a - 0.5) # 对每行的差值排序,取最小差值对应的列索引(argsort结果的第一列) closest_col_idx = np.argsort(diff, axis=1)[:, 0] # 花式索引提取目标元素 result = a[np.arange(a.shape[0]), closest_col_idx]
代码说明
np.argsort(diff, axis=1)会对每行的差值数组进行排序,返回的是原始元素的索引序列——差值越小的元素,对应的索引越靠前。- 取每行索引序列的第一个元素
[:, 0],就得到了每行中与0.5差值最小的元素的列索引,这和argmin(axis=1)的结果等价,但符合习题要求的argsort用法。 - 最后用
np.arange(a.shape[0])生成行索引,结合列索引做花式索引,就能提取出每行最接近0.5的元素。
内容的提问来源于stack exchange,提问作者levant pied
相关产品推荐
相关产品推荐

