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

PySpark实现带重置规则的窗口计数问题求助

PySpark实现带重置的计数列需求

原始数据

我有一个包含多国家数据的PySpark DataFrame,结构及数据如下:

df = spark.createDataFrame(
    data=[
    (1, "GERMANY", "20230606", True),
    (2, "GERMANY", "20230620", False),
    (3, "GERMANY", "20230627", True),
    (4, "GERMANY", "20230705", True),
    (5, "GERMANY", "20230714", False),
    (6, "GERMANY", "20230715", True),
    ],
    schema=["ID", "COUNTRY", "DATE", "FLAG"]
)
df.show()

展示结果:

+---+-------+--------+-----+
| ID|COUNTRY|    DATE| FLAG|
+---+-------+--------+-----+
|  1|GERMANY|20230606| true|
|  2|GERMANY|20230620|false|
|  3|GERMANY|20230627| true|
|  4|GERMANY|20230705| true|
|  5|GERMANY|20230714|false|
|  6|GERMANY|20230715| true|
+---+-------+--------+-----+

需求说明

需要新增一列COUNT_WITH_RESET,规则如下:

  • 当FLAG=False时,COUNT_WITH_RESET=0;
  • 当FLAG=True时,COUNT_WITH_RESET统计该国家从上一个FLAG=False的日期开始的行数。

预期输出:

+---+-------+--------+-----+----------------+
| ID|COUNTRY|    DATE| FLAG|COUNT_WITH_RESET|
+---+-------+--------+-----+----------------+
|  1|GERMANY|20230606| true|               1|
|  2|GERMANY|20230620|false|               0|
|  3|GERMANY|20230627| true|               1|
|  4|GERMANY|20230705| true|               2|
|  5|GERMANY|20230714|false|               0|
|  6|GERMANY|20230715| true|               1|
+---+-------+--------+-----+----------------+

尝试的代码及问题

我尝试用row_number()结合窗口函数,但无法实现计数重置,代码如下:

from pyspark.sql.window import Window
import pyspark.sql.functions as F

window_reset = Window.partitionBy("COUNTRY").orderBy("DATE")

df_with_reset = (
    df
    .withColumn("COUNT_WITH_RESET", F.when(~F.col("FLAG"), 0)
                .otherwise(F.row_number().over(window_reset)))
)

df_with_reset.show()

得到错误结果:

+---+-------+--------+-----+----------------+
| ID|COUNTRY|    DATE| FLAG|COUNT_WITH_RESET|
+---+-------+--------+-----+----------------+
|  1|GERMANY|20230606| true|               1|
|  2|GERMANY|20230620|false|               0|
|  3|GERMANY|20230627| true|               3|
|  4|GERMANY|20230705| true|               4|
|  5|GERMANY|20230714|false|               0|
|  6|GERMANY|20230715| true|               6|
+---+-------+--------+-----+----------------+

仅按国家分区的窗口不符合需求,请问是否思路正确?PySpark是否有内置函数实现该需求?是否需要使用UDF?


解决方案

不需要使用UDF,通过构建分组标识结合窗口函数即可实现。核心思路是:先为每个FLAG=False之后的行组生成唯一标识,再基于这个标识+国家进行分区,最后用row_number()实现组内计数。

具体代码实现

from pyspark.sql.window import Window
import pyspark.sql.functions as F

# 第一步:生成分组标识,每次遇到FLAG=False时分组ID递增
window_group = Window.partitionBy("COUNTRY").orderBy("DATE")
df_with_group = df.withColumn(
    "group_id",
    F.sum(F.when(~F.col("FLAG"), 1).otherwise(0)).over(window_group)
)

# 第二步:基于国家和分组ID构建窗口,计算组内行号
window_count = Window.partitionBy("COUNTRY", "group_id").orderBy("DATE")
df_result = df_with_group.withColumn(
    "COUNT_WITH_RESET",
    F.when(~F.col("FLAG"), 0).otherwise(F.row_number().over(window_count))
).drop("group_id")

df_result.show()

输出结果

+---+-------+--------+-----+----------------+
| ID|COUNTRY|    DATE| FLAG|COUNT_WITH_RESET|
+---+-------+--------+-----+----------------+
|  1|GERMANY|20230606| true|               1|
|  2|GERMANY|20230620|false|               0|
|  3|GERMANY|20230627| true|               1|
|  4|GERMANY|20230705| true|               2|
|  5|GERMANY|20230714|false|               0|
|  6|GERMANY|20230715| true|               1|
+---+-------+--------+-----+----------------+

原理说明

  • group_id的作用是把每个FLAG=False之后的连续FLAG=True行划分为同一个组:第一行group_id=0(无前置FLAG=False),第二行FLAG=False使group_id变为1,第三、四行属于group_id=1,第五行FLAG=False让group_id变为2,第六行属于group_id=2。
  • 基于COUNTRY和group_id分区后,row_number()会在每个组内重新开始计数,从而实现遇到FLAG=False时重置计数的效果。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.29 05:18:10