Python技术迷

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 也不至于炸

老板第二天早上随手点开报表,愣了一下: “诶?咋一下就出来了?”

我: “啊…昨天那个…就是那个…算发调优了一点点。”

他还在那儿感慨性能飞跃,我这边已经在想下一个坑了: ——要不要把这套东西改成按小时分区再跑一遍看看…

算了不说了,我去热个馒头,饿死了。