NumPy三维数组axis释义:np.stack函数axis=2工作原理咨询
NumPy三维数组axis=2轴工作机制详解
前置基础
首先定义两个参与堆叠的二维数组,两者形状均为(3,1)(3行1列):
import numpy as np a = np.array([[9],[9],[8]]) b = np.array([[1],[4],[6]])
使用np.stack沿不同轴堆叠的输出形状如下:
- axis=0:输出形状
(2,3,1) - axis=1:输出形状
(3,2,1) - axis=2:输出形状
(3,1,2)
核心逻辑先明确
np.stack的本质操作逻辑非常固定:
给所有参与堆叠的输入数组,统一在
axis参数指定的位置插入一个全新的维度,新维度的长度等于输入数组的总个数,最后沿这个新维度把所有数组合并为一个数组。
原数组a、b都是二维结构,维度顺序为[轴0(长度3), 轴1(长度1)],插入新维度的位置不同,最终的三维结构就不同:
- axis=0:新维度插在最前面,维度顺序变为
[新轴(长度2), 原轴0(3), 原轴1(1)],对应形状(2,3,1),也就是把a、b两个数组整体作为最外层的两个块。 - axis=1:新维度插在原轴0和原轴1中间,维度顺序变为
[原轴0(3), 新轴(2), 原轴1(1)],对应形状(3,2,1),也就是按行对齐,把两个数组同一行的内容沿新轴拼接。 - axis=2:新维度插在所有原有维度的末尾,维度顺序变为
[原轴0(3), 原轴1(1), 新轴(2)],对应形状(3,1,2)。
axis=2的直观结构解释
我们可以用熟悉的二维表格逻辑延伸理解三维轴:
- axis=0:代表有多少张独立的二维表
- axis=1:代表每张二维表有多少行
- axis=2:代表每张二维表每一行最内层的元素排列方向(可以理解为二维表的列方向)
沿axis=2堆叠时,会完全保留原数组的行数和单行列数,只在最内层的元素位置把两个数组对应坐标的数值拼在一起:
- 原数组a第0行第0列的值是9,原数组b第0行第0列的值是1,两个值就拼成输出第0行最内层的
[9,1] - 原数组a第1行第0列的值是9,原数组b第1行第0列的值是4,两个值就拼成输出第1行最内层的
[9,4] - 原数组a第2行第0列的值是8,原数组b第2行第0列的值是6,两个值就拼成输出第2行最内层的
[8,6]
最终输出结构和运行结果完全一致:
print(np.stack([a,b],axis=2)) # 输出 array([[[9, 1]], [[9, 4]], [[8, 6]]])
和axis=1的结果做对比就能明显看出差异:axis=1是把两个数组同位置的行作为新轴上的两个独立子元素(每个子元素还是长度1的数组[9]、[1]),而axis=2是直接把两个同位置的标量值,塞进最内层的同一个数组里。
内容的提问来源于stack exchange,提问作者Muhammad Waseem
相关产品推荐
相关产品推荐

