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

如何组合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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.24 01:52:14