如何用numpy.where或等效函数向量化重构迭代信号生成算法
用Numpy向量化重构状态保持型信号生成函数
你的原函数是一个典型的状态保持类迭代逻辑:根据输入数组r的阈值触发信号状态切换(<=30时切为1,>=60时切为0),未触发时保持上一次的信号状态。下面我给你一个完全基于Numpy向量化操作的高效实现,彻底摆脱循环:
核心思路
我们可以通过追踪每个位置最后一次触发1的索引和最后一次触发0的索引,然后比较这两个索引的先后,就能直接确定每个位置的信号状态:
- 如果最后一次触发1的时间晚于触发0,当前状态为1
- 反之则为0(初始状态为0,对应没有任何触发的情况)
实现代码
import numpy as np def generate_signal_vectorized(r): # 标记触发1的位置,并计算每个位置及之前最后一次触发1的索引(未触发则为-1) trigger_1 = r <= 30 last_trigger_1 = np.maximum.accumulate(np.where(trigger_1, np.arange(len(r)), -1)) # 标记触发0的位置,并计算每个位置及之前最后一次触发0的索引(未触发则为-1) trigger_0 = r >= 60 last_trigger_0 = np.maximum.accumulate(np.where(trigger_0, np.arange(len(r)), -1)) # 比较两次触发的先后,生成最终信号 return (last_trigger_1 > last_trigger_0).astype(int)
验证效果
用你给出的测试数组验证:
r = np.array([50, 52, 59, 69, 47, 33, 27, 26, 20, 30, 33, 35, 58, 55, 48, 60, 68, 55, 43, 49, 33, 30, 22, 28]) s_original = generate_signal(r) s_vectorized = generate_signal_vectorized(r) # 验证结果一致 print(np.array_equal(s_original, s_vectorized)) # 输出:True
效率对比
这种向量化方式的时间复杂度是O(n),但比原循环实现快得多(尤其是当r的长度很大时),因为Numpy的底层操作是用C实现的,避免了Python循环的开销。
内容的提问来源于stack exchange,提问作者php-junkie
相关产品推荐
相关产品推荐

