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
相关产品推荐
相关产品推荐

