基于Numpy实现三维数据批量按索引提取最后维度子集
Got it, let's ditch those nested loops and do this the proper NumPy way!
Your goal is to extract elements from the 3D values array using the 3D ind array as indices for the third dimension. The nested loop approach works, but it's not leveraging NumPy's vectorized operations—let's fix that.
The Clean NumPy Solution
For your specific array shapes (values is (4,5,10), ind is (4,5,5)), you can use advanced indexing directly with broadcasting:
import numpy as np # Your original data setup values = np.reshape(np.array([x/100 for x in range(4*5*10)]), (4, 5, 10)) ind = np.reshape(np.array([np.random.randint(0,10) for x in range(4*5*5)]), (4, 5, 5)) # The NumPy-style one-liner result = values[:, :, ind]
This works because:
values[:, :, ind]usesindas the index for the third dimension ofvalues.- NumPy automatically broadcasts the first two dimensions of
values(shape (4,5)) to match the first two dimensions ofind(shape (4,5)), then applies the 5 indices inindfor each (batch, word) pair, resulting in the desired (4,5,5) output array.
Verify It Matches the Loop Approach
To confirm this gives the same result as your nested loop code:
# Original loop implementation result_loop = np.zeros_like(ind, dtype=values.dtype) for batch_idx in range(len(values)): for word_idx in range(len(values[0])): result_loop[batch_idx][word_idx] = values[batch_idx,word_idx, ind[batch_idx, word_idx]] # Check if results are identical (within floating point tolerance) print(np.allclose(result, result_loop)) # Output: True
More General Case (For Any Number of Dimensions)
If you ever need to handle arrays with more dimensions, you can explicitly create indices for the leading dimensions using np.indices():
# Create indices for the first two dimensions batch_indices, word_indices = np.indices(values.shape[:2]) # Use all three index arrays to extract values result = values[batch_indices, word_indices, ind]
This is more verbose but works for any number of leading dimensions, making it a flexible approach for similar problems.
内容的提问来源于stack exchange,提问作者O K

