如何在NumPy中无需使用for循环获取二维数组每列最大值对应的行索引
获取每列最大值对应的行索引(无for循环)
针对你的需求,我有一个不用显式for循环的方案,能直接得到按列分组的最大值行索引,方便后续单独处理每一列。下面是具体实现:
import numpy as np A = np.array([[2,3,4], [2,0,4], [1,3,7]]) # 1. 计算每列的最大值 col_max = np.amax(A, axis=0) # 2. 获取所有最大值的行、列索引对 rows, cols = np.where(A == col_max) # 3. 按列索引排序,确保同一列的行索引连续 sorted_order = np.argsort(cols) sorted_rows = rows[sorted_order] sorted_cols = cols[sorted_order] # 4. 拆分得到每列对应的行索引组 split_points = np.where(np.diff(sorted_cols) != 0)[0] + 1 max_rowIndices_perColumn = np.split(sorted_rows, split_points) # 如果需要转为object类型的numpy数组(和你期望的格式一致) max_rowIndices_perColumn = np.array(max_rowIndices_perColumn, dtype=object)
运行后max_rowIndices_perColumn的结果就是:
array([array([0, 1]), array([0, 2]), array([2])], dtype=object)
方案解释:
- 第一步用
np.amax(A, axis=0)快速计算每列的最大值,这是向量化操作,没有循环。 np.where(A == col_max)会返回所有满足“元素等于对应列最大值”的位置,其中rows是行索引数组,cols是对应的列索引数组。比如你的例子中,rows是[0,1,0,2,2],cols是[0,0,1,1,2]。- 因为
np.where返回的索引是按行优先遍历的,所以我们需要按cols排序,让同一列的行索引聚在一起。 - 最后用
np.split根据列索引的分界点拆分排序后的行索引,就得到了每列独立的最大值行索引组,完全符合你后续单独处理每一列的需求。
为什么比你原来的方案更合适?
你之前用np.where(A== np.amax(A,axis=0))得到的是两个扁平数组,没法直接对应到每一列。而这个方案通过排序和拆分,直接把每列的索引分组,你可以像访问列表一样单独取出某一列的索引(比如max_rowIndices_perColumn[0]就是第一列的最大值行索引),非常方便后续处理。
内容的提问来源于stack exchange,提问作者user18089872
相关产品推荐
相关产品推荐

