如何快速获取满足条件的两个Numpy数组的有效索引?
优化方案:无需循环的Numpy向量化解法
首先我们可以把原判断条件化简,避开复杂的指数运算:
原条件f(x1, x2) > 0可以一步步简化为:
exp(x2 - x1)/1.2 - 1 > 0
→exp(x2 - x1) > 1.2
→x2 - x1 > np.log(1.2)
→x2 > x1 + np.log(1.2)
这个化简是核心,直接把函数运算转换成了简单的数值比较,接下来我们可以基于这个简化后的条件,用O(N)时间复杂度的方法得到结果,完全不需要循环,也不用生成10000×10000的巨型数组。
推导有效索引的逻辑
找a1的有效索引
对于a1中的元素x1,只要a2里存在至少一个x2满足x2 > x1 + np.log(1.2),x1就是有效元素:
- 如果x1 <
np.max(a2) - np.log(1.2),那a2的最大值必然大于x1 + np.log(1.2),满足条件; - 如果x1 ≥
np.max(a2) - np.log(1.2),那a2中所有元素都≤x1 + np.log(1.2),无法满足条件。
找a2的有效索引
对于a2中的元素x2,只要a1里存在至少一个x1满足x2 > x1 + np.log(1.2),x2就是有效元素:
- 如果x2 >
np.min(a1) + np.log(1.2),那a1的最小值必然满足x2 > x1 + np.log(1.2),满足条件; - 如果x2 ≤
np.min(a1) + np.log(1.2),那a1中所有元素都≥x2 - np.log(1.2),无法满足条件。
完整优化代码
import numpy as np a1 = np.random.rand(10000) a2 = np.random.rand(10000) # 预计算化简条件所需的常数 C = np.log(1.2) # 获取a1的有效索引 valid_i1 = np.where(a1 < (np.max(a2) - C))[0] # 获取a2的有效索引 valid_i2 = np.where(a2 > (np.min(a1) + C))[0] # 输出排序后的结果(和原代码输出格式一致) print(sorted(valid_i1)) print(sorted(valid_i2))
优化效果
- 时间复杂度从原代码的O(N²)降到O(N),对于10000个元素的数组,运行速度会提升几个数量级;
- 完全避免了生成大数组,内存占用可以忽略;
- 结果和原代码完全一致,逻辑严格等价。
内容的提问来源于stack exchange,提问作者Shaun Han
相关产品推荐
相关产品推荐

