如何在Polars的一次group_by操作中获取每组的head(n)和tail(n)行
一次分组获取每组的头n行和尾n行
针对需求,无需两次分组拼接,有两种高效实现方式:
方法1:窗口函数+筛选(推荐)
通过窗口函数给每个分组内的行标记位置,再筛选出头尾符合条件的行,全程仅需一次遍历:
import polars as pl df = pl.from_repr(""" ┌────────────┬────────┬──────────────────┐ │ date ┆ symbol ┆ ts_dom2secdom_op │ │ --- ┆ --- ┆ --- │ │ date ┆ str ┆ f64 │ ╞════════════╪════════╪══════════════════╡ │ 2000-01-04 ┆ AL ┆ -0.119165 │ │ 2000-01-04 ┆ RU ┆ 0.256691 │ │ 2000-01-05 ┆ AL ┆ -0.126549 │ │ 2000-01-05 ┆ RU ┆ 0.1851 │ │ 2000-01-06 ┆ CU ┆ -0.121354 │ │ 2000-01-06 ┆ RU ┆ 0.228452 │ │ 2000-01-07 ┆ AL ┆ -0.126013 │ │ 2000-01-07 ┆ RU ┆ 0.348729 │ │ 2000-01-10 ┆ AL ┆ -0.139447 │ │ 2000-01-10 ┆ RU ┆ 0.263048 │ └────────────┴────────┴──────────────────┘ """) n = 1 # 可按需调整n值 result = df.with_columns( # 给每个date分组内的行分配从0开始的位置序号 pl.int_range(0, pl.count()).over("date").alias("pos"), # 计算每个分组的总行数 pl.count().over("date").alias("group_size") ).filter( # 筛选前n行 或 后n行 (pl.col("pos") < n) | (pl.col("pos") >= pl.col("group_size") - n) ).drop("pos", "group_size") print(result)
这种方式性能最优,Polars会对窗口函数做优化,避免重复分组计算。
方法2:group_by.agg结合列表拼接
如果偏好使用group_by.agg语法,可将每组的头尾行拼接成列表后展开:
n = 1 result = df.group_by("date").agg( # 拼接每组的head(n)和tail(n)结果,再展平为一维列表 pl.concat_list(pl.all().head(n), pl.all().tail(n)).flatten().alias("combined") ).explode("combined").unnest("combined") print(result)
该方法通过agg内的head/tail获取目标行,再通过concat_list合并、explode+unnest还原为原表结构,代码更紧凑。
内容的提问来源于stack exchange,提问作者Hakase
相关产品推荐
相关产品推荐

