如何在DuckDB中复现Polars分组rolling_mean(min_periods=2)效果?
在DuckDB中复现Polars滚动均值(min_periods=2)的效果
数据集
使用如下测试数据集:
data = { 'id': ['a', 'a', 'a', 'b', 'b', 'b', 'b'], 'd': [1,2,3,0,1,2,3], 'sales': [5,1,3,4,1,2,3], }
需求:按id分组,添加一列窗口大小为2、min_periods=2的滚动均值列(窗口至少包含2个数据时才计算均值,否则返回null)。
Polars实现方案
已通过Polars完成需求,代码及结果如下:
import polars as pl df = pl.DataFrame(data) df.with_columns(sales_rolling = pl.col('sales').rolling_mean(2).over('id'))
执行结果:
shape: (7, 4) ┌─────┬─────┬───────┬───────────────┐ │ id ┆ d ┆ sales ┆ sales_rolling │ │ --- ┆ --- ┆ --- ┆ --- │ │ str ┆ i64 ┆ i64 ┆ f64 │ ╞═════╪═════╪═══════╪═══════════════╡ │ a ┆ 1 ┆ 5 ┆ null │ │ a ┆ 2 ┆ 1 ┆ 3.0 │ │ a ┆ 3 ┆ 3 ┆ 2.0 │ │ b ┆ 0 ┆ 4 ┆ null │ │ b ┆ 1 ┆ 1 ┆ 2.5 │ │ b ┆ 2 ┆ 2 ┆ 1.5 │ │ b ┆ 3 ┆ 3 ┆ 2.5 │ └─────┴─────┴───────┴───────────────┘
DuckDB初始尝试问题
用DuckDB实现时,初始代码返回的结果不符合min_periods=2的要求:窗口仅含单个值时仍计算了均值,而非返回null。
初始代码:
import duckdb duckdb.sql(""" select *, mean(sales) over ( partition by id order by d range between 1 preceding and 0 following ) as sales_rolling from df """).sort('id', 'd')
初始结果:
┌─────────┬───────┬───────┬───────────────┐ │ id │ d │ sales │ sales_rolling │ │ varchar │ int64 │ int64 │ double │ ├─────────┼───────┼───────┼───────────────┤ │ a │ 1 │ 5 │ 5.0 │ │ a │ 2 │ 1 │ 3.0 │ │ a │ 3 │ 3 │ 2.0 │ │ b │ 0 │ 4 │ 4.0 │ │ b │ 1 │ 1 │ 2.5 │ │ b │ 2 │ 2 │ 1.5 │ │ b │ 3 │ 3 │ 2.5 │ └─────────┴───────┴───────┴───────────────┘
DuckDB正确实现方案
要复现Polars中min_periods=2的效果,可结合COUNT()窗口函数判断窗口内的行数:只有当窗口内的行数≥2时,才返回均值,否则返回null。
方案一:直接嵌套窗口函数
import duckdb duckdb.sql(""" select *, CASE WHEN count(sales) over ( partition by id order by d range between 1 preceding and 0 following ) >= 2 THEN mean(sales) over ( partition by id order by d range between 1 preceding and 0 following ) ELSE NULL END as sales_rolling from df """).sort('id', 'd')
方案二:使用WINDOW子句简化
为避免重复定义窗口,可用WINDOW子句复用窗口逻辑:
import duckdb duckdb.sql(""" select *, CASE WHEN count(sales) over w >= 2 THEN mean(sales) over w ELSE NULL END as sales_rolling from df WINDOW w AS ( partition by id order by d range between 1 preceding and 0 following ) """).sort('id', 'd')
执行上述代码后,结果将与Polars完全一致:
┌─────────┬───────┬───────┬───────────────┐ │ id │ d │ sales │ sales_rolling │ │ varchar │ int64 │ int64 │ double │ ├─────────┼───────┼───────┼───────────────┤ │ a │ 1 │ 5 │ NULL │ │ a │ 2 │ 1 │ 3.0 │ │ a │ 3 │ 3 │ 2.0 │ │ b │ 0 │ 4 │ NULL │ │ b │ 1 │ 1 │ 2.5 │ │ b │ 2 │ 2 │ 1.5 │ │ b │ 3 │ 3 │ 2.5 │ └─────────┴───────┴───────┴───────────────┘
内容的提问来源于stack exchange,提问作者ignoring_gravity
相关产品推荐
相关产品推荐

