如何在Polars中高效计算新增分组数据的累积中位数?
高效计算分组增量累积中位数(Polars实现)
要解决仅针对新增数据计算累积中位数、复用历史数据的问题,核心思路是维护每个分组的有序数据集合,新增数据时直接插入到对应集合的正确位置,基于更新后的集合快速计算中位数,避免重复遍历整个分组的历史数据。
实现方案
Polars本身没有内置的增量中位数计算API,但可以结合Python的bisect模块实现有序插入,同时维护每个分组的历史状态:
- 初始化历史状态:先处理已有数据,为每个分组构建有序列表,并记录当前累积中位数。
- 处理新增数据:按分组拆分新增数据,对每个分组的新增值,逐个插入到对应有序列表的正确位置,然后直接计算当前中位数。
- 合并结果:将计算出的增量累积中位数关联回新增数据的DataFrame。
代码示例
假设我们已有历史数据的状态,现在处理新增数据:
import polars as pl import bisect # 初始化历史分组状态:key为分组ID,value为(有序值列表, 当前累积中位数) group_state = {} # 先处理原始历史数据(模拟首次初始化) original_df = pl.DataFrame( { "group": [0, 0, 0, 1, 1, 1, 2, 2, 2], "value": [20, 40, 30, 2, 4, 3, 200, 400, 300], } ) # 初始化每个分组的有序列表和中位数 for group, values in original_df.group_by("group").agg(pl.col("value")).iter_rows(): sorted_vals = sorted(values) n = len(sorted_vals) if n % 2 == 1: median = sorted_vals[n//2] else: median = (sorted_vals[n//2 -1] + sorted_vals[n//2])/2 group_state[group] = (sorted_vals, median) # 模拟新增数据 new_data = pl.DataFrame( { "group": [0, 1, 2], "value": [25, 5, 250], } ) # 处理新增数据,计算增量累积中位数 def compute_incremental_median(row): group, val = row sorted_vals, _ = group_state[group] # 插入到有序列表的正确位置 bisect.insort(sorted_vals, val) n = len(sorted_vals) # 计算当前中位数 if n % 2 == 1: new_median = sorted_vals[n//2] else: new_median = (sorted_vals[n//2 -1] + sorted_vals[n//2])/2 # 更新分组状态 group_state[group] = (sorted_vals, new_median) return new_median # 对新增数据应用计算 result_df = new_data.with_columns( pl.struct(["group", "value"]) .map_elements(compute_incremental_median, return_dtype=pl.Float64) .alias("median") ) print(result_df)
效率对比
- 原方案(
cumulative_eval):对每个分组的所有数据,每个累积窗口都要排序计算中位数,时间复杂度为O(n² log n)(n为分组数据量)。 - 增量方案:每个新增数据插入有序列表的时间是O(log m)(m为分组历史数据量),中位数计算是O(1),整体时间复杂度为O(k log m)(k为新增数据量),效率提升显著,尤其是分组数据量大、新增数据量小的场景。
注意事项
- 如果服务需要重启,要把
group_state持久化存储(比如用文件、数据库),避免每次重启都重新初始化历史数据。 - 若分组数量极大或单分组数据量超大规模,可考虑用更高效的有序数据结构(比如
SortedListfromsortedcontainers库,插入效率比bisect更高)。
内容的提问来源于stack exchange,提问作者Keethan
相关产品推荐
相关产品推荐

