如何在无循环的情况下获取二维NumPy数组每行首个正元素的索引
解决方案
你的代码问题出在两个地方:
row_idx用np.any得到的是布尔数组,不是期望的整数行索引数组np.argmax(data>0, axis=1)会给没有正元素的行返回0(因为全False的数组argmax默认取第一个位置),这会导致后续取到错误的元素
下面是修正后的代码,完全无需循环:
import numpy as np data = np.array([[1.0, -2.0, 3.0, 3], [0.0, 1.5, -2.0, 5], [-1.0, -2.0, -3.0, -5]]) # 标记每行是否存在正元素 has_positive = np.any(data > 0, axis=1) # 获取每行第一个正元素的列索引,过滤掉无正元素的行 col_idx = np.argmax(data > 0, axis=1)[has_positive] # 获取有正元素的行的整数索引 row_idx = np.where(has_positive)[0] # 验证结果 print('row indices: ', row_idx) # 输出 [0 1] print('col indices: ', col_idx) # 输出 [0 1] print('target elements: ', data[row_idx, col_idx]) # 输出 [1. 1.5]
关键逻辑说明
np.argmax(data>0, axis=1):对每行的布尔数组(正元素为True)取第一个True的位置,正好对应首个正元素的列索引np.where(has_positive)[0]:把布尔掩码转换成对应的整数行索引,只保留存在正元素的行- 通过
[has_positive]过滤列索引,去掉那些无正元素行的无效列值
内容的提问来源于stack exchange,提问作者xleData
相关产品推荐
相关产品推荐

