
本文介绍如何在不使用 for 循环的前提下,从形状为 (k, n, n) 的批量矩阵中提取每块矩阵的主对角线元素,并严格保持原始数组维度结构——即输出仍为 (k, n, n),仅非对角位置置零。
本文介绍如何在不使用 for 循环的前提下,从形状为 `(k, n, n)` 的批量矩阵中提取每块矩阵的主对角线元素,并**严格保持原始数组维度结构**——即输出仍为 `(k, n, n)`,仅非对角位置置零。
在科学计算与深度学习中,常需处理一批大小相同的方阵(如 k 个 n×n 矩阵),并希望高效提取每个矩阵的对角线元素,同时保留整体张量结构。虽然 np.diagonal(..., axis1=1, axis2=2) 可提取对角线,但它会压缩维度(返回 (k, n)),破坏原始 (k, n, n) 形状;而手动 reshape 重建既冗余又易出错。
最简洁、向量化且语义清晰的解法是利用广播乘法(broadcasting)结合单位矩阵:
import numpy as np # 示例:3 个 4×4 矩阵 n, k = 4, 3 arr = np.arange(k * n * n).reshape(k, n, n) # shape: (k, n, n) # ✅ 核心技巧:用 eye(n) 广播相乘,自动对齐最后两维 diags = arr * np.eye(n) # shape 保持为 (k, n, n)
np.eye(n) 生成 n×n 单位矩阵(对角线为 1,其余为 0)。当与 (k, n, n) 数组相乘时,NumPy 沿后两维自动广播:arr[i,:,:] 与 eye(n) 逐元素相乘,结果中仅保留第 i 个矩阵的对角线值,其余全为 0。
⚠️ 注意事项:
- 此方法适用于 主对角线(diag(arr[i])) 提取;若需副对角线或指定偏移,需改用 np.fliplr(np.eye(n)) 或索引技巧。
- 输入必须是 方阵批次(即最后两维相等);若为 (k, n, m) 非方阵,np.eye(n) 不兼容,应改用高级索引(如 np.arange(n) 构造索引元组)。
- 内存友好:全程纯向量化,无 Python 循环,性能接近底层 C 实现。
总结:相比 np.diagonal + reshape 组合,arr * np.eye(n) 更直观、更安全、更符合 NumPy 的广播哲学——它不是“提取对角线”,而是“掩码保留对角线”,天然维持张量结构,是批量矩阵对角操作的推荐范式。











