PySpark如何统计指定列特定值的连续出现次数
问题描述
我有一个名为info的列,数据集样例如下:
| Timestamp | info | +-------------------+----------+ |2016-01-01 17:54:30| 0 | |2016-02-01 12:16:18| 0 | |2016-03-01 12:17:57| 0 | |2016-04-01 10:05:21| 0 | |2016-05-11 18:58:25| 1 | |2016-06-11 11:18:29| 1 | |2016-07-01 12:05:21| 0 | |2016-08-11 11:58:25| 0 | |2016-09-11 15:18:29| 1 |
需求是统计info列中值为1的连续出现次数,其余位置填充0,最终结果的res列效果如下:
--------------------+----------+----------+ | Timestamp | info | res | +-------------------+----------+----------+ |2016-01-01 17:54:30| 0 | 0 | |2016-02-01 12:16:18| 0 | 0 | |2016-03-01 12:17:57| 0 | 0 | |2016-04-01 10:05:21| 0 | 0 | |2016-05-11 18:58:25| 1 | 1 | |2016-06-11 11:18:29| 1 | 2 | |2016-07-01 12:05:21| 0 | 0 | |2016-08-11 11:58:25| 0 | 0 | |2016-09-11 15:18:29| 1 | 1 |
之前尝试的代码无法得到正确结果:
df_input = df_input.withColumn( "res", F.when( df_input.info == F.lag(df_input.info).over(w1), F.sum(F.lit(1)).over(w1) ).otherwise(0) )
实现思路
之前的代码问题在于没有给连续相同值划分独立的分组,直接按全局排序开窗累加会把不连续的同值段算到一起。正确做法分三步:
- 先按时间排序开窗,判断当前行和上一行的
info值是否相等,不等的时候标记分组起点 - 累加分组标记生成唯一的连续段ID,把每一段连续的0或1划成独立分组
- 对
info=1的分组,按分组内排序做行号计数,info=0的位置直接填0即可
正确代码
首先定义按时间升序排列的窗口,替换原有逻辑即可:
from pyspark.sql import functions as F from pyspark.sql.window import Window # 全局按时间排序的基础窗口 w_order = Window.orderBy("Timestamp") # 连续段分组内排序的窗口,用于计数 w_partition = Window.partitionBy("group_id").orderBy("Timestamp") df_res = df_input.withColumn( # 标记值发生变化的行,作为连续段起点 "is_change", F.when(F.lag("info").over(w_order) == F.col("info"), 0).otherwise(1) ).withColumn( # 累加变化标记,为每个连续段生成唯一ID "group_id", F.sum("is_change").over(w_order.rowsBetween(Window.unboundedPreceding, Window.currentRow)) ).withColumn( # 值为1的段组内生成连续计数,其余位置填0 "res", F.when(F.col("info") == 1, F.row_number().over(w_partition)).otherwise(0) ).drop("is_change", "group_id")
执行后输出结果和预期的res列完全一致。
内容的提问来源于stack exchange,提问作者Babbara
相关产品推荐
相关产品推荐

