
本文介绍如何在不使用 for 循环的前提下,从形状为 (k, n, n) 的批量方阵中提取每块矩阵的主对角线元素,并保持原始批量结构——核心技巧是利用广播机制与单位矩阵逐元素相乘。
本文介绍如何在不使用 for 循环的前提下,从形状为 `(k, n, n)` 的批量方阵中提取每块矩阵的主对角线元素,并保持原始批量结构——核心技巧是利用广播机制与单位矩阵逐元素相乘。
在科学计算和深度学习中,常需处理批量(batch)方阵,例如形状为 (k, n, n) 的张量,其中 k 是样本数或批次大小,每个 (n, n) 子矩阵代表一个独立的方阵。若目标是提取每个子矩阵的主对角线元素,并保留其在原张量中的位置(即非压缩、非降维),则 np.diagonal 或 np.diag 并不直接满足需求:前者会改变维度顺序并压缩对角轴,后者仅支持 2D 输入且不支持 axis 参数。
一种简洁、纯向量化(zero-loop)的解决方案是:利用单位矩阵 np.eye(n) 与输入张量进行逐元素乘法(broadcasting)。由于 np.eye(n) 是 n×n 的布尔型单位矩阵(对角为 1,其余为 0),当它与形状为 (k, n, n) 的张量做 * 运算时,NumPy 自动沿前导维度广播,结果中仅对角位置被保留,其余位置置零,从而完美维持原始形状 (k, n, n)。
以下为完整示例代码:
import numpy as np
# 构造示例:3 个 4×4 矩阵
n = 4
k = 3
arr = np.arange(k * n * n).reshape(k, n, n)
# 向量化提取对角线(保持形状)
diags = arr * np.eye(n) # 自动广播:(k,n,n) * (n,n) → (k,n,n)
print("原始张量 shape:", arr.shape) # (3, 4, 4)
print("对角掩码后 shape:", diags.shape) # (3, 4, 4)
print("第一个矩阵的对角线(非零值):", diags[0].diagonal()) # [0. 5. 10. 15.]✅ 优势总结:
- 完全避免显式循环,充分利用 NumPy 广播机制;
- 输出形状严格等于输入形状,便于后续张量操作(如求和、归一化、拼接);
- 计算高效,底层由 C 实现,适用于大规模批量矩阵。
⚠️ 注意事项:
- 此方法返回的是带零填充的对角矩阵(即 (k, n, n)),若仅需对角线值组成的 (k, n) 向量,请改用 np.diagonal(arr, axis1=1, axis2=2);
- np.eye(n) 默认为 float64,若输入为整型且需保持 dtype,可显式指定:np.eye(n, dtype=arr.dtype);
- 该技巧仅适用于提取主对角线;若需副对角线或其他偏移对角线,需配合 np.fliplr 或索引切片,但不再具备同等简洁性。
综上,arr * np.eye(n) 是提取批量方阵主对角线并保形的最简向量化范式,兼顾可读性、性能与通用性,推荐作为 NumPy 批量矩阵处理的标准实践之一。


















