如何使用numpy.where结合双索引获取数组符号变化的元素索引
用NumPy找出数组中符号变化的元素索引
你要实现的是找出数组中相邻元素发生符号变化的位置索引(对应示例中的v2),用NumPy可以不用循环,直接通过向量操作高效完成,步骤如下:
- 先把列表转为NumPy数组:
import numpy as np v = np.array([1, 2, -1, 2, 3, -1, 3, -10, -10, -10])
- 计算数组每个元素的符号,然后比较相邻元素的符号是否不同:
# 获取每个元素的符号(正为1,负为-1,0为0) signs = np.sign(v) # 找出前一个元素和后一个元素符号不同的位置索引 v2 = np.where(signs[:-1] != signs[1:])[0]
运行后v2就是array([1, 2, 4, 5, 6]),和你要的结果一致。
逻辑说明
np.sign(v)快速得到每个元素的符号,避免手动判断正负;signs[:-1]取除了最后一个元素的所有符号,signs[1:]取除了第一个元素的所有符号,两者对应位置比较就能找出相邻元素符号变化的位置;np.where()返回符合条件的索引数组,正好是你需要的结果。
另外你原来的循环代码逻辑有问题:for i in range(len(v)-1)搭配v[i] * v[i-1] < 0会在i=0时取到数组最后一个元素(v[-1]),导致逻辑错误。正确的循环逻辑应该是遍历前n-1个元素,比较当前元素和下一个元素的符号,比如:
v = [1, 2, -1, 2, 3, -1, 3, -10, -10, -10] v2 = [] for i in range(len(v)-1): if v[i] * v[i+1] < 0: v2.append(i)
这样得到的v2也是[1,2,4,5,6],和NumPy方法的结果一致。
内容的提问来源于stack exchange,提问作者Dragos Efrim
相关产品推荐
相关产品推荐

