NumPy ndarray指定axis参数时min/max方法运算规则咨询
NumPy ndarray min/max方法axis参数运算规则
核心规则
指定axis=n时,运算逻辑是沿着第n个轴的方向做聚合计算,计算完成后移除第n个维度,而非“对第n个轴下的每个子元素单独求全局最值”。
针对测试用例的拆解
测试数组x形状为(2,3,4),维度索引从0开始计数,各轴对应含义:
- 轴0(axis=0):长度为2,对应最外层2个独立的3×4二维矩阵块
- 轴1(axis=1):长度为3,对应每个二维矩阵内的行方向
- 轴2(axis=2):长度为4,对应每个二维矩阵内的列方向
调用x.min(axis=0)的实际计算逻辑:
- 沿着轴0的方向,对两个矩阵块相同坐标位置的元素两两取最小值:即对任意坐标
(i,j)(i范围0-2,j范围0-3),计算min(x[0,i,j], x[1,i,j]) - 计算完成后轴0被移除,输出数组的形状为原形状去掉轴0的长度,即
(3,4),和实际运行得到的输出形状完全吻合。
可以用输出值做验证:
- 输出数组第一行第一列的
0.4139181,是x[0,0,0]=0.4139181和x[1,0,0]=0.79760899的最小值 - 输出数组第二行第一列的
0.54778544,是x[0,1,0]=0.63775691和x[1,1,0]=0.54778544的最小值
和打印的输出结果完全匹配。
预期效果的实现方式
如果需要得到长度为2的数组,分别对应x[0]、x[1]两个矩阵块的全局最小值,本质是对轴0上的每个子元素,沿着剩下的所有轴做min聚合,对应参数需要传入轴元组:
x.min(axis=(1,2))
该调用会保留轴0,消除轴1和轴2,最终返回形状为(2,)的数组,就是预期的结果。
常见调用效果参考
针对形状为(2,3,4)的数组x,不同axis参数的输出效果:
x.min(axis=1):沿行方向逐列取最小值,消除轴1,返回形状为(2,4)的数组x.min(axis=2):沿列方向逐行取最小值,消除轴2,返回形状为(2,3)的数组x.min():不指定axis时,对全数组所有元素取最小值,返回单个标量
内容的提问来源于stack exchange,提问作者westcoaststudent
相关产品推荐
相关产品推荐

