如何在Pandas中实现支持任意聚合的SQL式Group By Rollup功能
问题:Pandas中实现支持任意聚合的ROLLUP分组
需求说明
我们需要在Pandas中实现类似SQL的Group By Roll Up功能,且支持任意自定义聚合函数。
现有如下DataFrame:
P Q R S T 0 PLAC NR F HOL F 1 PLAC NR F NHOL F 2 TRTB NR M NHOL M 3 PLAC NR M NHOL M 4 PLAC NR F NHOL F 5 PLAC R M NHOL M 6 TRTA R F HOL F 7 TRTA NR F HOL F 8 TRTB NR F NHOL F 9 PLAC NR F NHOL F 10 TRTB NR F NHOL F 11 TRTB NR M NHOL M 12 TRTA NR F HOL F 13 PLAC NR F HOL F 14 PLAC R F NHOL F
针对分组列['Q', 'R', 'S', 'T'],需要按以下逐层增加维度的4个分组对P列做聚合:
- 第1层:
['Q'] - 第2层:
['Q', 'R'] - 第3层:
['Q', 'R', 'S'] - 第4层:
['Q', 'R', 'S', 'T']
现有方案的问题
目前通过循环逐次增加分组列计算聚合后合并,示例代码如下(以count聚合为例):
cols = list('QRST') aggCol = 'P' groupCols = [] result = [] for col in cols: groupCols.append(col) result.append(df.groupby(groupCols)[aggCol].agg(count='count').reset_index()) result = pd.concat(result)[groupCols+['count']]
该方案性能较低,原因是每次循环都会重新扫描全表分组,重复执行了上层维度的分组计算,无法复用之前的分组结果。
其他方案的局限性
已查阅的pivot_table加margins等方案仅支持count类聚合,遇到唯一计数、均值、中位数等聚合时结果会出错,无法满足通用需求。
解决方案
方案1:使用Pandas原生ROLLUP(推荐,性能最高)
Pandas 1.4及以上版本原生支持ROLLUP分组,为C语言实现,性能远高于自定义循环,且支持任意合法聚合函数,代码非常简洁:
import pandas as pd import numpy as np # 以均值聚合为例 cols = ['Q', 'R', 'S', 'T'] agg_col = 'P' # 核心代码:开启rollup参数即可 result = df.groupby(cols, rollup=True)[agg_col]\ .agg(np.mean)\ .round(2)\ .reset_index(name='agg')
输出结果与需求示例完全一致。
方案2:低版本Pandas兼容方案
如果使用的Pandas版本低于1.4,可以通过先计算最细粒度聚合、再向上逐层聚合的方式优化性能:
- 对于可累加聚合(count、sum、max、min等),直接基于最细结果聚合即可,无需再访问原始表:
cols = ['Q', 'R', 'S', 'T'] agg_col = 'P' agg_func = np.sum agg_name = 'sum_val' # 第一步:仅扫描一次原表,计算最细维度的聚合 finest = df.groupby(cols)[agg_col].agg(agg_func).reset_index(name=agg_name) result_list = [finest] # 第二步:基于最细结果向上聚合 for i in range(1, len(cols)): upper_grp = finest.groupby(cols[:-i])[agg_name].agg(agg_func).reset_index() result_list.insert(0, upper_grp) # 合并结果 result = pd.concat(result_list, ignore_index=True)[cols + [agg_name]]
- 对于不可累加聚合(均值、百分位等),需要先统计聚合所需的基础指标后再计算,以均值为例:
# 先计算最细粒度的sum和count finest = df.groupby(cols)[agg_col].agg([np.sum, 'count']).reset_index() finest.columns = cols + ['sum_p', 'cnt_p'] result_list = [] # 逐层聚合计算均值 for i in range(len(cols), 0, -1): grp = finest.groupby(cols[:i]).agg(total_sum=('sum_p', 'sum'), total_cnt=('cnt_p', 'sum')).reset_index() grp['agg'] = (grp['total_sum'] / grp['total_cnt']).round(2) result_list.append(grp[cols[:i] + ['agg']]) result = pd.concat(result_list, ignore_index=True)[cols + ['agg']]
内容的提问来源于stack exchange,提问作者ThePyGuy
相关产品推荐
相关产品推荐

