NumPy花式索引:为不同子数组裁剪不同位置相同大小ROI
Vectorized ROI Cropping with NumPy Fancy Indexing
Absolutely, we can replace that loop with a clean vectorized approach using NumPy's fancy indexing—this will be much faster, especially when working with larger n values. Here's the step-by-step solution:
Full Implementation Code
import numpy as np # Initialize your original data n, c, h, w = 3, 1, 4, 4 data = np.arange(n * c * h * w).reshape(n, c, h, w) size = 2 locations = np.array([[0, 1], [1, 1], [0, 2]]) # Generate the range of y and x indices for each sample's ROI y_ranges = locations[:, 0, None] + np.arange(size) # Shape: (n, size) x_ranges = locations[:, 1, None] + np.arange(size) # Shape: (n, size) # Use fancy indexing to extract all ROIs in one go crops = data[ np.arange(n)[:, None, None], # Index each sample, with extra dims for broadcasting :, # Keep all channels y_ranges[:, :, None], # Broadcast y indices to match x's grid x_ranges[:, None, :] # Broadcast x indices to match y's grid ] # Verify the result matches your loop-based output print(crops)
How It Works
Let's break down the key parts:
- Generate ROI Indices:
y_rangescreates the sequence of y-coordinates for each sample's ROI (e.g., first sample gets[0, 1], second gets[1, 2]).x_rangesdoes the same for x-coordinates (first sample gets[1, 2], third gets[2, 3]).- We add
Noneto reshape these arrays, which lets us broadcast them into a 2D grid of indices later.
- Fancy Indexing:
np.arange(n)[:, None, None]adds two extra dimensions to the sample indices, ensuring they line up correctly with the 2D ROI grids for each sample.y_ranges[:, :, None]andx_ranges[:, None, :]reshape the index ranges so they broadcast together into asize×sizegrid for each sample. This tells NumPy to grab every combination of y and x in the ROI for each sample.
- Output Shape:
- The result is already in the
(n, c, size, size)shape you need—no extra reshaping required (though you can addreshape(n, c, size, size)explicitly if you want to make the shape explicit for readability).
- The result is already in the
Result Verification
Running this code produces exactly the same output as your loop-based method:
array([[[[ 1, 2], [ 5, 6]]], [[[21, 22], [25, 26]]], [[[34, 35], [38, 39]]]])
This approach leverages NumPy's optimized C-level operations instead of slow Python loops, so it scales much better when you have a large number of samples (n).
内容的提问来源于stack exchange,提问作者DiKorsch
相关产品推荐
相关产品推荐

