如何使用NumPy获取每行中n个最小元素的索引?
获取NumPy数组每行n个最小元素的索引(含行号)
没问题,这事儿用NumPy可以轻松实现,我给你分步骤讲解并给出代码:
实现代码
import numpy as np n = 2 p1 = np.asarray([[20, 30, 10], [10, 20, 30], [30, 20, 10]]) # 1. 获取每行最小的n个元素的索引(按元素值从小到大排序) min_value_indices = np.argsort(p1, axis=1)[:, :n] # 2. 将索引按从小到大排序(匹配你示例中的输出格式) sorted_indices = np.sort(min_value_indices, axis=1) # 3. 生成行号数组,转为列向量以便合并 row_numbers = np.arange(p1.shape[0]).reshape(-1, 1) # 4. 合并行号和索引,得到最终结果 result = np.hstack([row_numbers, sorted_indices]) print(result)
运行这段代码后,输出结果正好是你期望的:
[[0 0 2] [1 0 1] [2 1 2]]
代码解释
np.argsort(p1, axis=1):对数组的每一行执行排序,返回的是每行元素从小到大排列对应的原始索引。axis=1指定按行操作。[:, :n]:截取每行的前n个索引,也就是该行最小的n个元素对应的索引。np.sort(min_value_indices, axis=1):将每行的索引按从小到大排序,这一步是为了匹配你示例中索引的排列顺序(如果不需要排序,直接跳过这一步即可)。np.arange(p1.shape[0]).reshape(-1, 1):生成从0开始的行号数组,reshape(-1,1)把它转为列向量,这样才能和索引数组(每行n列)横向合并。np.hstack([row_numbers, sorted_indices]):横向合并行号列和索引数组,得到每行开头为行号、后面为n个最小元素索引的结果。
如果你不需要索引按从小到大排序,只需要按元素值从小到大的顺序排列索引,那去掉np.sort那一步就行,此时第一行的结果会是[0,2,0],对应元素10(索引2)和20(索引0)的顺序。
内容的提问来源于stack exchange,提问作者Никита Михалков
相关产品推荐
相关产品推荐

