PySpark DataFrame分组统计分类列状态变更次数方法
问题说明
现有按year、month、day、hour升序排序的PySpark DataFrame,时间步长为2小时,样例数据构造代码如下:
data = [(2010, 3, 12, 0, 'p1', 'state1'), (2010, 3, 12, 0, 'p2', 'state2'), (2010, 3, 12, 0, 'p3', 'state1'), (2010, 3, 12, 0, 'p4', 'state2'), (2010, 3, 12, 2, 'p1', 'state3'), (2010, 3, 12, 2, 'p2', 'state1'), (2010, 3, 12, 2, 'p3', 'state3'), (2010, 3, 12, 4, 'p1', 'state1'), (2010, 3, 12, 6, 'p1', 'state1')] columns = ['year', 'month', 'day', 'hour', 'process_id','state'] df = spark.createDataFrame(data=data, schema=columns)
需求为按process_id、year、month、day维度分组,统计每个进程当日state的状态变更总次数,期望输出如下:
+----+-----+---+----------+----------+ |year|month|day|process_id| chg_count| +----+-----+---+----------+----------+ |2010| 3| 12| p1| 2| |2010| 3| 12| p2| 1| |2010| 3| 12| p3| 1| |2010| 3| 12| p4| 0| +----+-----+---+----------+----------+
实现方案
方案1:直接在groupby+agg框架内实现
不需要新增预处理步骤,直接通过内置数组函数在聚合逻辑中完成统计,核心思路是分组内按小时排序收集状态序列,对比相邻状态的差异计数。
需要先导入pyspark函数:
from pyspark.sql import functions as F
聚合代码如下:
chg_count_df = df.groupby('process_id', 'year', 'month', 'day').agg( F.expr(""" aggregate( transform(sort_array(collect_list(struct(hour, state))), x -> x.state), named_struct('prev', element_at(transform(sort_array(collect_list(struct(hour, state))), x -> x.state), 1), 'cnt', 0), (acc, curr) -> named_struct( 'prev', curr, 'cnt', acc.cnt + if(curr != acc.prev, 1, 0) ), acc -> acc.cnt ).cast('int') as chg_count """) )
注意:即使原始数据全局按时间排序,groupby shuffle过程会打乱分组内数据顺序,必须通过
sort_array对分组内的时间、状态结构体排序,不能直接收集state列,否则会因顺序错乱导致统计错误。
方案2:生产环境最优实现(窗口函数+分组聚合)
该方案性能更稳定,逻辑更易维护,避免了单分组数据量过大时数组收集的内存开销,是优先推荐的写法:
from pyspark.sql import functions as F from pyspark.sql.window import Window # 定义分组排序窗口 w = Window.partitionBy('process_id', 'year', 'month', 'day').orderBy('hour') chg_count_df = df.withColumn('prev_state', F.lag('state').over(w)) \ .withColumn('is_change', F.when(F.col('state') != F.col('prev_state'), 1).otherwise(0)) \ .groupby('process_id', 'year', 'month', 'day') \ .agg(F.sum('is_change').cast('int').alias('chg_count'))
逻辑说明:
- 按分组维度分区、按小时排序,通过
lag函数取上一个时间点的状态 - 标记当前状态和上一状态不一致的记录为变更1,否则为0(分组内第一条记录无上一状态,标记为0)
- 分组求和变更标记,得到总变更次数,结果和预期完全一致
方案对比
- 方案1完全适配预设的groupby写法,不需要额外的列预处理步骤,但单分组下记录量过大时(如时间粒度细化到分钟级、单组记录超千条),collect_list生成的大数组会带来额外内存开销
- 方案2shuffle次数和纯groupby逻辑一致,无大对象收集开销,代码可读性更强,在全量数据规模下性能表现更稳定,是生产环境首选方案
内容的提问来源于stack exchange,提问作者Tristan Tran
相关产品推荐
相关产品推荐

