如何从2D数组高效生成每行填充对角线的NumPy 3D数组?
高效将2D NumPy数组转换为每行对应对角矩阵的3D数组
针对你提出的需求——将2D数组的每一行映射为3D数组对应维度的对角矩阵,且第一维度规模较大时保证效率,以下是两种高效的向量化实现方案:
方法1:利用np.einsum构造
np.einsum通过爱因斯坦求和约定实现高效的张量操作,无需显式循环即可完成对角矩阵的批量生成:
import numpy as np # 输入示例数组 arr = np.array([[1, 2], [3, 4], [5, 6], [7, 8]]) # 生成目标3D数组 result = np.einsum('ij,jk->ijk', arr, np.eye(arr.shape[1]))
原理
np.eye(arr.shape[1])生成与输入数组列数一致的单位矩阵(M×M)- 通过
einsum的ij,jk->ijk规则,将输入数组的每一行(ij维度)与单位矩阵(jk维度)相乘,等价于把每行元素批量填充到对应位置的对角线上,最终得到形状为(N, M, M)的3D数组。
方法2:广播索引直接赋值
通过预分配零数组,利用NumPy的广播索引直接给对角线位置赋值,同样是无循环的高效操作:
import numpy as np arr = np.array([[1, 2], [3, 4], [5, 6], [7, 8]]) N, M = arr.shape # 预分配全零3D数组 result = np.zeros((N, M, M), dtype=arr.dtype) # 批量赋值对角线元素 result[:, np.arange(M), np.arange(M)] = arr
原理
[:, np.arange(M), np.arange(M)]是三维数组的索引,:匹配第一维度的所有元素,np.arange(M)分别匹配第二、第三维度的对角位置- 直接将输入数组的元素广播赋值到这些对角位置,操作全程由NumPy底层优化,性能优异。
性能说明
两种方案均为向量化操作,避免了Python层面的循环,当第一维度规模极大(如10^5级)时,相比循环实现能提升几个数量级的效率,完全适配大规模数据场景。
内容的提问来源于stack exchange,提问作者user8036269
相关产品推荐
相关产品推荐

