如何访问形状为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
相关产品推荐
相关产品推荐

