如何在高维NumPy数组中遍历所有长度为m的一维子数组?
如何在高维NumPy数组中遍历所有长度为m的一维子数组?
我完全理解你的需求啦——你有一个n维的NumPy数组,每个维度的长度都是m,想要找出所有长度为m的一维子数组,就像2维单位矩阵里要包含行、列、两条对角线那样。这个问题在高维场景下确实需要点巧思,因为切片的位置得灵活变换,咱们一步步拆解来解决:
一、先明确要覆盖的两类一维子数组
首先得理清,你要找的一维子数组其实分为两种核心类型:
- 单轴切片型:固定n-1个维度的索引,剩下一个维度取完整切片(比如2维里的行、列);
- 对角线序列型:所有维度的索引按同步规则变化,逐个取元素组成一维数组(比如2维里的主、副对角线)。
接下来咱们分别搞定这两类,再把结果合并起来。
二、处理单轴切片型的子数组
这类子数组的核心逻辑是:遍历每个轴,生成其他所有轴的索引组合,然后对当前轴取全切片。比如3维数组中,要生成x[i,j,:]、x[i,:,j]、x[:,i,j]这三种形式的所有切片。
可以用itertools.product来生成固定维度的索引组合,再动态构建切片对象:
import numpy as np from itertools import product def get_single_axis_slices(arr): n = arr.ndim m = arr.shape[0] slices = [] # 遍历每个轴,作为要取全切片的目标轴 for axis in range(n): # 确定需要固定的其他维度 fixed_dims = [d for d in range(n) if d != axis] # 生成所有固定维度的索引笛卡尔积(每个维度取0到m-1的所有值) for indices in product(range(m), repeat=len(fixed_dims)): # 构建切片元组:固定维度用对应索引,目标轴用:(即slice(None)) slice_tuple = list(indices) slice_tuple.insert(axis, slice(None)) slice_tuple = tuple(slice_tuple) # 取出对应的一维数组并加入结果 slices.append(arr[slice_tuple]) return slices
测试2维单位矩阵的情况:
x = np.identity(4) single_axis_slices = get_single_axis_slices(x) # 这里会包含4行+4列,共8个一维数组,和你例子里的前两部分完全一致
三、处理对角线序列型的子数组
这类的核心是找到所有“索引随i同步变化”的序列,i从0到m-1,每个维度的索引可以是i(正序)或m-1-i(逆序),对应高维场景下的各种对角线方向。
def get_diagonal_arrays(arr): n = arr.ndim m = arr.shape[0] diagonals = [] # 生成所有可能的索引规则组合:每个位置True表示用i,False表示用m-1-i from itertools import product as bool_product for signs in bool_product([True, False], repeat=n): diag = [] for i in range(m): # 按当前规则生成每个维度的索引 indices = tuple(i if s else (m-1 - i) for s in signs) diag.append(arr[indices]) diagonals.append(np.array(diag)) # 如果你想去掉完全反向的重复对角线(比如正序主对角线和逆序主对角线),可以加去重逻辑 # 比如:diagonals = list({tuple(d) for d in diagonals}) return diagonals
测试2维单位矩阵的情况:
diag_arrays = get_diagonal_arrays(x) # 这里会包含4种对角线:(i,i)、(i,3-i)、(3-i,i)、(3-i,3-i) # 对应你例子里的两条对角线,以及它们的反向版本,按需去重即可
四、合并两类结果
把上面两部分的结果合并,就是你要的所有长度为m的一维子数组:
def get_all_1d_subarrays(arr): single_axis = get_single_axis_slices(arr) diagonals = get_diagonal_arrays(arr) # 按需去重(比如某些特殊场景下对角线和单轴切片可能重复) return single_axis + diagonals
测试2维单位矩阵:
x = np.identity(4) all_subarrays = get_all_1d_subarrays(x) # 包含4行+4列+4条对角线(未去重),和你例子的核心需求匹配
五、高维数组的测试
比如3维的“单位立方体”数组(每个维度长度3):
# 创建3维数组,仅当i=j=k时为1,其余为0 x_3d = np.zeros((3,3,3)) for i in range(3): x_3d[i,i,i] = 1 all_3d_subarrays = get_all_1d_subarrays(x_3d) # 结果包含: # 单轴切片:9个x[i,j,:] + 9个x[i,:,j] +9个x[:,i,j] → 共27个 # 对角线:8种规则组合生成的数组 → 共8个 # 总共有35个长度为3的一维数组
这样就完美实现高维数组的遍历需求啦!
备注:内容来源于stack exchange,提问作者Matt
相关产品推荐
相关产品推荐

