如何利用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
相关产品推荐
相关产品推荐

