基于一维数组生成含连续对角值的多维数组的高效实现
高效生成带连续对角值的多维数组
投影操作后需基于一维数组生成含连续对角值的多维数组,要求方案高效(实际数据量远大于示例的12个值),原“紧凑”轴不再使用但数据值保留。
输入
>>> import numpy as np >>> a = np.arange(1,13).astype(float) array([ 1., 2., 3., 4., 5., 6., 7., 8., 9., 10., 11., 12.]) >>> a.shape (12,)
期望输出
array([[[ 1., 0., 0., 0.], [ 0., 2., 0., 0.], [ 0., 0., 3., 0.], [ 0., 0., 0., 4.]], [[ 5., 0., 0., 0.], [ 0., 6., 0., 0.], [ 0., 0., 7., 0.], [ 0., 0., 0., 8.]], [[ 9., 0., 0., 0.], [ 0., 10., 0., 0.], [ 0., 0., 11., 0.], [ 0., 0., 0., 12.]]]) shape: (3, 4, 4)
尝试过的无效方法
>>> np.tile(np.identity(4),(3,1)).reshape(3,4,4) array([[[1., 0., 0., 0.], [0., 1., 0., 0.], [0., 0., 1., 0.], [0., 0., 0., 1.]], [[1., 0., 0., 0.], [0., 1., 0., 0.], [0., 0., 1., 0.], [0., 0., 0., 1.]], [[1., 0., 0., 0.], [0., 1., 0., 0.], [0., 0., 1., 0.], [0., 0., 0., 1.]]])
高效解决方案
方法1:直接索引赋值(内存效率最高)
通过索引定位对角位置,直接赋值,无中间数组开销,适合超大数组:
import numpy as np a = np.arange(1,13).astype(float) # 初始化结果数组 result = np.zeros((3, 4, 4), dtype=a.dtype) # 获取对角索引 idx = np.arange(4) # 赋值到对角位置 result[:, idx, idx] = a.reshape(3, 4)
方法2:利用np.diag与转置
简洁直观,适合中等规模数据:
a_reshaped = a.reshape(3, 4) result = np.diag(a_reshaped).transpose(0, 2, 1)
方法3:einsum实现矩阵乘法
适合需要显式矩阵运算逻辑的场景:
a_reshaped = a.reshape(3, 4) result = np.einsum('ij,jk->ijk', a_reshaped, np.eye(4))
三种方法均能得到符合要求的输出,其中方法1的内存效率最优,适合处理大规模投影数据。
内容的提问来源于stack exchange,提问作者Tobias
相关产品推荐
相关产品推荐

