如何在单次索引调用中正确实现NumPy多维数组混合索引
NumPy 3维数组单次多维度索引实现
问题复现
给定维度结构为 samples * rows * columns 的3维数组:
import numpy as np arr_3d = np.array([ [ [ 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] ] ])
目标是选取索引为1、2的样本,所有行,索引为0、1的列,分步索引可以得到正确结果:
>>> arr_3d[[1,2],:,:][:,:,[0,1]] array([ [ [10, 11], [13, 14], [16, 17] ], [ [19, 20], [22, 23], [25, 26] ] ])
直接将索引合并为单次调用arr_3d[[1,2],:,[0,1]]会得到不符合预期的结果,这是NumPy高级索引的混合索引特性导致的:当多个维度传入一维数组形式的高级索引时,索引会按位置配对取值,而非独立对每个维度做切片选择,上述错误写法实际是按(1, :, 0)、(2, :, 1)的配对规则取值,最终得到形状为(2,3)的错误结果:
>>> arr_3d[[1,2],:,[0,1]] array([ [10, 13, 16], [20, 23, 26] ])
正确单次索引写法
最简洁、可读性最高的实现方式是使用np.ix_函数,该函数会自动将各维度的索引数组调整为适配广播的维度形状,避免触发索引配对逻辑,实现每个维度独立选取范围:
>>> arr_3d[np.ix_([1,2], np.arange(arr_3d.shape[1]), [0,1])] array([ [ [10, 11], [13, 14], [16, 17] ], [ [19, 20], [22, 23], [25, 26] ] ])
np.ix_的入参按维度顺序传入每个维度要选取的索引列表即可,行维度需要全选时,传入对应长度的np.arange生成的序列即可,返回结果和分步索引完全一致。
如果不想调用额外函数,也可以手动给索引数组增加维度,让不同维度的高级索引形状满足广播规则、不触发位置配对:
>>> arr_3d[np.array([1,2])[:, None, None], :, np.array([0,1])[None, None, :]] array([ [ [10, 11], [13, 14], [16, 17] ], [ [19, 20], [22, 23], [25, 26] ] ])
这种写法通过给样本索引增加两个维度、列索引增加两个维度,让两个高级索引的维度分别落在结果轴的两端,中间保留全选的行维度,也能得到正确结果,但写法繁琐可读性差,日常使用优先选择np.ix_实现。
内容的提问来源于stack exchange,提问作者LayneSadler
相关产品推荐
相关产品推荐

