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

基于Polars DataFrame计算窗口化事件链的优化方案问询

简化Polars中事件连续streak统计的实现

需求说明

给定如下Polars DataFrame:

import polars as pl

data = pl.DataFrame({
    "user_id": [1, 1, 1, 1, 1, 2, 2, 2, 2], 
    "event": [False, True, True, False, True, True, True, False, False]
})

需要计算event_chain列,用于统计用户的事件连续streak,规则为:

  • 当前事件发生且前4行存在事件时,streak计数器递增
  • 若后续4行无事件则重置为0

规则示例

user_ideventevent_chain取值原因
1False0尚未发生任何事件
1True0当前行前4行无事件(不含当前行)
1True1当前行有事件,且前4行存在1次事件
1False1因后续4行有事件,不重置为0
1True2当前行有事件且前4行有事件,streak递增
2True0无历史事件
2True1当前行有事件且用户前4行有事件
2False0当前行无事件且后续4行无事件,重置为0
2False0

原可行代码(待简化)

data.with_columns(
    rows_since_last_event=pl.int_range(pl.len()).over("user_id")
    - pl.when("event").then(pl.int_range(pl.len())).forward_fill()
    .over("user_id"),
    rows_till_next_event=pl.when("event").then(pl.int_range(pl.len()))
    .backward_fill().over("user_id") - pl.int_range(pl.len()).over("user_id")  # 修正原代码笔误:athlete_id改为user_id
)
.with_columns(
    chain_event=pl.when(
        pl.col("event")
        .fill_null(0)
        .rolling_sum(window_size=4, min_periods=1)
        .over("user_id")
        - pl.col("event").fill_null(0)
        > 0
    )
    .then(1)
    .otherwise(0)
)
.with_columns(
    chain_event_change=pl.when(
        pl.col("chain_event").eq(1),
        pl.col("chain_event").shift().eq(0),
        pl.col("rows_since_last_event").fill_null(5) > 3,
    )
    .then(1)
    .when(
        pl.col("chain_event").eq(0),  # 修正原代码笔误:congested_event改为chain_event
        pl.col("chain_event").shift().eq(1),
        pl.col("rows_till_next_event").fill_null(5) > 3,
    )
    .then(1)
    .otherwise(0)
)
.with_columns(
    chain_event_identifier=pl.col("chain_event_change")
    .cum_sum()
    .over("user_id")
)
.with_columns(
    event_chain=pl.col("chain_event")
    .cum_sum()
    .over("user_id", "chain_event_identifier")
)

简化实现方案

通过窗口函数+状态标记的方式可大幅简化代码,核心思路是先标记事件链分组,再在分组内完成计数与规则适配:

result = (
    data
    .with_row_index(name="idx")
    .over("user_id")
    .with_columns(
        # 判断当前行前4行是否存在事件
        has_prev_event=pl.col("event").rolling_sum(window_size=4, min_periods=1).shift(1).fill_null(0) > 0,
        # 判断当前行后4行是否存在事件
        has_next_event=pl.col("event").rolling_sum(window_size=4, min_periods=1).shift(-1).fill_null(0) > 0,
        # 标记事件链起始点,生成链分组ID
        chain_start=pl.when(pl.col("event") & ~pl.col("has_prev_event")).then(1).otherwise(0).cum_sum()
    )
    .with_columns(
        # 每个链内的事件累计计数
        chain_count=pl.col("event").cum_sum().over("user_id", "chain_start")
    )
    .with_columns(
        # 按规则生成最终event_chain
        event_chain=pl.when(
            ~pl.col("event") & ~pl.col("has_next_event")
        ).then(0).otherwise(
            pl.when(pl.col("event") & ~pl.col("has_prev_event")).then(0).otherwise(pl.col("chain_count") - 1)
        )
    )
    .drop("idx", "has_prev_event", "has_next_event", "chain_start", "chain_count")
)

print(result)

代码说明

  1. has_prev_event:通过滚动求和判断当前行前4行是否有事件,用于触发streak递增条件
  2. has_next_event:反向滚动求和判断当前行后4行是否有事件,用于决定是否需要重置为0
  3. chain_start:标记每个事件链的起始位置,通过累加起始点生成唯一的链分组ID
  4. chain_count:在每个用户的事件链分组内,累计事件发生的次数
  5. event_chain:根据规则调整计数:初始事件(无前序事件)计数为0;无后续事件的非事件行重置为0;其余情况取链内计数减1修正起始点偏移

运行结果将与示例表格完全匹配。

内容的提问来源于stack exchange,提问作者cdkdrf

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.22 14:37:02