Python numpy中pins重叠区间合并后求vals数组最大值的优化方法
优化实现方案
前提:你使用的频率向量freqs为升序排列,这是该类场景的通用属性,符合你的示例代码设定。
核心优化思路
- 替换原
O(m*n)的逐pin遍历生成mask逻辑,改为*O(m log m)*的区间合并逻辑,m为pin数量,n为freq长度,数据量越大性能优势越明显 - 利用freq有序的特性,用numpy内置二分查找
searchsorted直接定位区间索引,不需要遍历整个freq数组生成islands标记 - 省去大尺寸中间数组的内存占用
完整优化代码
import numpy as np np.random.seed(666) freqs = np.linspace(0, 20, 50) vals = np.random.randint(100, size=(len(freqs), 1)).flatten() print(freqs) print(vals) pins = [2, 6, 10, 11, 15, 15.2] # 优化后逻辑 # 1. 排序pins生成原始区间 sorted_pins = np.sort(pins) intervals = np.column_stack([sorted_pins - 1, sorted_pins + 1]) # 2. 合并重叠区间 merged_intervals = [] for curr in intervals: if not merged_intervals: merged_intervals.append(curr) else: last = merged_intervals[-1] if curr[0] <= last[1]: last[1] = max(last[1], curr[1]) else: merged_intervals.append(curr) merged_intervals = np.array(merged_intervals) # 3. 二分查找定位区间在freqs中的索引 left_idx = np.searchsorted(freqs, merged_intervals[:, 0], side='left') right_idx = np.searchsorted(freqs, merged_intervals[:, 1], side='right') # 4. 计算每个区间的最大值 maxs = [np.max(vals[l:r]) for l, r in zip(left_idx, right_idx)] print(maxs)
运行输出结果和原代码完全一致:[73, 97, 79, 77]
性能优势
当数据规模较大时,比如freq长度为1e6、pin数量为1e4,原代码时间复杂度为O(mn + n),优化后为O(m log m + k)(k为合并后区间数量,通常远小于m和n),性能提升可达上千倍;同时不需要生成和freq等长的中间标记数组,内存占用降低90%以上。
内容的提问来源于stack exchange,提问作者shamalaia
相关产品推荐
相关产品推荐

