无需for-loop提取nxn矩阵及nxkxk矩阵数组对角线元素(保原形状)
NumPy 对角线提取问题解决方案
1. 无for循环提取n×n矩阵的对角线元素
直接使用NumPy内置的np.diag()或np.diagonal()函数,无需任何循环:
np.diag(matrix):专为二维矩阵设计,直接返回主对角线的一维数组np.diagonal(matrix):功能与np.diag()一致,同时支持高维数组的轴参数设置
代码示例:
import numpy as np # 生成3×3测试矩阵 test_mat = np.array([[1, 2, 3], [4, 5, 6], [7, 8, 9]]) # 两种方法提取对角线 diag_result1 = np.diag(test_mat) diag_result2 = np.diagonal(test_mat) print(diag_result1) # 输出: [1 5 9]
2. 无for循环提取n×k×k数组的对角线并保持三维形状
如果需要从形状为(n, k, k)的数组(包含n个k×k矩阵)中提取对角线,同时保持三维结构(而非压缩为二维),可以用以下两种简洁方法:
方法1:索引切片+维度扩展
利用np.arange(k)生成对角线的位置索引,直接定位每个子矩阵的对角线元素,再通过np.newaxis快速扩展维度:
import numpy as np # 生成2个3×3矩阵组成的测试数组(形状(2, 3, 3)) test_arr = np.array([ [[1, 2, 3], [4, 5, 6], [7, 8, 9]], [[10, 11, 12], [13, 14, 15], [16, 17, 18]] ]) n, k, _ = test_arr.shape # 提取对角线并保持(2, 3, 1)的三维形状 diag_3d = test_arr[:, np.arange(k), np.arange(k)][:, :, np.newaxis] print(diag_3d.shape) # 输出: (2, 3, 1) # 若需要(2, 1, 3)的形状,调整扩展维度的位置即可 diag_3d_alt = test_arr[:, np.arange(k), np.arange(k)][:, np.newaxis, :] print(diag_3d_alt.shape) # 输出: (2, 1, 3)
方法2:np.diagonal+np.expand_dims
如果习惯使用np.diagonal,可以直接对其返回的二维结果扩展维度,一步完成:
# 先提取对角线得到(2, 3)的二维数组,再扩展为(2, 3, 1)的三维数组 diag_3d = np.expand_dims(np.diagonal(test_arr, axis1=1, axis2=2), axis=-1) print(diag_3d.shape) # 输出: (2, 3, 1)
两种方法均无需for循环,且比“提取后手动重构形状”的写法更简洁。
内容的提问来源于stack exchange,提问作者Rutger Versteegden
相关产品推荐
相关产品推荐

