You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何在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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.05.25 08:12:27