Numpy数组逐元素比较生成符合条件输出数组的高效实现方法
NumPy多输入数组符合条件输出的高效实现方案
用NumPy原生向量化运算实现,全程无Python层级循环,性能远高于遍历写法,同时支持任意数量的输入数组,适配性强。
实现步骤
- 将所有输入的一维数组建堆为二维数组,统一做逐元素运算
- 逐元素判断两个必要条件:
- 数值大于0.5
- 数值为同位置所有输入数组中的最大值
- 两个条件同时满足的位置标记为1,否则为0,拆分后得到对应每个输入的输出数组
代码示例
import numpy as np # 示例输入 in_1 = np.array([0.4, 0.7, 0.8, 0.3]) in_2 = np.array([0.9, 0.8, 0.6, 0.4]) # 所有输入数组存入列表,新增输入只需扩展该列表即可 in_arrays = [in_1, in_2] # 沿第0轴堆叠所有输入数组 stacked_arr = np.stack(in_arrays, axis=0) # 计算两个条件 over_threshold = stacked_arr > 0.5 is_max = stacked_arr == stacked_arr.max(axis=0, keepdims=True) # 生成结果并拆分 res = (over_threshold & is_max).astype(int) out_1, out_2 = res
运行后得到的out_1为array([0, 0, 1, 0]),out_2为array([1, 1, 0, 0]),和示例要求完全一致。
内容的提问来源于stack exchange,提问作者sixtytrees
相关产品推荐
相关产品推荐

