You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何用numpy.argmax从三维数组提取另一同形数组对应最大值

利用numpy.argmax结果从同形状数组提取对应值

已知两个形状相同的三维numpy数组ndarr1和ndarr2,通过np.argmax(ndarr1, axis=0)计算出了第一维度(axis=0)上最大值的索引数组argmax1,需要基于该索引数组从ndarr2中提取对应位置的元素,得到目标结果。

示例代码如下:

import numpy as np

# 初始化示例数组
ndarr1 = np.array([[[0.89, 0.79, 0.64],
                    [0.03, 0.53, 0.1 ]],

                   [[0.21, 0.76, 0.99],
                    [0.47, 0.08, 0.48]],

                   [[0.67, 0.94, 0.99],
                    [0.98, 0.75, 0.59]],

                   [[0.09, 0.96, 0.98],
                    [0.43, 0.98, 0.71]]])

argmax1 = np.argmax(ndarr1, axis=0)
# argmax1结果:array([[0, 3, 1], [2, 3, 3]], dtype=int64)

ndarr2 = np.array([[[0.79, 0.72, 0.82],
                    [0.25, 0.7 , 0.56]],

                   [[0.46, 0.11, 0.31],
                    [0.55, 0.76, 0.13]],

                   [[0.09, 0.23, 0.35],
                    [0.3 , 0.42, 0.06]],

                   [[0.24, 0.1 , 0.92],
                    [0.82, 0.52, 0.7 ]]])

需要提取得到的目标数组:

# array([[0.79, 0.1 , 0.31],
#        [0.3 , 0.52, 0.7 ]])

解决方法

方法1:使用np.take_along_axis(推荐)

np.take_along_axis专门用于根据指定轴上的索引提取元素,只需先给argmax1增加一个维度,使其与ndarr2在axis=0上的形状兼容:

# 给argmax1增加维度,形状从(2,3)变为(1,2,3)
argmax1_expanded = argmax1[np.newaxis, ...]
# 沿axis=0提取对应索引的元素
result = np.take_along_axis(ndarr2, argmax1_expanded, axis=0)
# 去除多余维度,得到(2,3)的结果
result = result.squeeze(axis=0)

print(result)

方法2:索引广播

通过生成其他维度的索引,结合numpy的广播机制提取元素:

# 生成第二、第三维度的索引数组
i, j = np.indices(argmax1.shape)
# 利用索引广播提取对应位置的元素
result = ndarr2[argmax1, i, j]

print(result)

两种方法都能得到目标结果,其中np.take_along_axis更简洁直观,适合多维数组的索引提取场景。

内容的提问来源于stack exchange,提问作者Alex Essilfie

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.28 07:05:11