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

如何用向量化方法获取布尔数组axis-1上值为1的行索引

Vectorized Solution for Collecting Row Indices by Column in NumPy

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 are 0s)
  • 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:

  1. Locate all 1s in the array
    Use np.where() to get the row and column indices of every 1 in the array. This function returns two arrays: one with row positions, and one with corresponding column positions.

  2. Sort row indices by their column
    Use np.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 is 1. 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 a 1)
    • 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 our cols array, this index is [1,4,5,6,0,7,8,2,9,10,3,11,12]
  • When we use this index to slice the rows array, 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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.13 07:48:51