Polars优化:如何对同一分组仅执行一次分组操作?
Polars优化分组移位操作:确保分组仅执行一次
先看K=2的示例,这个问题在分组键g基数较高且K远大于1的场景下,性能影响会非常显著:
import polars as pl df = pl.DataFrame(dict( g=[1, 2, 1, 2, 1, 2], v=[1, 2, 3, 4, 5, 6], )) K = 2 # 原写法:每个shift(k).over('g')都会单独触发一次分组 df.with_columns((pl.col("v").shift(k+1).over('g').alias(f's{k}') for k in range(K)))
执行结果:
╭─────┬─────┬──────┬──────╮ │ g ┆ v ┆ s0 ┆ s1 │ │ i64 ┆ i64 ┆ i64 ┆ i64 │ ╞═════╪═════╪══════╪══════╡ │ 1 ┆ 1 ┆ null ┆ null │ │ 2 ┆ 2 ┆ null ┆ null │ │ 1 ┆ 3 ┆ 1 ┆ null │ │ 2 ┆ 4 ┆ 2 ┆ null │ │ 1 ┆ 5 ┆ 3 ┆ 1 │ │ 2 ┆ 6 ┆ 4 ┆ 2 │ ╰─────┴─────┴──────┴──────╯
问题
原写法中,Polars会为每个over('g')子句单独执行分组操作,当K很大时会产生大量重复分组,导致性能下降。如何让分组仅执行一次,达到与以下group_by.agg写法相当的执行速度?
df.group_by('g').agg((pl.col("v").shift(k+1).alias(f's{k}') for k in range(K)))
解决方案
方法1:使用group_by.map_groups
通过map_groups对每个分组仅处理一次,在分组内部生成所有需要的移位列,全程只执行一次分组操作:
def add_shift_columns(group): return group.with_columns( (pl.col("v").shift(k+1).alias(f"s{k}") for k in range(K)) ) # maintain_order=True 保证结果顺序与原表一致 result = df.group_by("g", maintain_order=True).map_groups(add_shift_columns)
方法2:先聚合再展开
先通过group_by.agg一次性生成每个分组的所有移位结果(以列表形式存储),再通过explode展开列表,还原为原表的行结构:
# 聚合生成每个分组的v及所有移位列的列表 agg_df = df.group_by("g", maintain_order=True).agg( pl.col("v"), *(pl.col("v").shift(k+1).alias(f"s{k}") for k in range(K)) ) # 展开所有列表列,得到与原写法一致的结果 result = agg_df.explode(["v"] + [f"s{k}" for k in range(K)])
这两种方法都能确保分组仅执行一次,性能与group_by.agg写法相当,完美解决高基数分组+大K场景下的性能问题。
内容的提问来源于stack exchange,提问作者Cingonius Varro
相关产品推荐
相关产品推荐

