Numpy高级索引问题:如何用索引数组提取目标形状数组?
解决Numpy高级索引中获取(K,A,B)形状结果的问题
你的问题出在索引方式的理解上:当你直接传入形状为(K,N)的数组作为索引时,Numpy并不会将每个行视为一个N维索引元组,而是会把这个数组当作对原数组第一个维度的索引,进而触发广播机制,导致输出形状不符合预期。
正确的实现方式
要实现预期的(K,A,B)形状结果,你需要将索引数组转换为N个长度为K的一维数组组成的元组,分别对应原数组的前N个维度。具体来说:
- 将你的
indices数组转置(从(5,3)变为(3,5)) - 将转置后的数组拆分为元组,传入索引
修改后的代码如下:
import numpy as np arr = np.random.uniform(size=(4, 4, 4, 3, 3)) indices = [(0, 1, 2), (2, 1, 1), (0, 0, 1), (1, 1, 2), (2, 2, 2)] indices = np.array(indices) # 转换为形状(5,3)的数组 selected_arr = arr[tuple(indices.T)] print(selected_arr.shape) # 输出 (5, 3, 3)
原理说明
当你传入由N个一维数组组成的元组时,Numpy会执行整数数组索引:
- 每个一维数组对应原数组的一个维度(比如
indices.T[0]对应原数组的第0维,indices.T[1]对应第1维,以此类推) - 每个位置上的索引值会被组合起来,提取原数组中对应位置的元素
- 最终结果的形状为
(K,)加上原数组剩余的维度(A,B),正好符合你的预期
另一种等价写法
如果你觉得转置不够直观,也可以通过列表推导式拆分索引数组:
idxs = [indices[:, i] for i in range(indices.shape[1])] selected_arr = arr[tuple(idxs)]
内容的提问来源于stack exchange,提问作者Ufuk Can Bicici
相关产品推荐
相关产品推荐

