
本文介绍一种基于广播机制的向量化方法,使用 np.eye() 与批量矩阵相乘,直接提取每个 n×n 子矩阵的主对角线元素,同时严格保持输入张量的原始形状(如 k×n×n),避免显式循环和冗余重构。
本文介绍一种基于广播机制的向量化方法,使用 np.eye() 与批量矩阵相乘,直接提取每个 n×n 子矩阵的主对角线元素,同时严格保持输入张量的原始形状(如 k×n×n),避免显式循环和冗余重构。
在科学计算与深度学习中,常需处理形如 (k, n, n) 的批量方阵(即 k 个 n×n 矩阵),并希望高效提取每个矩阵的主对角线元素,同时维持原始张量结构——即输出应为相同 batch 维度的对角矩阵形式:每个对角矩阵仅在主对角线上保留原值,其余位置置零(形状仍为 k×n×n),而非压缩为 (k, n) 形状。
NumPy 的 np.diagonal() 虽支持指定轴提取对角线(如 np.diagonal(arr, axis1=1, axis2=2)),但其返回结果会降维(变为 (k, n)),破坏原始三维结构;而 np.diag() 仅适用于一维或二维输入,不支持轴参数,无法直接用于批量操作。
✅ 推荐解法:利用广播乘法 + 单位矩阵掩码
核心思想是构造一个 n×n 的单位矩阵 I = np.eye(n),再通过 NumPy 广播机制将其与 (k, n, n) 张量逐矩阵相乘:
- arr * np.eye(n) 中,np.eye(n) 自动广播至 (1, n, n),再扩展为 (k, n, n);
- 乘法按元素进行,仅当 i == j 时保留 arr[:, i, j],其余位置乘以 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.]
⚠️ 注意事项:
- 此方法输出为浮点型(因 np.eye(n) 默认 float64),若需保持原始 dtype,可显式指定:np.eye(n, dtype=arr.dtype);
- 仅适用于提取主对角线(i == j);若需副对角线或其他偏移对角线,需改用 np.eye(n, k=offset) 或索引技巧;
- 内存友好:全程向量化,无 Python 循环,适合 GPU 加速框架(如 CuPy)无缝迁移。
总结:相比先用 np.diagonal() 提取再 np.expand_dims()/np.repeat() 重建形状,该广播掩码法更简洁、可读性更高,且天然兼容任意 batch size,是处理批量矩阵对角操作的推荐实践。











