如何缩短Pandas中分组滚动窗口apply函数的执行耗时?
Pandas滚动均值计算性能优化方案
核心优化方向:避免逐组apply,改用向量化分组操作
原代码依赖groupby('id').apply()逐组处理,这是性能瓶颈的核心——apply本质是循环调用函数,无法利用Pandas的向量化优化能力。以下是具体落地的优化方案:
1. 预处理异常值+直接调用groupby.rolling(最优推荐)
先将异常值对应的fuel_level标记为NaN,再直接使用分组后的rolling方法,全程用向量化操作替代循环:
import pandas as pd # 关键前提:确保timestamp是datetime类型(非该类型会导致rolling耗时暴增) df['timestamp'] = pd.to_datetime(df['timestamp']) # 预处理:将异常值的fuel_level设为NaN,后续rolling会自动忽略 df['fuel_level_clean'] = df['fuel_level'].where(df['flag_outlier'] == 0) # 设置时间索引,分组后直接计算滚动均值(结果自动与原数据对齐) df = df.set_index('timestamp') df['fuel_level_mean'] = df.groupby('id')['fuel_level_clean'].rolling( window='120s', closed='both', min_periods=5 ).mean().values # 恢复原数据结构并清理临时列 df = df.reset_index().drop('fuel_level_clean', axis=1)
2. 用transform简化映射逻辑
如果需要快速保留原数据所有列,可使用groupby.transform自动将分组计算结果映射回原数据:
df['timestamp'] = pd.to_datetime(df['timestamp']) df['fuel_level_clean'] = df['fuel_level'].where(df['flag_outlier'] == 0) # 用transform替代apply,直接生成与原数据对齐的结果列 df['fuel_level_mean'] = df.set_index('timestamp').groupby('id')['fuel_level_clean'].transform( lambda x: x.rolling('120s', closed='both', min_periods=5).mean() ).reset_index(level=0, drop=True) df = df.drop('fuel_level_clean', axis=1)
3. 极端性能需求:Numba编译自定义滚动逻辑
若上述优化仍无法满足速度要求,可使用Numba编译手动实现的滚动函数,绕过Pandas内置rolling的额外开销:
from numba import jit import numpy as np # Numba编译的快速滚动均值函数(基于秒级时间戳计算) @jit(nopython=True) def fast_rolling_mean(values, timestamps, window_sec=120, min_periods=5): n = len(values) result = np.full(n, np.nan) for i in range(n): end_ts = timestamps[i] start_ts = end_ts - window_sec # 筛选窗口内的有效非NaN数据 mask = (timestamps >= start_ts) & (timestamps <= end_ts) valid_data = values[mask][~np.isnan(values[mask])] if len(valid_data) >= min_periods: result[i] = valid_data.mean() return result # 预处理时间戳为秒级整数,转换数值列为数组 df['timestamp_sec'] = df['timestamp'].astype(np.int64) // 10**9 df['fuel_level_clean'] = df['fuel_level'].where(df['flag_outlier'] == 0).values # 分组应用Numba函数并映射回原数据 df['fuel_level_mean'] = df.groupby('id').apply( lambda x: fast_rolling_mean(x['fuel_level_clean'], x['timestamp_sec']) ).explode().values # 清理临时列 df = df.drop(['timestamp_sec', 'fuel_level_clean'], axis=1)
额外性能小贴士
- 强制检查时间类型:确保
timestamp是datetime64[ns]类型,字符串/object类型会导致rolling隐式转换,耗时翻倍。 - 精简数据列:计算前仅保留必要列(id、timestamp、fuel_level、flag_outlier),减少内存占用。
- 精度换速度:若业务允许,将
fuel_level转为float32类型,降低内存开销间接提升计算效率。
内容的提问来源于stack exchange,提问作者Sangeetha R
相关产品推荐
相关产品推荐

