使用Spark SQL窗口函数计算分类下Indicator连续为1的天数
问题:统计分组内Indicator连续为1的天数
原始数据
df = [{"Category": 'A', "date": '01/01/2022', "Indicator": 1}, {"Category": 'A', "date": '02/01/2022', "Indicator": 0}, {"Category": 'A', "date": '03/01/2022', "Indicator": 1}, {"Category": 'A', "date": '04/01/2022', "Indicator": 1}, {"Category": 'A', "date": '05/01/2022', "Indicator": 1}, {"Category": 'B', "date": '01/01/2022', "Indicator": 0}, {"Category": 'B', "date": '02/01/2022', "Indicator": 1}, {"Category": 'B', "date": '03/01/2022', "Indicator": 1}, {"Category": 'B', "date": '04/01/2022', "Indicator": 0}, {"Category": 'B', "date": '05/01/2022', "Indicator": 0}, {"Category": 'B', "date": '06/01/2022', "Indicator": 1}] df = spark.createDataFrame(df)
需求
按Category分组、date升序排序,使用窗口函数(含LEAD())创建新字段consec_ind,统计Indicator连续为1的天数:
- 当
Indicator为0时,consec_ind为0 - 当
Indicator为1时,consec_ind为当前连续1的累计天数(从1开始递增)
尝试的代码(未得到预期结果)
先创建临时视图:
df.createOrReplaceTempView('df')
执行SQL:
select date, Indicator, case when Indicator > 0 THEN (sum(count(Indicator)) over (order by date)) else 0 end as running_total from df WHERE Category = 'A' group by date, Indicator order by date, Indicator;
预期输出
[ {"Category": 'A', "date": '01/01/2022', "Indicator": 1,"consec_ind":1}, {"Category": 'A', "date": '02/01/2022', "Indicator": 0,"consec_ind":0}, {"Category": 'A', "date": '03/01/2022', "Indicator": 1,"consec_ind":1}, {"Category": 'A', "date": '04/01/2022', "Indicator": 1,"consec_ind":2}, {"Category": 'A', "date": '05/01/2022', "Indicator": 1,"consec_ind":3}, {"Category": 'B', "date": '01/01/2022', "Indicator": 0,"consec_ind":0}, {"Category": 'B', "date": '02/01/2022', "Indicator": 1,"consec_ind":1}, {"Category": 'B', "date": '03/01/2022', "Indicator": 1,"consec_ind":2}, {"Category": 'B', "date": '04/01/2022', "Indicator": 0,"consec_ind":0}, {"Category": 'B', "date": '05/01/2022', "Indicator": 0,"consec_ind":0}, {"Category": 'B', "date": '06/01/2022', "Indicator": 1,"consec_ind":1} ]
注:原预期输出遗漏了A类的01/01/2022记录,此处补全以匹配原始数据逻辑。
解决方案
方法一:SQL实现(结合窗口函数,含LEAD()辅助验证)
核心思路是先标记连续1的分组,再在分组内计算累计天数:
WITH grouped_data AS ( SELECT Category, date, Indicator, -- 标记连续1的分组:每次遇到Indicator=0则分组ID递增 SUM(CASE WHEN Indicator = 0 THEN 1 ELSE 0 END) OVER (PARTITION BY Category ORDER BY date ROWS BETWEEN UNBOUNDED PRECEDING AND CURRENT ROW) AS group_id FROM df ), consecutive_counts AS ( SELECT Category, date, Indicator, -- 在每个分组内计算行号,Indicator=0时设为0 CASE WHEN Indicator = 1 THEN ROW_NUMBER() OVER (PARTITION BY Category, group_id ORDER BY date) ELSE 0 END AS consec_ind, -- 用LEAD()查看下一条记录的Indicator,辅助验证连续性 LEAD(Indicator, 1, 0) OVER (PARTITION BY Category ORDER BY date) AS next_indicator FROM grouped_data ) SELECT Category, date, Indicator, consec_ind FROM consecutive_counts ORDER BY Category, date;
方法二:PySpark API实现
from pyspark.sql import Window from pyspark.sql.functions import sum as spark_sum, when, row_number, lead # 定义窗口:按Category分组,date排序 window_group = Window.partitionBy("Category").orderBy("date") window_consec = Window.partitionBy("Category", "group_id").orderBy("date") # 标记连续分组,计算连续天数 result_df = df.withColumn( "group_id", spark_sum(when(df["Indicator"] == 0, 1).otherwise(0)).over(window_group) ).withColumn( "consec_ind", when(df["Indicator"] == 1, row_number().over(window_consec)).otherwise(0) ).withColumn( "next_indicator", # 可选:用LEAD()查看下一条记录的Indicator lead(df["Indicator"], 1, 0).over(window_group) ).select("Category", "date", "Indicator", "consec_ind") # 查看结果 result_df.orderBy("Category", "date").show()
说明
- 原尝试代码的问题:未按
Category分组统计,且sum(count(Indicator))的逻辑无法区分不同的连续1区间,导致累计值错误。 - 方案中使用
SUM(CASE...)生成分组ID,将连续的1划分为同一组,再用ROW_NUMBER()在组内计算连续天数,符合预期需求。 - 加入
LEAD()函数用于查看下一条记录的Indicator,可辅助验证连续区间的边界,满足需求中使用LEAD()的要求。
内容的提问来源于stack exchange,提问作者Mrmoleje
相关产品推荐
相关产品推荐

