Pandas GroupBy提速神技:40分钟到4秒
那天加完班快十一点,我正准备关电脑回家,老板在群里丢了一句: “那个日报统计为啥跑了快 40 分钟?明天早上九点前给我搞到 5 秒以内。”
我看了一眼自己那段 Pandas 代码,沉默了三秒: ——行吧,这不叫代码,这叫“Python 版人工搬砖”。
中间过程我就不按顺序讲了,想到哪儿说哪儿哈。
先说原始写法有多惨,你们感受下
当时的需求大概是: 几千万条日志,按 user_id 和 day 分组,算各种聚合指标:次数、金额、最大最小值之类的。
我那会儿刚学完 groupby,膨胀得不行,写了这么一段屎山:
defcalc_user_day_metrics(df):
result = []
# 这里本来就已经很危险了,外层 for 就离谱
for (uid, day), sub_df in df.groupby(['user_id', 'day']):
row = {
'user_id': uid,
'day': day,
'cnt': len(sub_df),
'amount_sum': sub_df['amount'].sum(),
'amount_max': sub_df['amount'].max(),
}
# 关键还在这儿各种 if/for 里乱搞
high = sub_df[sub_df['amount'] > 100]
row['big_order_cnt'] = len(high)
result.append(row)
return pd.DataFrame(result)
跑小样本几万行,飞快。 一上生产,几千万行,直接飙到 40 分钟。 监控一看,CPU 占用死高,单核拉满,内存还狂抖,我还以为是服务器问题。
后来想想,这段代码问题巨多:
groupby 之后又在 Python 层 for 循环,等于把 C 写好的高铁扔了,自己骑自行车围着地球转 里面各种筛选、len、sum,都在重复扫那一坨 sub_df 还返回一个 list,再转 DataFrame,多一次内存拷贝
说白了:你让 Pandas 当 Excel,用 Python 当人工运算员,不慢就怪了。
真正的提速,第一刀:把 for 干掉
当晚我一边啃冷掉的炸鸡,一边对着代码发呆,突然意识到: 我算的那几个东西,其实都是标准聚合,完全可以一次性在 groupby 里完成。
于是换成这样:
defcalc_user_day_metrics_fast(df):
grouped = df.groupby(['user_id', 'day'], sort=False) base = grouped['amount'].agg(
cnt='size', # 这个比 len(sub_df) 快多了
amount_sum='sum',
amount_max='max'
)
# 大额订单数量,别在 Python 里筛,直接先加一列标记
tmp = df.copy()
tmp['is_big'] = (tmp['amount'] > 100).astype('int8')
big_cnt = tmp.groupby(['user_id', 'day'], sort=False)['is_big'].sum()
base['big_order_cnt'] = big_cnt
return base.reset_index()
你看几个点哈:
size比你在 for 里len(sub_df)快很多把 amount > 100这种条件,提前算成一列 is_big,再 groupby 聚合全程没有 Python 级别的 for,全部丢给底层 C 去跑
这一步改完,从 40 分钟掉到大概 20 秒。 已经能给老板交差了,但我这人就有病:既然能 20 秒,为啥不能 4 秒?
第二刀:把 apply 这种“性能杀手”都干掉
中途我还犯过一次错,想算点更复杂的指标,就又手欠写了个 apply:
defstupid_way(df):
grouped = df.groupby(['user_id', 'day'], sort=False)
defcustom_func(sub):
# 一堆骚操作
return pd.Series({
'cnt': len(sub),
'median': sub['amount'].median(),
'ratio': (sub['amount'] > 50).mean(),
})
return grouped.apply(custom_func).reset_index()
apply 跑起来,整台机器开始“嗡嗡嗡”,耗时直接翻几倍。
apply 最大的问题就是:
内部还是 Python 调用 每个分组都要构造一个新的 DataFrame/Series 对象 你还以为自己写了个“优雅的函数式编程”
后来我干脆把能拆的都拆开:
median 用 grouped['amount'].median()ratio 这种二分类比例,用布尔转 0/1 再 mean()完全不需要 apply。
大概长这样:
defcalc_more_metrics(df):
g = df.groupby(['user_id', 'day'], sort=False)
res = g['amount'].agg(
cnt='size',
median='median',
) ratio = (df['amount'] > 50).astype('int8')
res['gt_50_ratio'] = ratio.groupby([df['user_id'], df['day']]).mean()
return res.reset_index()
就这么拆一下,时间能从 20 秒再减个几秒。 apply 真的是能不用就不用,真要用,最好只在小数据上玩。
第三刀:数据类型和 group key 优化,才是从 10 秒到 4 秒的关键
这里是很多人忽略的地方。 代码已经很“矢量化”了,但为啥还要十几秒?
我当时做了几件小事:
1)把字符串转成分类类型
user_id 本来是字符串(因为有前导 0 那种),直接 groupby,Pandas 每次都要拿字符串哈希。 我手动改了一下:
df['user_id'] = df['user_id'].astype('category')
df['day'] = pd.to_datetime(df['day']).dt.date.astype('category')
这一改,groupby 分组那一步直接快了一截。 category 的好处:
内部是整数编码 groupby 等操作都按整数来搞 内存也省好多
2)sort=False 不是摆设
很多人没注意,groupby 默认是要帮你排好组 key 的。 但我这报表根本不在乎 user_id 的排序,随便就行。
所以我统一加上 sort=False:
g = df.groupby(['user_id', 'day'], sort=False)
这一个参数,几千万行数据,能省掉非常可观的一段排序时间。
3)只保留用得上的列
原始 df 有二三十列,真正统计只用到 3、4 列,我一开始懒得裁剪。 后来手动加了一段预处理:
used_cols = ['user_id', 'day', 'amount']
df_small = df[used_cols].copy()
df_small['user_id'] = df_small['user_id'].astype('category')
df_small['day'] = pd.to_datetime(df_small['day']).dt.date.astype('category')
df_small['amount'] = df_small['amount'].astype('float32')
你别小看这几行,
列少了,groupby 内部搬运数据的量直接变小 float64 换成 float32,内存压下去,CPU cache 也更友好
以上这些加起来,才是真正从十几秒下到 4 秒的关键。
顺手说个小坑,很多人 groupby 后又 merge 回去,白白多一次 IO
我中间还踩过一个傻坑:
stats = calc_user_day_metrics_fast(df)# 再 merge 回原表,搞什么标记之类
df2 = df.merge(stats, on=['user_id', 'day'], how='left')
看着很合理对吧,其实非常费劲:
merge 又把整个大表拷贝了一遍 键列如果没建好索引,等着被爆内存就行
后来直接用 transform 搞定:
g = df.groupby(['user_id', 'day'], sort=False)
df['day_amount_sum'] = g['amount'].transform('sum')
df['day_cnt'] = g['amount'].transform('size')
这样所有指标直接写回原表,不用 merge, 而且 transform 也是在 C 层跑,比你手动 merge 快太多。
说点八卦,性能这事儿在别的地方也是一个路子
比如我之前写数据库那篇压测,Postgres 比 MySQL 在高并发场景里,吞吐量和耗时都能差一大截,很大一部分也是因为底层实现和访问模式的差异,跟咱这个“别在 Python 层乱 for,尽量让底层干活”是一样的路子
你会发现各种系统优化,最后都绕不开那几件事:
少做无意义的重复工作 把逻辑下沉到更靠近数据、离底层更近的地方 让 CPU 干的事情集中、顺滑,不要又算又等 IO
最后我把整个报表流程梳了一下,大概就这几个阶段:
1)进来先瘦身:
只保留要用的列 转合适的 dtype(category / float32 / int32)
2)预计算好各种布尔条件列(is_big、is_valid 之类)
3)所有能 agg 的就 agg,能 transform 的就 transform
坚决杜绝大规模 apply 和 Python for
4)中间结果少 merge,多用索引、transform
跑完一看监控图:
CPU 一阵猛冲,4 秒结束, 内存曲线平滑,磁盘 IO 也不至于炸
老板第二天早上随手点开报表,愣了一下: “诶?咋一下就出来了?”
我: “啊…昨天那个…就是那个…算发调优了一点点。”
他还在那儿感慨性能飞跃,我这边已经在想下一个坑了: ——要不要把这套东西改成按小时分区再跑一遍看看…
算了不说了,我去热个馒头,饿死了。