np.apply_along_axis输出形状解析:为何维度存在差异?
关于np.apply_along_axis输出维度差异的解释
核心规则:输出维度由原数组维度位置与函数输出维度拼接而成
np.apply_along_axis 的最终输出形状遵循明确的拼接逻辑:
输出形状 = (原数组中指定axis之前的维度) + (函数的输出形状) + (原数组中指定axis之后的维度)
我们用你的两个案例验证这个规则:
第一个案例拆解
- 原数组:
np.asarray([[1,1,1],[2,3,4]]),形状为(2, 3) - 指定axis=0:该axis是原数组的第一个维度,axis之前无维度,axis之后的维度是
(3) - 函数输出形状:
(2, 2) - 拼接后输出形状:
() + (2,2) + (3) = (2,2,3),与你得到的结果一致
第二个案例拆解
- 原数组:
np.asarray([[1,2,3]]).T,形状为(3, 1) - 指定axis=1:该axis是原数组的最后一个维度,axis之前的维度是
(3),axis之后无维度 - 函数输出形状:
(2, 2) - 拼接后输出形状:
(3) + (2,2) + () = (3,2,2),也和你得到的结果一致
简单来说,apply_along_axis 会把函数的输出结果“嵌入”到原数组中被遍历的axis位置,原数组中axis前后的维度会分别保留在输出形状的前后两端,这就是两个案例输出形状不同的根本原因。
内容的提问来源于stack exchange,提问作者Konstantin
相关产品推荐
相关产品推荐

