如何用向量化方法获取布尔数组axis-1上值为1的行索引
Awesome question! Ditching loops for vectorized operations in NumPy is always a smart move—let's walk through how to solve this cleanly and efficiently, no loops required.
First, let's recap your problem to make sure we're on the same page:
- You have a 2D NumPy array where each row contains exactly one
1(rest are0s) - You need to collect row indices of all
1s, processing columns in order (left to right), and group all row indices for a column before moving to the next - Your desired output is
array([ 1, 4, 5, 6, 0, 7, 8, 2, 9, 10, 3, 11, 12])
Step-by-Step Vectorized Approach
Here's how to achieve this with pure NumPy vectorized operations:
Locate all
1s in the array
Usenp.where()to get the row and column indices of every1in the array. This function returns two arrays: one with row positions, and one with corresponding column positions.Sort row indices by their column
Usenp.argsort()on the column indices to get an ordering that groups all rows from the same column together (starting with column 0, then 1, etc.). Apply this ordering to the row indices array to get your final result.
Full Code Implementation
import numpy as np # Your input array arr = np.array([[0, 1, 0, 0], [1, 0, 0, 0], [0, 0, 1, 0], [0, 0, 0, 1], [1, 0, 0, 0], [1, 0, 0, 0], [1, 0, 0, 0], [0, 1, 0, 0], [0, 1, 0, 0], [0, 0, 1, 0], [0, 0, 1, 0], [0, 0, 0, 1], [0, 0, 0, 1]], dtype=np.uint8) # Get row and column indices of all 1s rows, cols = np.where(arr == 1) # Sort rows by their column to group them in column order result = rows[np.argsort(cols)] print(result) # Output: array([ 1, 4, 5, 6, 0, 7, 8, 2, 9, 10, 3, 11, 12])
Why This Works
np.where(arr == 1)scans the array and collects every position where the value is1. For your input, this gives us:rows:array([0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12])(all row indices with a1)cols:array([1, 0, 2, 3, 0, 0, 0, 1, 1, 2, 2, 3, 3])(corresponding column for each row)
np.argsort(cols)generates an index array that sorts the column indices in ascending order. For ourcolsarray, this index is[1,4,5,6,0,7,8,2,9,10,3,11,12]- When we use this index to slice the
rowsarray, we reorder the rows to group all entries from column 0 first, then column 1, etc.—exactly matching your desired output.
Bonus: Efficiency
This approach is fully vectorized, meaning all operations happen in optimized C-level code rather than slow Python loops. For large arrays, this will be drastically faster than any loop-based implementation.
内容的提问来源于stack exchange,提问作者kmario23

