如何基于Numpy实现更高效的移动中位数计算?
看起来你在处理大规模信号数据的移动中位数计算时遇到了性能瓶颈——10000条长度750的信号,用现有的滑动窗口方法确实容易慢下来。我来分享几个能显著提速的优化思路,亲测有效!
先明确你的场景
你的数据集是每行对应一条信号列表,当前用np.lib.stride_tricks.sliding_window_view配合np.nanmedian实现移动中位数,再通过apply遍历每条信号。这个方法的问题在于:滑动窗口会创建大量临时数组,加上np.nanmedian内部的全排序操作,在数据量上来后开销会非常大。
优化方案1:用Numba JIT编译加速(最推荐)
中位数计算属于CPU密集型任务,Numba的即时编译能把Python循环的开销降到接近C语言的水平,对这类场景提升特别明显。我们可以把移动中位数的逻辑改成循环实现,再用Numba装饰:
import numpy as np from numba import jit @jit(nopython=True) def moving_median_numba(signal, n=150): signal_len = len(signal) result = np.full(signal_len, np.nan) # 按照你的逻辑填充前n-1个值:用整个信号的中位数 fill_val = np.nanmedian(signal) for i in range(n-1): result[i] = fill_val # 处理滑动窗口的中位数 for i in range(n-1, signal_len): window = signal[i - n + 1 : i + 1] # 如果你确定信号里没有NaN,可以去掉下面这行,进一步提速 window = window[~np.isnan(window)] window_len = len(window) if window_len == 0: continue half = window_len // 2 if window_len % 2 == 1: # 奇数长度:取中间值 med = np.partition(window, half)[half] else: # 偶数长度:取中间两个值的平均 lower = np.partition(window, half - 1)[half - 1] upper = np.partition(window, half)[half] med = (lower + upper) / 2 result[i] = med return signal - result
使用的时候直接替换原来的apply即可:
df['result'] = df['values'].apply(moving_median_numba)
这里我用了np.partition代替全排序——它只需要把数组分成比中位数小和大的两部分,不需要完全排序,速度比np.nanmedian默认的全排序快很多,尤其是窗口较大的时候。
优化方案2:批量处理二维数组,减少循环开销
如果不想引入Numba依赖,可以先把所有信号转成二维数组,再用向量化操作批量处理:
# 把列中的列表转成二维数组 signal_arr = np.array(df['values'].tolist()) window_size = 150 signal_count, signal_len = signal_arr.shape # 计算前window_size-1个位置的填充值(按你的逻辑用整个信号的中位数) fill_vals = np.nanmedian(signal_arr, axis=1).reshape(-1, 1) fill_part = np.repeat(fill_vals, window_size - 1, axis=1) # 批量计算滑动窗口中位数 rolling_medians = [] for row in signal_arr: # 用pandas的rolling快速计算窗口中位数 med_series = pd.Series(row).rolling(window_size).median() rolling_medians.append(med_series.dropna().values) # 拼接填充部分和计算结果,得到完整的中位数数组 full_medians = np.concatenate([fill_part, np.array(rolling_medians)], axis=1) # 最终结果 result_arr = signal_arr - full_medians # 转回到DataFrame的列格式 df['result'] = list(result_arr)
这个方法利用了pandas rolling的内部优化,比纯Python循环快,但还是不如Numba方案高效,适合不能用Numba的场景。
小提醒:检查你的前n-1个值逻辑
你原代码里前n-1个值用的是整个信号的中位数,这和常规的移动中位数逻辑(前k个值用前k个元素的中位数)不太一样。如果是故意这么设计的没问题,但如果是笔误,可以调整填充逻辑,比如把fill_val = np.nanmedian(signal)改成循环计算每个前k窗口的中位数。
备注:内容来源于stack exchange,提问作者Jesper Mølgaard

