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

如何在Python中根据输入形状与块形状生成n维数组切片索引

Hey there! Let's tackle your question with two straightforward parts, just like you requested.

1. How to split an n-dimensional NumPy array into blocks of a specified size

Splitting an n-dimensional array into blocks boils down to generating valid slice indices for each block, then using those slices to extract subarrays from the original array. Here's a step-by-step breakdown:

  • Calculate valid starting positions: For each dimension, if the input size is input_dim and the block size is block_dim, valid starting indices range from 0 to input_dim - block_dim (inclusive). This guarantees every block fits entirely within the original array.
  • Combine starting positions: Generate all combinations of these starting indices across all dimensions (this is the Cartesian product of the starting points for each axis).
  • Create slice tuples: For each combination of starting positions, build a tuple of slice objects (e.g., slice(start, start + block_dim) for each axis), then use this tuple to index the original array and extract the block.

Here's a concrete example using your sample shapes:

import numpy as np
import itertools

# Create a sample array with shape (2, 2, 3)
arr = np.arange(12).reshape(2,2,3)

# Define target block shape
block_shape = (2,2,2)

# We'll use the `foo` function we build next to generate slices
def foo(input_shape, block_shape):
    starts_per_dim = []
    for in_dim, block_dim in zip(input_shape, block_shape):
        if block_dim > in_dim:
            raise ValueError(f"Block dimension {block_dim} can't be larger than input dimension {in_dim}")
        num_starts = in_dim - block_dim + 1
        starts_per_dim.append(range(num_starts))
    
    block_slices = []
    for start_tuple in itertools.product(*starts_per_dim):
        slice_tuple = tuple(slice(s, s + bd) for s, bd in zip(start_tuple, block_shape))
        block_slices.append(slice_tuple)
    return block_slices

# Get all block slices
slices = foo(arr.shape, block_shape)

# Extract each block from the array
blocks = [arr[sl] for sl in slices]

# Print results to verify
for i, block in enumerate(blocks):
    print(f"Block {i+1}:\n{block}\n")

Running this will output two blocks of shape (2,2,2), exactly as expected.

2. Implementing the foo function

The foo function takes an input array shape and block shape, then returns a list of slice tuples for each valid block. Here's the full implementation with validation checks and clear logic:

import itertools

def foo(input_shape, block_shape):
    # Validate dimensionality match
    if len(input_shape) != len(block_shape):
        raise ValueError("Input shape and block shape must have the same number of dimensions")
    
    starts_per_dim = []
    for axis_idx, (in_dim, block_dim) in enumerate(zip(input_shape, block_shape)):
        # Validate positive block dimensions
        if block_dim <= 0:
            raise ValueError(f"Block dimension for axis {axis_idx} must be a positive integer")
        # Ensure block fits within input dimension
        if block_dim > in_dim:
            raise ValueError(f"Block dimension {block_dim} exceeds input dimension {in_dim} for axis {axis_idx}")
        
        # Calculate all valid starting indices for this axis
        num_starts = in_dim - block_dim + 1
        starts_per_dim.append(range(num_starts))
    
    # Generate all slice tuples from starting position combinations
    block_slices = []
    for start_coords in itertools.product(*starts_per_dim):
        current_slice = tuple(
            slice(start, start + block_dim)
            for start, block_dim in zip(start_coords, block_shape)
        )
        block_slices.append(current_slice)
    
    return block_slices

Testing the function

Let's test it with your sample input:

result = foo((2,2,3), (2,2,2))
print(result)

This will output exactly what you requested:

[(slice(0, 2, None), slice(0, 2, None), slice(0, 2, None)), (slice(0, 2, None), slice(0, 2, None), slice(1, 3, None))]

A few key notes about this function:

  • It includes validation to catch incompatible shapes or invalid block sizes early.
  • itertools.product efficiently handles generating all combinations of starting positions, making the function work seamlessly for any number of dimensions.
  • Each slice tuple is built dynamically to match the block size for each axis.

内容的提问来源于stack exchange,提问作者Pranjal Sharma

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.13 08:37:25