未知迭代维度数量时,如何遍历数组的内部轴?
动态遍历数组任意内部维度(固定保留最后两维)
核心思路
利用itertools.product生成待遍历维度的所有索引组合,再通过动态拼接切片元组的方式,无需提前知晓遍历维度数量即可实现任意内部维度的遍历,同时固定保留最后两个维度。
实现步骤
- 确定待遍历维度的形状:通过数组形状的切片获取待遍历维度的长度,比如
x.shape[1:-2]表示从第1轴到倒数第3轴(最后两轴-2、-1固定保留)。 - 生成索引组合:用
itertools.product遍历待遍历维度的所有索引组合,每个组合是一个元组(长度等于待遍历维度的数量)。 - 构建完整切片索引:将前缀切片(
slice(None)等价于:)、索引组合元组、后缀切片(最后两维的:)拼接成完整的索引元组,直接用于数组索引。
代码示例
示例1:4维数组,遍历第1轴
import itertools import numpy as np # 定义4维数组 I, J, K, L = 2, 3, 4, 5 x = np.random.rand(I, J, K, L) # 遍历第1轴(对应J维度) for idx_tuple in itertools.product(*map(range, x.shape[1:-2])): # 拼接完整索引:[:, idx, :, :] full_idx = (slice(None),) + idx_tuple + (slice(None), slice(None)) current_slice = x[full_idx] print(f"当前切片形状: {current_slice.shape}") # 输出 (2,4,5),符合预期
示例2:5维数组,遍历第1、2轴
I, J, K, L, M = 2, 3, 4, 5, 6 x = np.random.rand(I, J, K, L, M) # 遍历第1、2轴(对应J、K维度) for idx_tuple in itertools.product(*map(range, x.shape[1:-2])): # 拼接完整索引:[:, idx1, idx2, :, :] full_idx = (slice(None),) + idx_tuple + (slice(None), slice(None)) current_slice = x[full_idx] print(f"当前切片形状: {current_slice.shape}") # 输出 (2,5,6),符合预期
示例3:6维数组,遍历第1、2、3轴
I, J, K, L, M, N = 2, 3, 4, 5, 6, 7 x = np.random.rand(I, J, K, L, M, N) # 遍历第1、2、3轴(对应J、K、L维度) for idx_tuple in itertools.product(*map(range, x.shape[1:-2])): # 拼接完整索引:[:, idx1, idx2, idx3, :, :] full_idx = (slice(None),) + idx_tuple + (slice(None), slice(None)) current_slice = x[full_idx] print(f"当前切片形状: {current_slice.shape}") # 输出 (2,6,7),符合预期
灵活调整遍历维度
如果需要遍历的不是从第1轴开始的内部维度,只需调整x.shape的切片范围即可。比如5维数组要遍历第2、3轴,只需将x.shape[1:-2]改为x.shape[2:-2],同时前缀切片调整为(slice(None), slice(None)):
for idx_tuple in itertools.product(*map(range, x.shape[2:-2])): full_idx = (slice(None), slice(None)) + idx_tuple + (slice(None), slice(None)) current_slice = x[full_idx] # 此时切片为[:, :, idx1, idx2, :](对应5维数组)
内容的提问来源于stack exchange,提问作者bayes2021
相关产品推荐
相关产品推荐

