如何在NumPy或PyTorch中提取矩阵的对角条带?
提取K×N(K>N)矩阵的对角条带
嘿,刚好碰到过类似的需求,给你两种实用的实现方式,包括你提到的布尔掩码方法,还有更省心的numpy内置技巧~
方案一:手动构造布尔掩码
你说的布尔掩码思路完全可行,先把完整的正确代码放出来,再给你拆解逻辑:
import numpy as np # 生成示例矩阵 X = np.arange(12).reshape(4, 3) print("原矩阵:") print(X) # 构造匹配需求的布尔掩码 mask = np.asarray([ [True, False, False], [True, True, False], [False, True, True], [False, False, True] ]) # 提取掩码选中的元素,先重塑为N×(K-N+1),再转置得到目标形状 result = X[mask].reshape(X.shape[1], X.shape[0] - X.shape[1] + 1).T print("\n提取后的对角条带矩阵:") print(result)
掩码逻辑拆解
这个掩码的设计是让每条斜向的“条带”按行优先被提取:
- 第一列的前两个元素(0,3)、第二列的中间两个元素(4,7)、第三列的后两个元素(8,11)分别被标记为
True - 提取后得到一维数组
[0,3,4,7,8,11],先把它reshape成N×(K-N+1)的矩阵(这里是3×2),再转置就得到(K-N+1)×N的目标矩阵(2×3),正好对应每条对角线作为一行的结果。
方案二:用numpy内置函数简化操作
如果不想手动写掩码,用np.diag配合np.vstack可以一步到位,代码更简洁还通用:
import numpy as np X = np.arange(12).reshape(4, 3) n_rows, n_cols = X.shape # 提取所有长度为N的对角线,再垂直堆叠成矩阵 result = np.vstack([np.diag(X, k=-i) for i in range(n_rows - n_cols + 1)]) print(result)
代码解释
np.diag(X, k=i):这个函数用来提取矩阵的指定对角线,k=0是主对角线,k=-1是主对角线下方的第一条对角线,以此类推- 对于K×N(K>N)的矩阵,我们需要提取
K-N+1条对角线(这里就是2条),循环生成这些对角线数组后,用np.vstack垂直堆叠起来,直接得到目标矩阵。运行后会直接输出[[0 4 8],[3 7 11]],完全符合需求。
内容的提问来源于stack exchange,提问作者Ivan Bilan
相关产品推荐
相关产品推荐

