You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

无需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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.25 07:20:20