Numpy一维数组分块重塑及轴变换高效实现技术问询
Got it, let's break down this reshaping problem clearly—this is a classic axis manipulation task where getting the order right is key, especially since you need efficiency for large n/h and iterative runs.
First, let's restate the problem with concrete variables to make it easier to follow:
import numpy as np n = 2 h = 2 k = 2 * h + 1 # This is 5 in your example m = n * k # 10 in your example input_data = np.arange(m ** 2) # Generates 0-99 as in your example
The Core Idea
Your goal is to split the 1D array into k×k blocks (each n×n), then stack those blocks row-wise to form an m×m matrix. The trick is to structure the axes so we can merge them without copying data (critical for performance).
Step-by-Step Efficient Solution
We can do this in 3 view-based operations (no data copying, super fast even for large arrays):
Reshape directly into a 4D array of blocks
Skip the intermediate(k², n, n)step—we can go straight from 1D to a 4D array that represents(number of block rows, number of block columns, block rows, block columns):blocks_4d = input_data.reshape(k, k, n, n)For your example, this creates a
(5,5,2,2)array whereblocks_4d[i,j]is then×nblock at rowi, columnjof the block grid.Rearrange axes to group block rows with inner block rows
Right now, the axes are(block_row, block_col, inner_row, inner_col). We need to swap theblock_colandinner_rowaxes so we can merge adjacent axes later:rearranged = blocks_4d.transpose(0, 2, 1, 3)Now the axes are
(block_row, inner_row, block_col, inner_col)—this aligns the dimensions we want to merge next.Merge axes to get the final
m×mmatrix
We just need to combine the first two axes (block_row+inner_row=mrows) and the last two axes (block_col+inner_col=mcolumns):final_matrix = rearranged.reshape(m, m)
Verify the Result
Let's check the first two rows to match your expected output:
print(final_matrix[:2]) # Output: # [[ 0 1 4 5 8 9 12 13 16 17] # [ 2 3 6 7 10 11 14 15 18 19]]
If You Already Have the (k², n, n) Array
If you're starting from the intermediate blocks = input_data.reshape(k**2, n, n) you mentioned, just add one extra reshape step first:
blocks_4d = blocks.reshape(k, k, n, n) # Then proceed with transpose and reshape as above
Why This Is Efficient
All operations here (reshape and transpose) are view operations when possible (which they are here, since our input array is C-contiguous). This means NumPy doesn't copy any data—it just changes how it interprets the underlying memory. This is crucial for large n/h and iterative runs, as it keeps the operation nearly instantaneous regardless of array size.
内容的提问来源于stack exchange,提问作者yacola

