Python NumPy索引用法困惑:多维数组索引结果维度解析
理解NumPy中用多维数组索引一维数组的逻辑
嘿,作为NumPy新手碰到这种索引逻辑确实容易懵,我来一步步给你拆解清楚~
先从你的简化示例说起
首先看你贴的这段测试代码:
import numpy as np a = np.ones([1,1,5,5], dtype='int64') b = np.ones([11], dtype='float64') x = b[a] print(x.shape) # (1, 1, 5, 5)
这其实是NumPy里整数数组索引的典型用法,核心逻辑很简单:
b是一个一维数组,长度为11;a是一个4维的整数数组,所有元素都是1(因为np.ones生成的全1数组)。- 当你用多维整数数组去索引一维数组时,NumPy会把索引数组
a里的每一个元素,都当作b的下标去取值,最后返回的数组形状和索引数组a的形状完全一致。
举个更直观的小例子帮你理解:
import numpy as np b = np.array([10, 20, 30]) # 一维数组 a = np.array([[0, 1], [2, 0]]) # 2维索引数组 print(b[a]) # 输出 [[10 20], [30 10]],形状和a一样是(2,2)
这里a里的每个数字都是b的下标,b[a]就是把这些下标对应的元素按a的结构排列起来。
回到你的测试代码,a里全是1,所以b[a]就是把b[1]这个值重复填充成(1,1,5,5)的形状,这就是为什么x.shape是这个结果。
再分析你的实际业务代码
现在看你补充的这段实际代码:
def gausslabel(length=180, stride=2): gaussian_pdf = signal.gaussian(length+1, 3) label = np.reshape(np.arange(stride/2, length, stride), [1,1,-1,1]) y = np.reshape(np.arange(stride/2, length, stride), [1,1,1,-1]) delta = np.array(np.abs(label - y), dtype=int) delta = np.minimum(delta, length-delta)+length/2 return gaussian_pdf[delta]
这里的核心逻辑和上面的测试代码完全一致:
gaussian_pdf是一个一维数组,长度为length+1(默认是181),存的是高斯分布的概率值。delta是一个4维的整数数组(形状是(1,1, N, N),其中N是np.arange(stride/2, length, stride)的元素个数,默认stride=2时是90),它的每个元素都是合法的下标值(通过np.minimum和后续的加法操作,确保不会超出gaussian_pdf的索引范围)。- 最后
gaussian_pdf[delta]就是把delta里的每个整数作为下标,去gaussian_pdf中取出对应位置的高斯值,返回的数组形状和delta完全一致。
找资料的关键词
如果你想深入学习这种索引逻辑,可以去NumPy官方文档里搜索**「Integer array indexing」**,这部分内容会详细讲解各种整数数组索引的场景和规则,包括多维索引数组的情况。
内容的提问来源于stack exchange,提问作者Sumsuddin Shojib
相关产品推荐
相关产品推荐

