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

如何利用Polars的map_batches优化连续非零值最大计数代码?

优化Polars计算连续非零值最大计数的方案

需求说明

每列对应一个地点,每行对应一个时间戳,需要计算各列中连续非零值的最大计数(可转换为布尔值判断是否为零)。现有代码可实现需求但效率偏低,希望通过map_batches()优化。

模拟数据

import polars as pl

pivoted_df = pl.from_repr("""
┌─────────────────────┬────────────┬────────────┐
│ Date                ┆ Location 1 ┆ Location 2 │
│ ---                 ┆ ---        ┆ ---        │
│ datetime[ns]        ┆ i64        ┆ i64        │
╞═════════════════════╪════════════╪════════════╡
│ 2023-01-01 00:00:00 ┆ 0          ┆ 1          │
│ 2023-01-01 01:00:00 ┆ 1          ┆ 1          │
│ 2023-01-01 02:00:00 ┆ 1          ┆ 1          │
│ 2023-01-01 03:00:00 ┆ 0          ┆ 1          │
│ 2023-01-01 04:00:00 ┆ 1          ┆ 1          │
│ 2023-01-01 05:00:00 ┆ 1          ┆ 0          │
│ 2023-01-01 06:00:00 ┆ 1          ┆ 0          │
└─────────────────────┴────────────┴────────────┘
""")

预期输出

┌────────────┬───────┐
│ Location   ┆ Value │
│ ---        ┆ ---   │
│ str        ┆ i32   │
╞════════════╪═══════╡
│ Location 1 ┆ 3     │
│ Location 2 ┆ 5     │
└────────────┴───────┘

现有代码问题

原代码通过循环逐列处理,每次都要对整个DataFrame执行选择、列计算等操作,当列数较多时,重复的DataFrame操作会导致效率低下:

for col in pivoted_df.drop("Date").columns:
    xy_cont_df_a = (
        pivoted_df.select(pl.col(col))
        .with_columns(
            pl.when(
                pl.col(col).cast(pl.Boolean)
                & pl.col(col)
                .cast(pl.Boolean)
                .shift(-1, fill_value=False)
                .not_()
            ).then(
                pl.count().over(
                    (
                        pl.col(col).cast(pl.Boolean)
                        != pl.col(col).cast(pl.Boolean).shift()
                    ).cum_sum()
                )
            )
        )
        .max()
    )

优化方案:使用map_batches()

通过map_batches()批量处理所有列,结合自定义函数实现向量化计算,避免循环带来的冗余操作:

步骤1:定义计算连续非零最大长度的函数

def max_consecutive_non_zero(s: pl.Series) -> int:
    # 将列转换为布尔值(非零为True)
    bool_series = s != 0
    # 生成分组标识:当前值与前一个值不同时,分组ID递增
    group_ids = (bool_series != bool_series.shift()).cum_sum()
    # 仅统计非零分组的长度,取最大值(若全为零则返回0)
    non_zero_lengths = bool_series.group_by(group_ids).len().filter(bool_series)
    return non_zero_lengths.max() if not non_zero_lengths.is_empty() else 0

步骤2:用map_batches()批量处理列并重塑结果

result = (
    pivoted_df.drop("Date")
    # 批量对所有列应用自定义函数
    .map_batches(lambda df: df.select(pl.all().map(max_consecutive_non_zero)))
    # 将宽表转为长表,匹配预期输出格式
    .melt(variable_name="Location", value_name="Value")
)

print(result)

优化原理

  • map_batches()利用Polars的并行处理能力,一次性处理所有列,避免了循环中重复的DataFrame初始化与操作开销。
  • 自定义函数内使用Polars原生的向量化方法(shift()、cum_sum()、group_by()),比逐行计算效率更高。

验证结果

运行优化后的代码,输出与预期完全一致:

┌────────────┬───────┐
│ Location   ┆ Value │
│ ---        ┆ ---   │
│ str        ┆ i64   │
╞════════════╪═══════╡
│ Location 1 ┆ 3     │
│ Location 2 ┆ 5     │
└────────────┴───────┘

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.26 08:44:53