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

如何在Polars中逐行应用掩码并计算每行连续1的最大个数?

如何在Polars中逐行应用掩码并计算每行连续1的最大个数?

我明白你现在需要用Polars实现和Pandas一样的功能——逐行找出连续1的最大个数。先回顾一下你给出的示例DataFrame:

import polars as pl

df = pl.from_repr("""
┌─────┬─────┬─────┬─────┬─────┐
│ 0   ┆ 1   ┆ 2   ┆ 3   ┆ 4   │
│ --- ┆ --- ┆ --- ┆ --- ┆ --- │
│ i64 ┆ i64 ┆ i64 ┆ i64 ┆ i64 │
╞═════╪═════╪═════╪═════╪═════╡
│ 0   ┆ 0   ┆ 0   ┆ 0   ┆ 0   │
│ 1   ┆ 0   ┆ 0   ┆ 1   ┆ 1   │
│ 0   ┆ 0   ┆ 1   ┆ 0   ┆ 1   │
│ 1   ┆ 0   ┆ 0   ┆ 0   ┆ 0   │
│ 1   ┆ 0   ┆ 1   ┆ 1   ┆ 1   │
│ 0   ┆ 0   ┆ 0   ┆ 0   ┆ 0   │
│ 1   ┆ 1   ┆ 1   ┆ 0   ┆ 0   │
│ 0   ┆ 0   ┆ 0   ┆ 0   ┆ 0   │
│ 1   ┆ 1   ┆ 0   ┆ 1   ┆ 1   │
│ 1   ┆ 1   ┆ 1   ┆ 0   ┆ 1   │
└─────┴─────┴─────┴─────┴─────┘
""")

你已经有了Pandas的实现方案:

在Pandas中你是这样实现的:

m = df.to_pandas().eq(1)
m.cumsum(axis=1).sub(m.cumsum(axis=1).mask(m).ffill(axis=1).fillna(0)).max(axis=1)

输出结果:

0    0.0
1    2.0
2    1.0
3    1.0
4    3.0
5    0.0
6    3.0
7    0.0
8    2.0
9    3.0
dtype: float64

下面我给你两种Polars的实现方案,分别适合不同的场景:

方案一:逐行分组统计(适合小数据集)

这个方案用itertools.groupby对每行的元素进行分组,直接统计连续1的最长长度,逻辑直观易懂:

import itertools

result = df.select(
    pl.row(pl.all()).map_elements(
        lambda row: max(
            sum(1 for _ in group) if val == 1 else 0
            for val, group in itertools.groupby(row)
        ),
        return_dtype=pl.Int64
    ).alias("max_consecutive_1s")
)

print(result)

输出结果:

shape: (10, 1)
┌───────────────────┐
│ max_consecutive_1s│
│ ---               │
│ i64               │
╞═══════════════════╡
│ 0                 │
│ 2                 │
│ 1                 │
│ 1                 │
│ 3                 │
│ 0                 │
│ 3                 │
│ 0                 │
│ 2                 │
│ 3                 │
└───────────────────┘

方案二:向量化实现(高效,适合大数据集)

这个方案和你Pandas的逻辑完全对齐,用Polars的向量化操作实现,性能远高于逐行循环,适合处理大规模数据:

# 1. 创建等于1的掩码
mask = df.select(pl.all().eq(1))
# 2. 逐行计算累计和
cumsum = mask.select(pl.all().cumsum(axis=1))
# 3. 对掩码为0的位置填充前向的累计和,无前置值则填0
masked_cumsum = cumsum.select(
    pl.all().where(mask).forward_fill(axis=1).fill_null(0)
)
# 4. 计算连续1的计数并取每行最大值
max_consec = (cumsum - masked_cumsum).select(
    pl.max_horizontal(pl.all()).alias("max_consecutive_1s")
)

print(max_consec)

这个方案的输出和方案一完全一致,核心逻辑和你Pandas的代码对应:

  • 先标记所有1的位置
  • 计算累计和来追踪连续序列的起点
  • 通过前向填充和差值计算,得到每个连续1序列的长度
  • 最后取每行的最大值就是我们要的结果

备注:内容来源于stack exchange,提问作者rhug123

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.04.13 18:13:02