You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何在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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.18 20:14:51