如何使用np.take索引多维numpy数组 按指定行索引获取目标结果
问题描述
我有一个shape为(n,x,y)的多维数组,本次示例使用的数组如下:
import numpy as np A = np.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]]])另有一个存储索引值的多维数组
Row_values,shape为(z,2),里面的值对应要从原数组A中提取的行索引:Row_values = np.array([[0,1], [0,2], [1,2], [1,3]])我需要将
Row_values中的所有索引值分别应用到A的3个子数组上,最终得到shape为(12,2,3)的结果数组,预期结果如下:Result = np.array([[[0,1,2], [3,4,5]], [[0,1,2], [6,7,8]], [[3,4,5], [6,7,8]], [[3,4,5], [9,10,11]], [[12,13,14], [15,16,17]], [[12,13,14], [18,19,20]], [[15,16,17], [18,19,20]], [[15,16,17], [21,22,23]], [[24,25,26], [27,28,29]], [[24,25,26], [30,31,32]], [[27,28,29], [30,31,32]], [[27,28,29], [33,34,35]]])我尝试用
np.take()实现该需求但没有成功,请问有没有其他更简便的numpy函数可以实现,或者要怎么正确使用np.take()达到上述效果?
实现方法
以下两种方法都可以直接得到符合要求的结果:
方法1:高级索引(最直观)
直接构造对应维度的索引数组即可取值,不需要额外函数:
# 构造第一维度索引:A的每个子数组对应4组Row_values,所以每个索引重复4次 first_idx = np.repeat(np.arange(A.shape[0]), len(Row_values)) # 构造第二维度索引:把Row_values复制3份,对应A的3个子数组 second_idx = np.tile(Row_values, (A.shape[0], 1)) # 直接索引取值,结果shape默认就是(12,2,3) result = A[first_idx, second_idx]
方法2:正确使用np.take实现
你之前调用np.take失败大概率是缺少维度调整步骤,正确用法如下:
# axis=1指定沿着A的第二个维度(行维度)取索引,得到的临时结果shape为(3,4,2,3) temp = np.take(A, Row_values, axis=1) # 把前两个维度合并,就得到目标shape(12,2,3) result = temp.reshape(-1, 2, A.shape[-1])
内容的提问来源于stack exchange,提问作者E. hendrix
相关产品推荐
相关产品推荐

