如何将[N*1*M]数组值快速赋值给[N*M*M]数组各层的对角线位置
Numpy批量赋值多维数组每层对角线的优化方案
完全可以通过向量化操作替代循环,性能提升非常显著,且输出结果和原有实现完全一致。
优化后代码
import numpy as np # 保留原有输入格式 a = np.arange(6).reshape(2,3,1) I = np.zeros((a.shape[0], 3, 3)) # 核心向量化赋值,无任何Python层循环 diag_idx = np.arange(3) I[:, diag_idx, diag_idx] = a.squeeze(-1) print(I)
代码说明
- 利用numpy多维高级索引特性,
:选中所有N个数据层,连续两次传入[0,1,2]索引,即可定位到每个3×3矩阵的三个对角线位置(0,0)、(1,1)、(2,2) - 所有运算都在numpy底层C层面执行,当N的规模达到万级以上时,执行效率比循环实现高数百倍
- 如果需要适配动态的对角线长度,可以改用更通用的
np.diag_indices生成索引:dim = a.shape[1] diag_idx = np.diag_indices(dim) I[:, diag_idx[0], diag_idx[1]] = a.squeeze(-1)
内容的提问来源于stack exchange,提问作者oddroj
相关产品推荐
相关产品推荐

