如何高效计算满足行级条件的Rolling Mean?
优化基于行级条件的滚动均值计算(适配大型DataFrame)
需求说明
需要对DataFrame中满足finished_time < created_time的行,计算最近window个target值的滚动均值,现有方法时间复杂度为O(n²),无法适配大型数据集,需要优化。
数据重建代码
import numpy as np import pandas as pd np.random.seed(21) created_time = pd.to_datetime('2023-07-21 00:10:00') + pd.timedelta_range(start='0 minutes', periods=10, freq='15T') finished_time = created_time + pd.to_timedelta(np.random.randint(1, 30, size=10), unit='minutes') target = np.random.randint(1,100, size=10) df = pd.DataFrame(data={'created_time': created_time, 'finished_time': finished_time, 'target': target})
原始低效代码
def custom_roll(row, df, window): rolling_mean = ( df.target [df.finished_time < row.created_time] .iloc[:window] .mean() ) return rolling_mean df.apply(lambda row: custom_roll(row, df, window=3), axis=1)
优化方案(O(n log n)时间复杂度)
核心思路:利用排序+二分查找+前缀和实现向量化计算,避免逐行遍历。
步骤说明
- 按
finished_time排序数据,将时间和目标值转为数组,方便后续二分查找和前缀和计算; - 对每一行的
created_time,用二分查找快速定位所有满足finished_time < created_time的行的边界; - 通过前缀和数组快速计算最近window个元素的总和,再求均值。
实现代码
window = 3 # 1. 按finished_time排序,提取数组 sorted_df = df.sort_values('finished_time').reset_index(drop=True) sorted_finished = sorted_df['finished_time'].values sorted_target = sorted_df['target'].values # 2. 计算前缀和数组,用于快速求和 prefix_sum = np.concatenate([[0], np.cumsum(sorted_target)]) # 3. 二分查找定位每个created_time对应的边界 indices = np.searchsorted(sorted_finished, df['created_time'].values, side='left') # 4. 确定每个边界对应的起始位置(最多取window个元素) start_indices = np.maximum(0, indices - window) # 5. 计算区间和与均值 rolling_sums = prefix_sum[indices] - prefix_sum[start_indices] counts = indices - start_indices # 无满足条件的元素时返回NaN rolling_means = np.where(counts > 0, rolling_sums / counts, np.nan) # 赋值回原DataFrame df['rolling_mean'] = rolling_means
效果验证
运行后得到的rolling_mean列与原始方法结果一致,但计算效率提升几个数量级,可轻松处理十万级甚至百万级数据集。
内容的提问来源于stack exchange,提问作者malwin
相关产品推荐
相关产品推荐

