如何加速分组后重复提取每组前n行的操作?
如何加速分组后重复提取每组前n行的操作?
嘿,这个场景我太熟了!1亿行数据+近30个分组列,重复跑分组取前N行确实会白耗不少算力,咱们来一步步把速度提上去~
首先得先指出你原代码里的一个小bug:你创建DataFrame时用的schema是[str(col_i) for col_i in range(n_cols)],所以列名是"0"、"1"..."29",但后面group_by用的是"col_1"这类名称,Polars会找不到这些列直接报错,得先把列名统一好——要么创建时用f"col_{i}"当列名,要么group_by时用str(i)引用列。
接下来看核心优化点:你现在的代码循环3次,每次都要执行一次group_by.head(),相当于把分组、取前N行的重操作重复了3次,这在1亿行数据上是极大的浪费。咱们换个思路:既然要取的是每组前1、2、3行,那只需要一次性获取每组前3行(也就是所有需求里的最大N值),之后再基于这个结果拆分出不同N的子集就行,这样只需要做一次核心的分组计算,剩下的都是轻量筛选,能省掉大量时间。
优化后的代码是这样的:
import polars as pl import numpy as np rows = 100000000 n_cols = 30 # 先统一列名,避免报错 df = pl.DataFrame( np.random.randint(0, 100, size=(n_cols, rows)), schema=[f"col_{i}" for i in range(n_cols)] ) First_n_rows_list = [1,2,3] max_n = max(First_n_rows_list) # 一次性获取每组前max_n行,只做一次分组+head操作 df_sorted = df.sort("col_0") grouped_top_max = df_sorted.group_by([f"col_{i}" for i in range(1, n_cols)]).head(max_n) # 给每个分组内的行加个序号,方便后续筛选前n行 grouped_top_max = grouped_top_max.with_row_index(group=True, name="group_row_num") # 针对每个n生成对应的结果 result_dfs = [] for n in First_n_rows_list: # 筛选分组内序号小于n的行(序号从0开始) subset = grouped_top_max.filter(pl.col("group_row_num") < n) subset = subset.with_columns(pl.lit(n).alias('First_n_rows').cast(pl.Int8)) result_dfs.append(subset) # 合并结果并去掉临时的序号列 result = pl.concat(result_dfs).drop("group_row_num")
为什么这样更快?
- 把3次重复的分组+head操作压缩成1次,这是最关键的优化——分组本身是重计算,尤其是近30个分组列的场景,重复执行完全是做无用功。
- 用
with_row_index(group=True)给每个分组内的行加序号,后续筛选前n行就是简单的数值比较,几乎没开销。 - 最终结果和你原代码的输出结构完全一致,没有任何额外的字段残留。
额外的可选优化
如果全局排序1亿行的开销实在太大,你也可以跳过全局排序,改成分组内排序后取前max_n行,避免全量数据排序的成本:
# 分组内对col_0排序,取前max_n行 grouped_top_max = df.group_by([f"col_{i}" for i in range(1, n_cols)]).agg( pl.all().sort_by("col_0").head(max_n) ).explode(pl.all().exclude([f"col_{i}" for i in range(1, n_cols)]))
这个思路适合分组基数很大的场景,具体选哪种可以根据你的实际数据测试一下性能差异。
备注:内容来源于stack exchange,提问作者user28199045
相关产品推荐
相关产品推荐

