使用Numpy获取二维数组各列绝对值最大值及所有对应索引
问题描述
给定二维NumPy数组:
import numpy as np A = np.array([[ 100, -5, 3, 200], [ 20, -100, 4, 8], [ 12, -10, 10, 4], [-100, 80, 4, 14]])
需要获取每一列的绝对值最大值,以及该最大值对应的所有行索引。
目前已通过abs(A).max(axis=0)得到每列绝对值最大值:
max_col = abs(A).max(axis=0) print(max_col) # 输出: [100 100 10 200]
但使用abs(A).argmax(axis=0)只能得到每列第一个绝对值最大值的索引:
maxValueIndex = abs(A).argmax(axis=0) print(maxValueIndex) # 输出: [0 1 2 0]
无法捕获第一列中值为-100的行索引3,需要获取所有符合条件的索引。
解决方案
使用np.where()可以找出所有满足条件的元素索引,具体步骤如下:
- 计算每列的绝对值最大值(和你已有的代码一致)
- 利用广播机制,对比数组中每个元素的绝对值是否等于对应列的最大值
- 通过
np.where()提取所有符合条件的行、列索引,再按列整理结果
完整代码示例:
import numpy as np A = np.array([[ 100, -5, 3, 200], [ 20, -100, 4, 8], [ 12, -10, 10, 4], [-100, 80, 4, 14]]) # 计算每列绝对值最大值 max_col = abs(A).max(axis=0) # 获取所有绝对值等于列最大值的元素索引 row_indices, col_indices = np.where(abs(A) == max_col) # 按列分组整理行索引 col_to_rows = {} for col, row in zip(col_indices, row_indices): col_to_rows.setdefault(col, []).append(row) # 输出结果 for col_idx in range(A.shape[1]): print(f"列{col_idx}: 绝对值最大值={max_col[col_idx]},对应行索引={col_to_rows[col_idx]}")
输出结果
列0: 绝对值最大值=100,对应行索引=[0, 3] 列1: 绝对值最大值=100,对应行索引=[1] 列2: 绝对值最大值=10,对应行索引=[2] 列3: 绝对值最大值=200,对应行索引=[0]
说明
np.where(abs(A) == max_col)会返回两个数组:row_indices是符合条件的行索引集合,col_indices是对应的列索引集合- 因为
max_col是一维数组,和二维的abs(A)比较时会自动广播,实现逐列对比 - 最后通过字典按列分组,能清晰展示每一列的所有目标索引
内容的提问来源于stack exchange,提问作者pchi
相关产品推荐
相关产品推荐

