如何根据指定第二维度索引使用np.take从3D NumPy矩阵提取目标行
NumPy 按轴匹配索引选取3D数组元素的实现
给定条件
- 3D测试数组定义:
import numpy as np a = np.arange(3*4*5).reshape(3,4,5)
数组结构如下:
array([[[ 0, 1, 2, 3, 4], [ 5, 6, 7, 8, 9], [10, 11, 12, 13, 14], [15, 16, 17, 18, 19]], [[20, 21, 22, 23, 24], [25, 26, 27, 28, 29], [30, 31, 32, 33, 34], [35, 36, 37, 38, 39]], [[40, 41, 42, 43, 44], [45, 46, 47, 48, 49], [50, 51, 52, 53, 54], [55, 56, 57, 58, 59]]])
- 索引规则:索引数组
idx = [1,0,2,2],对轴1(第二维度)的第i个位置,从轴0(第一维度)选取idx[i]对应的整行(保留轴2全部元素),最终输出形状为(4,5)的矩阵,目标结果如下:
array([[20, 21, 22, 23, 24], [ 5, 6, 7, 8, 9], [50, 51, 52, 53, 54], [55, 56, 57, 58, 59]])
正确实现
直接使用NumPy原生高级整数索引即可,写法简洁且执行效率最高:
idx = np.array([1, 0, 2, 2]) result = a[idx, np.arange(a.shape[1]), :]
逻辑说明
NumPy多维数组高级索引支持按位置匹配不同轴的下标:
- 第一个参数
idx对应轴0的选取下标,长度和轴1长度一致 - 第二个参数
np.arange(a.shape[1])生成序列[0,1,2,3],对应轴1的每个位置顺序 - 第三个参数
:表示选取轴2的全部元素 - 索引会自动按位置配对:依次取
a[1,0,:]、a[0,1,:]、a[2,2,:]、a[2,3,:],拼接后正好得到目标结果
结果验证
执行代码后打印result,输出和目标完全一致:
array([[20, 21, 22, 23, 24], [ 5, 6, 7, 8, 9], [50, 51, 52, 53, 54], [55, 56, 57, 58, 59]])
注:该方法无额外内存拷贝开销,适配任意满足
len(idx)==a.shape[1]且idx取值在轴0长度范围内的索引场景。
内容的提问来源于stack exchange,提问作者InsDSt
相关产品推荐
相关产品推荐

