基于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_id | event | event_chain | 取值原因 |
|---|---|---|---|
| 1 | False | 0 | 尚未发生任何事件 |
| 1 | True | 0 | 当前行前4行无事件(不含当前行) |
| 1 | True | 1 | 当前行有事件,且前4行存在1次事件 |
| 1 | False | 1 | 因后续4行有事件,不重置为0 |
| 1 | True | 2 | 当前行有事件且前4行有事件,streak递增 |
| 2 | True | 0 | 无历史事件 |
| 2 | True | 1 | 当前行有事件且用户前4行有事件 |
| 2 | False | 0 | 当前行无事件且后续4行无事件,重置为0 |
| 2 | False | 0 |
原可行代码(待简化)
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)
代码说明
has_prev_event:通过滚动求和判断当前行前4行是否有事件,用于触发streak递增条件has_next_event:反向滚动求和判断当前行后4行是否有事件,用于决定是否需要重置为0chain_start:标记每个事件链的起始位置,通过累加起始点生成唯一的链分组IDchain_count:在每个用户的事件链分组内,累计事件发生的次数event_chain:根据规则调整计数:初始事件(无前序事件)计数为0;无后续事件的非事件行重置为0;其余情况取链内计数减1修正起始点偏移
运行结果将与示例表格完全匹配。
内容的提问来源于stack exchange,提问作者cdkdrf
相关产品推荐
相关产品推荐

