关于numpy.apply_along_axis与np.diag组合使用的结果不符及替代方法咨询
解决numpy apply_along_axis(axis=0)生成三维对角数组的问题
首先看你的数组定义:
import numpy as np b = np.array([[1,2,3], [4,5,6], [7,8,9]])
为什么np.apply_along_axis(np.diag, 0, b)得不到预期结果?
因为np.apply_along_axis设计时默认期望每个函数调用返回1D数组或标量,当传入axis=0时,它会逐列取一维数组(比如[1,4,7])传给np.diag得到3x3矩阵,但它会尝试将这些矩阵扁平化后堆叠成二维数组,而非保留三维结构,所以输出不符合预期。
两种可行的解决方法:
方法一:列表推导式逐列生成对角矩阵
直接遍历转置后的数组(原数组的列变为行),对每个行生成对角矩阵,最后合并为三维数组:
result = np.array([np.diag(col) for col in b.T])
方法二:利用numpy广播机制构造
通过广播直接生成目标三维数组,效率更高:
# 生成单位矩阵并扩展维度,与转置后的数组广播相乘 result = np.eye(3)[None, :, :] * b.T[:, None, :]
验证结果:
打印result会得到你预期的输出:
[[[1 0 0] [0 4 0] [0 0 7]] [[2 0 0] [0 5 0] [0 0 8]] [[3 0 0] [0 6 0] [0 0 9]]]
内容的提问来源于stack exchange,提问作者Daniel Esnady
相关产品推荐
相关产品推荐

