如何用索引数组对Numpy数组第二维度切片得到指定结果?
解决NumPy三维数组按指定索引切片的问题
首先,我们先明确原数组的结构:
import numpy as np a = np.arange(24).reshape(4,3,2) index_dim2 = np.array([0,1,2,2])
a的形状是(4,3,2),可以拆分为4个形状为(3,2)的子数组:
a[0] = [[0,1], [2,3], [4,5]]a[1] = [[6,7], [8,9], [10,11]]a[2] = [[12,13], [14,15], [16,17]]a[3] = [[18,19], [20,21], [22,23]]
我们的目标是对每个第一维度的子数组(共4个),提取第二维度中对应index_dim2索引的元素,并且保持结果为(4,1,2)的三维结构(和目标输出一致)。
正确的切片表达式
你需要的切片表达式是:
a[np.arange(4), index_dim2[:, np.newaxis], :]
或者更简洁的写法(None等价于np.newaxis):
a[range(4), index_dim2[:, None], :]
为什么这样写?
np.arange(4):遍历第一维度的所有索引(0到3),对应4个子数组。index_dim2[:, np.newaxis]:把原本形状为(4,)的索引数组转换为(4,1),这样在索引时会和第一维度的索引广播,让每个第一维度的元素只提取第二维度中指定的单个索引,同时保留一个长度为1的第二维度,确保结果是三维的。::提取第三维度的所有元素(每个子元素的两个数值)。
验证结果
运行这段代码:
import numpy as np a = np.arange(24).reshape(4,3,2) index_dim2 = np.array([0,1,2,2]) b = a[np.arange(4), index_dim2[:, np.newaxis], :] print(b)
输出结果完全符合你的目标:
array([[[ 0, 1]], [[ 8, 9]], [[16, 17]], [[22, 23]]])
另外,如果你先得到(4,2)的二维数组再扩展维度,也可以达到同样效果,比如:
b = a[np.arange(4), index_dim2, :][:, np.newaxis]
但第一种方式更直接,一步到位完成索引和维度保留。
内容的提问来源于stack exchange,提问作者cass
相关产品推荐
相关产品推荐

