如何简化3D NumPy数组各层最大值所在列的索引获取代码?
简化3D NumPy数组各层最大值列索引的获取
咱们现在有个形状为 (n_layers, n_rows, n_cols) 的3D NumPy数组,需要给每个二维层(也就是每个矩阵)找出整个层里最大值所在的列索引。
示例输入
import numpy as np arr = np.array([[[0.05, 0.05, 0.9 ], [0.4 , 0.5 , 0.1 ], [0.7 , 0.2 , 0.1 ], [0.1 , 0.2 , 0.7 ]], [[0.98, 0.01, 0.01], [0.2 , 0.3 , 0.95], [0.33, 0.33, 0.34], [0.33, 0.33, 0.34]]])
预期输出
array([2, 0])
简化实现方案
不用搞复杂的中间步骤,直接通过两次维度操作就能搞定:
# 先拿到每个层里最大值在展平后的索引 flat_idx = arr.argmax(axis=(1, 2)) # 转成列索引:展平索引对列数取余即可 indices = flat_idx % arr.shape[-1]
嫌麻烦的话,一行代码就能写完:
indices = arr.argmax(axis=(1, 2)) % arr.shape[-1]
原理说明
arr.argmax(axis=(1, 2)):针对第0轴的每个二维层,直接沿着行和列的维度找最大值的位置,得到的是把二维层展平成一维后的索引(比如第一层的0.9在展平后的第2位,第二层的0.98在第0位)。- 用这个索引对列数取余,是因为展平后的索引计算公式是
行索引 * 列数 + 列索引,取余后就直接拿到列索引了。
如果你的场景里可能存在多个相同的最大值,还可以用这种方式(会返回所有最大值的列索引):
# 先找出每个层的最大值,保留维度方便后续匹配 max_vals = arr.max(axis=(1, 2), keepdims=True) # 找出所有等于最大值的位置,提取列索引 indices = np.where(arr == max_vals)[2]
要是每个层只有一个最大值,结果和上面的方法完全一致。
和原代码对比
原代码先找每行的最大值索引,再找这些行索引里对应最大值的位置,步骤绕了一圈。上面的简化方法直接针对整个层操作,代码更简洁,运行效率也更高。
内容的提问来源于stack exchange,提问作者Riccardo Bucco
相关产品推荐
相关产品推荐

