
本文介绍如何在不使用 for 循环的前提下,从形状为 (k, n, n) 的批量矩阵数组中提取每个矩阵的主对角线元素,并保持原始三维结构——即返回一个同样形状的数组,仅保留对角线位置的值,其余置零。
本文介绍如何在不使用 for 循环的前提下,从形状为 `(k, n, n)` 的批量矩阵数组中提取每个矩阵的主对角线元素,并**保持原始三维结构**——即返回一个同样形状的数组,仅保留对角线位置的值,其余置零。
在科学计算与深度学习中,常需批量处理多个方阵(如一批样本的协方差矩阵、注意力权重矩阵等),并单独操作其对角线(例如提取特征尺度、正则化对角项、构造对角矩阵)。虽然 np.diagonal(..., axis1=1, axis2=2) 可提取对角线,但它会压缩维度(输出形状为 (k, n)),丢失原始 (k, n, n) 结构;而 np.diag 不支持多维轴指定,无法直接用于批量场景。
一种简洁、完全向量化(zero-loop)的解决方案是利用广播乘法 + 单位矩阵掩码:
将输入数组 arr(形状 (k, n, n))与 np.eye(n)(形状 (n, n))逐元素相乘。得益于 NumPy 的广播机制,(k, n, n) * (n, n) 会自动将 eye 沿第 0 轴广播,结果仍为 (k, n, n),且仅对角线位置保留原值,其余位置因与 0 相乘变为 0。
以下是完整示例:
import numpy as np
# 构造示例:3 个 4×4 矩阵
n, k = 4, 3
arr = np.arange(k * n * n).reshape(k, n, n)
# 向量化提取对角线(保持形状)
diags = arr * np.eye(n)
print("原始数组 shape:", arr.shape) # (3, 4, 4)
print("对角掩码后 shape:", diags.shape) # (3, 4, 4)
print("第一个矩阵的对角线:", np.diag(diags[0])) # [ 0. 5. 10. 15.]
✅ 优势说明:
- 无显式循环:纯向量化,性能高,尤其适用于 GPU 加速(如 CuPy)或大批次场景;
- 结构保全:输出与输入同形,便于后续张量运算(如与原矩阵相减、叠加掩码等);
- 内存友好:np.eye(n) 是轻量常量,不随 k 增长。
⚠️ 注意事项:
- 此方法返回的是「对角矩阵形式」(非降维结果),若你实际需要一维对角线序列(如 (k, n)),仍应使用 np.diagonal(arr, axis1=1, axis2=2);
- 若需提取其他对角线(如次对角线),可改用 np.eye(n, k=1) 或索引切片(如 np.einsum('iij->ij', arr)),但 einsum 在此场景下不保持零填充结构;
- 对于超大规模 k,确保 np.eye(n) 能放入内存(通常 n 较小,影响极小)。
综上,arr * np.eye(n) 是兼顾简洁性、可读性与工程效率的推荐方案——它用一行代码替代循环,在保持张量结构的同时精准“点亮”每一张矩阵的对角线。











