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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.28 12:51:19