Polars按组提取前N个元素:head报错原因及优化方案
问题描述
构造的Polars DataFrame代码
import numpy as np import polars as pl pl.Config(tbl_rows=20) # 显示完整输出 df = (pl .DataFrame(dict( j=np.random.randint(10, 99, 20), )) .with_row_index() .select( g=pl.col('index') // 4, j='j', ) )
数据结构
shape: (20, 2) ┌─────┬─────┐ │ g ┆ j │ │ --- ┆ --- │ │ u32 ┆ i64 │ ╞═════╪═════╡ │ 0 ┆ 95 │ │ 0 ┆ 80 │ │ 0 ┆ 51 │ │ 0 ┆ 68 │ │ 1 ┆ 71 │ │ 1 ┆ 92 │ │ 1 ┆ 44 │ │ 1 ┆ 97 │ │ 2 ┆ 36 │ │ 2 ┆ 64 │ │ 2 ┆ 70 │ │ 2 ┆ 80 │ │ 3 ┆ 75 │ │ 3 ┆ 69 │ │ 3 ┆ 54 │ │ 3 ┆ 16 │ │ 4 ┆ 88 │ │ 4 ┆ 89 │ │ 4 ┆ 97 │ │ 4 ┆ 37 │ └─────┴─────┘
需求
提取每个g分组中的前2个元素,目标结果如下:
shape: (10, 2) ┌─────┬─────┐ │ g ┆ j │ │ --- ┆ --- │ │ u32 ┆ i64 │ ╞═════╪═════╡ │ 0 ┆ 95 │ │ 0 ┆ 80 │ │ 1 ┆ 71 │ │ 1 ┆ 92 │ │ 2 ┆ 36 │ │ 2 ┆ 64 │ │ 3 ┆ 75 │ │ 3 ┆ 69 │ │ 4 ┆ 88 │ │ 4 ┆ 89 │ └─────┴─────┘
尝试的代码及报错
尝试使用head(2).over('g')实现:
dfj = (df .select( pl.all().head(2).over('g') ) ) print(dfj)
抛出异常:
ComputeError: the length of the window expression did not match that of the group Error originated in expression: 'col("g").slice(offset=0, length=2).over([col("g")])'
想了解两个问题:
- 为什么
head结合over会报错 - 是否有更好的解决方案(尤其无需使用
group_by)
解答
1. 报错原因
Polars的窗口函数(over)核心逻辑是给原数据的每一行生成对应的值,因此要求每个分组返回的结果长度必须和原分组的行数完全一致。
当你对pl.all()调用head(2)时,每个分组仅返回2条数据,但原分组有4行,导致窗口函数返回的结果长度(2)和分组行数(4)不匹配,无法对应到原分组的每一行,因此抛出长度不匹配的错误。
2. 无需group_by的最优解决方案
推荐使用filter结合窗口函数row_number()实现,代码简洁高效:
result = df.filter(pl.row_number().over('g') <= 2) print(result)
原理
row_number().over('g')会给每个g分组内的行按原始顺序标记序号(从1开始),过滤出序号≤2的行,就能精准得到每个分组的前2条数据,完全符合需求。
另外还有一种写法,用slice+over配合flatten,但可读性稍差:
result = (df .select(pl.struct(['g', 'j']).head(2).over('g')) .flatten() .unnest('struct') )
内容的提问来源于stack exchange,提问作者levant pied
相关产品推荐
相关产品推荐

