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

如何访问形状为N*M*M的numpy数组中各M*M矩阵的下三角元素

提取三维numpy数组中每个二维矩阵的下三角元素解决方案

错误原因分析

  • 第一种写法arr[:,np.tril_indices(M,-1)]:np.tril_indices返回两个长度为M*(M-1)/2的行、列索引数组,直接放在第二维索引会生成形状为(N, 2, M*(M-1)/2)的中间数组,当M数值较大时会占用过量内存导致内核崩溃,且输出形状不符合需求。
  • 第二种写法arr[:][np.tril_indices(M,-1)]:arr[:]仍然是完整的三维数组,此时索引会作用于长度为N的第0轴,而tril_indices生成的索引值超过N的取值范围,因此抛出轴0越界的IndexError。
  • np.tril(arr)仅会将每个二维矩阵的上三角元素置为0,不会改变原数组形状,因此无法直接得到目标格式的输出。

正确实现方案

直接使用np.tril_indices生成的行、列索引,对三维数组的最后两维进行索引即可一步得到目标结果,该实现为矢量化操作,无Python层循环,计算效率高:

import numpy as np

# 获取下三角(偏移量-1,不含对角线)对应的行、列索引
tri_row_idx, tri_col_idx = np.tril_indices(M, k=-1)
# 对三维数组的后两维按索引提取元素,输出形状自动为 (N, M*(M-1)//2)
result = arr[:, tri_row_idx, tri_col_idx]

验证示例

# 测试参数与输入
N = 6
M = 4
test_arr = np.random.rand(N, M, M)

# 执行提取
tri_r, tri_c = np.tril_indices(M, k=-1)
res = test_arr[:, tri_r, tri_c]

print(res.shape) 
# 输出为 (6, 6),符合 M*(M-1)/2 = 4*3/2 = 6 的预期

内容的提问来源于stack exchange,提问作者Aayush Desai

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.26 15:42:04