PySpark如何按分区键生成滚动窗口聚合的Map类型结果列
问题描述
借助pyspark.sql函数实现基于指定Window()窗口规范的滚动聚合,生成存储键值对的Map类型新字段,字段值为窗口内各维度键对应的累计聚合结果。
可复现测试数据
df = spark.createDataFrame( [ ('AK', "2022-05-02", 1651449600, 'US', 3), ('AK', "2022-05-03", 1651536000, 'ON', 1), ('AK', "2022-05-04", 1651622400, 'CO', 1), ('AK', "2022-05-06", 1651795200, 'AK', 1), ('AK', "2022-05-06", 1651795200, 'US', 5) ], ["state", "ds", "ds_num", "region", "count"] ) df.show()
运行输出:
+-----+----------+----------+------+-----+ |state| ds| ds_num|region|count| +-----+----------+----------+------+-----+ | AK|2022-05-02|1651449600| US| 3| | AK|2022-05-03|1651536000| ON| 1| | AK|2022-05-04|1651622400| CO| 1| | AK|2022-05-06|1651795200| AK| 1| | AK|2022-05-06|1651795200| US| 5| +-----+----------+----------+------+-----+
已实现的部分方案
- 收集窗口帧范围内的region集合
import pyspark.sql.functions as F from pyspark.sql.window import Window days = lambda i: i * 86400 df.withColumn('regions_4W', F.collect_set('region').over( Window().partitionBy('state').orderBy('ds_num').rangeBetween(-days(27),0)))\ .sort('ds')\ .show()
运行输出:
+-----+----------+----------+------+-----+----------------+ |state| ds| ds_num|region|count| regions_4W| +-----+----------+----------+------+-----+----------------+ | AK|2022-05-02|1651449600| US| 3| [US]| | AK|2022-05-03|1651536000| ON| 1| [US, ON]| | AK|2022-05-04|1651622400| CO| 1| [CO, US, ON]| | AK|2022-05-06|1651795200| AK| 1|[CO, US, ON, AK]| | AK|2022-05-06|1651795200| US| 5|[CO, US, ON, AK]| +-----+----------+----------+------+-----+----------------+
- 按state、ds维度聚合生成当日count值Map
df\ .groupby('state', 'ds', 'ds_num')\ .agg(F.map_from_entries(F.collect_list(F.struct("region", "count"))).alias("count_rolling_4W"))\ .sort('ds')\ .show()
运行输出:
+-----+----------+----------+------------------+ |state| ds| ds_num| count_rolling_4W| +-----+----------+----------+------------------+ | AK|2022-05-02|1651449600| {US -> 3}| | AK|2022-05-03|1651536000| {ON -> 1}| | AK|2022-05-04|1651622400| {CO -> 1}| | AK|2022-05-06|1651795200|{AK -> 1, US -> 5}| +-----+----------+----------+------------------+
预期输出效果
需要得到指定滚动窗口内所有维度键对应聚合值的Map类型列,相同region的count值在窗口内累加,效果如下:
+-----+----------+----------+------------------------------------+ |state| ds| ds_num| count_rolling_4W| +-----+----------+----------+------------------------------------+ | AK|2022-05-02|1651449600| {US -> 3}| | AK|2022-05-03|1651536000| {US -> 3, ON -> 1}| | AK|2022-05-04|1651622400| {US -> 3, ON -> 1, CO -> 1}| | AK|2022-05-06|1651795200|{US -> 8, ON -> 1, CO -> 1, AK -> 1}| +-----+----------+----------+------------------------------------+
实现代码
核心思路是通过窗口函数收集滚动范围内所有维度键和计数值的列表,再用高阶函数遍历列表累加相同键的数值,最终直接生成累计Map,无需额外关联补全维度(要求Spark版本 >= 3.0,支持内置高阶函数):
import pyspark.sql.functions as F from pyspark.sql.window import Window days = lambda i: i * 86400 # 定义滚动窗口规范:按state分区,按时间戳排序,范围为过去27天到当天 roll_window = Window.partitionBy('state').orderBy('ds_num').rangeBetween(-days(27), 0) result = ( df # 逐行计算窗口内累计Map .withColumn('count_rolling_4W', F.aggregate( # 收集窗口内所有(region, count)结构 F.collect_list(F.struct('region', 'count')).over(roll_window), # 初始值为空Map F.create_map().cast("map<string,bigint>"), # 遍历合并逻辑:相同region的count累加 lambda acc, item: F.map_concat( acc, F.create_map(item.region, F.coalesce(acc[item.region], F.lit(0)) + item.count) ) ) ) # 同一天的多行计算结果完全一致,去重保留每日一行 .dropDuplicates(['state', 'ds', 'ds_num']) .select('state', 'ds', 'ds_num', 'count_rolling_4W') .sort('ds') ) result.show(truncate=False)
运行输出完全匹配预期:
+-----+----------+----------+------------------------------------+ |state|ds |ds_num |count_rolling_4W | +-----+----------+----------+------------------------------------+ |AK |2022-05-02|1651449600|{US -> 3} | |AK |2022-05-03|1651536000|{US -> 3, ON -> 1} | |AK |2022-05-04|1651622400|{US -> 3, ON -> 1, CO -> 1} | |AK |2022-05-06|1651795200|{US -> 8, ON -> 1, CO -> 1, AK -> 1}| +-----+----------+----------+------------------------------------+
注:Map类型的键展示顺序由内部存储逻辑决定,不影响键值对的实际取值。
内容的提问来源于stack exchange,提问作者Adrian
相关产品推荐
相关产品推荐

