如何组合Numpy非连续切片提取指定子矩阵?
提取Numpy协方差矩阵数组的非连续子矩阵
问题背景
有一个形状为(M, N, N)的Numpy数组,其中包含M个(N,N)的协方差矩阵,需要从中提取形状为(M, P, P)的非连续索引子矩阵。目前通过高级索引可以实现需求,但希望找到更直观的切片相关解决方案。
示例代码与输出
import numpy as np # Display all the columns np.set_printoptions(threshold=False, edgeitems=50, linewidth=200) # Create a 6 x 6 matrix. x = np.arange(36).reshape(6,6) # Now make multiple copies to practice with. y = np.array([x, x]) print(f"{y.shape=}\n") print(f"{y=}\n") # We want to extract the submatrices containting the first 2 indices # and the last 2 indices. There are an "unknown" number of intermediate # indices - in this example 2. Thus I'm using negative indices to get the # last two indicies. # Extraction using advanced indexing s = np.array([[0, 1] + [-2, -1]]) subset = y[:, s.T, s] print(f"{subset=}\n") # Now try it with numpy slices. This approach doesn't work first_slice = np.s_[0:2] second_slice = np.s_[-2:] combined_slice = np.r_[first_slice, second_slice] subset = y[:, combined_slice, combined_slice] print(subset)
运行输出:
y.shape=(2, 6, 6) y=array([[[ 0, 1, 2, 3, 4, 5], [ 6, 7, 8, 9, 10, 11], [12, 13, 14, 15, 16, 17], [18, 19, 20, 21, 22, 23], [24, 25, 26, 27, 28, 29], [30, 31, 32, 33, 34, 35]], [[ 0, 1, 2, 3, 4, 5], [ 6, 7, 8, 9, 10, 11], [12, 13, 14, 15, 16, 17], [18, 19, 20, 21, 22, 23], [24, 25, 26, 27, 28, 29], [30, 31, 32, 33, 34, 35]]]) subset=array([[[ 0, 1, 4, 5], [ 6, 7, 10, 11], [24, 25, 28, 29], [30, 31, 34, 35]], [[ 0, 1, 4, 5], [ 6, 7, 10, 11], [24, 25, 28, 29], [30, 31, 34, 35]]]) [[0 7] [0 7]]
核心原因
Numpy的原生切片(slice对象)仅能表示连续、步长固定的索引范围,无法直接描述“前2个+最后2个”这类非连续的索引集合。np.r_虽然可以合并多个切片,但它最终返回的是索引数组,而非切片对象;直接使用该数组进行双维度索引时,会触发高级索引的“配对行为”——将两个一维数组按位置一一对应提取元素,导致结果不符合矩阵子提取的预期。
可行解决方案
1. 用np.ix_生成网格索引
np.ix_可以将一维索引数组转换为二维网格索引,避免配对陷阱,写法更直观:
# 合并切片得到索引数组 idx = np.r_[0:2, -2:] # 使用np.ix_生成矩阵索引,确保提取(M, P, P)子矩阵 subset = y[:, np.ix_(idx, idx)]
np.ix_(idx, idx)会把一维索引数组转为(P,1)和(1,P)的二维数组,触发广播后生成(P,P)的网格索引,从而正确提取每个协方差矩阵的对应子矩阵。
2. 分步索引替代
如果不想用np.ix_,可以通过两次索引实现需求:
idx = np.r_[0:2, -2:] # 先提取目标行,再提取目标列 subset = y[:, idx, :][:, :, idx]
第一步y[:, idx, :]得到(M, P, N)的数组,第二步在列维度上再次索引,最终得到(M, P, P)的子矩阵。
3. 封装可复用函数
若需频繁提取此类子矩阵,可封装为函数:
def get_cov_submatrix(arr, *slices): """从(M,N,N)的协方差数组中提取非连续子矩阵""" idx = np.r_[slices] return arr[:, np.ix_(idx, idx)] # 使用示例:传入多个切片 subset = get_cov_submatrix(y, slice(0,2), slice(-2, None))
总结
- 非连续索引无法通过单一切片对象实现,必须借助索引数组(可通过
np.r_合并切片生成)。 np.ix_是确保正确提取二维子矩阵的关键工具,它能让索引逻辑符合矩阵子提取的直觉,避免高级索引的配对行为。
内容的提问来源于stack exchange,提问作者Tom Johnson
相关产品推荐
相关产品推荐

