PySpark中计算分组窗口下各子组最大值的滚动平均值实现方法
实现思路
你需要的逻辑无法直接用普通窗口聚合实现,核心是需要先对窗口范围内的数据按子组取最大值,再对这些最大值求平均,我们可以通过窗口函数+高阶聚合函数的组合实现该逻辑。
完整实现代码
from pyspark.sql import functions as F from pyspark.sql.window import Window # 你的示例数据构造 df = spark.createDataFrame( [(1, 'a', 1, 5.0), (1, 'a', 2, 10.0), (1, 'a', 3, 25.0), (1, 'a', 4, 50.0), (1, 'a', 5, 75.0), (1, 'b', 3, 100.0), (1, 'b', 4, 30.0), (1, 'b', 5, 60.0), (1, 'b', 6, 90.0), (1, 'b', 7, 120.0), (2, 'c', 1, 200.0), (2, 'c', 2, 400.0), (2, 'c', 3, 600.0), (2, 'c', 4, 800.0), (2, 'c', 5, 1000.0), (2, 'c', 6, 1200.0), (2, 'c', 7, 1300.0), (2, 'c', 8, 1400.0), (2, 'd', 5, 150.0), (2, 'd', 6, 250.0), (2, 'd', 7, 350.0)], ("group", "sub-group","time", "value")) # 你定义的窗口规则 w = Window.partitionBy('group').orderBy('time').rangeBetween(-2, -1) # 核心计算逻辑 df = df.withColumn( "avg_max_value", F.aggregate( # 收集当前窗口内所有(子组、值)对 F.collect_list(F.struct("sub-group", "value")).over(w), # 初始状态:空Map,存储每个子组的最大值 F.create_map().cast("map<string, double>"), # 遍历更新Map:相同子组保留最大值 lambda acc, row: F.map_concat( F.map_filter(acc, lambda k, _: k != row["sub-group"]), F.create_map( row["sub-group"], F.greatest(acc.get_item(row["sub-group"]), row["value"]) ) ), # 最终计算所有子组最大值的平均值(PySpark3.0+支持array_avg) lambda max_map: F.array_avg(F.map_values(max_map)) ) ) # 验证结果 df.orderBy("group", "sub-group", "time").show()
低版本兼容说明
如果你使用的PySpark版本低于3.0没有array_avg函数,可以把最终处理逻辑替换为如下表达式:
lambda max_map: F.expr(""" aggregate( map_values(max_map), 0D, (acc, v) -> acc + v, acc -> if(size(map_values(max_map)) = 0, null, acc / size(map_values(max_map))) ) """)
内容的提问来源于stack exchange,提问作者Vincent
相关产品推荐
相关产品推荐

