PySpark中GroupBy结合滚动平均值实现报错问题排查
问题分析与解决
你的错误在于混淆了窗口函数和GroupBy聚合的用法:窗口函数已经通过partitionBy("group")实现了分组逻辑,同时依赖order列确定滚动窗口的顺序,此时再对group做GroupBy会导致Spark无法处理order列——因为GroupBy会合并同一组的行,而order列既不在GroupBy字段里,也没有被聚合处理。
滚动平均值是针对分组内的每一行计算的(每行对应自己的滚动窗口结果),不需要额外的GroupBy操作,直接用窗口函数后选择所需字段即可。
修正后的代码
import pandas as pd from pyspark.sql import SparkSession from pyspark.sql.window import Window from pyspark.sql.functions import col, avg # 初始化SparkSession(未初始化时需执行) spark = SparkSession.builder.appName("RollingAvgDemo").getOrCreate() data = pd.DataFrame({ 'group':['A']*5+['B']*5, 'order':[1,2,3,4,5, 1,2,3,4,5], 'value':[23, 54, 65, 64, 78, 98, 78, 76, 77, 57] }) spark_df = spark.createDataFrame(data) # 定义窗口:按group分区,按order排序,窗口覆盖当前行与前一行 window_spec = Window.partitionBy("group").orderBy("order").rowsBetween(-1, 0) # 计算滚动平均值 rolling_avg = avg(col("value")).over(window_spec).alias("value_rolling_avg") # 直接选择原字段+滚动平均值字段,无需GroupBy spark_df.select("group", "order", "value", rolling_avg).show()
输出结果
+-----+-----+-----+------------------+ |group|order|value|value_rolling_avg | +-----+-----+-----+------------------+ | A| 1| 23| 23.0| # 第一行无前置行,仅取自身值 | A| 2| 54| 38.5| | A| 3| 65| 59.5| | A| 4| 64| 64.5| | A| 5| 78| 71.0| | B| 1| 98| 98.0| | B| 2| 78| 88.0| | B| 3| 76| 77.0| | B| 4| 77| 76.5| | B| 5| 57| 67.0| +-----+-----+-----+------------------+
补充说明
如果需要对每个分组的滚动平均值做进一步聚合(比如取每个分组的最后一个滚动值),可在窗口计算后再做GroupBy:
# 先计算滚动平均值,再按group分组取最后一行的滚动值 spark_df.withColumn("value_rolling_avg", rolling_avg) \ .groupBy("group") \ .agg({"value_rolling_avg": "last"}) \ .show()
内容的提问来源于stack exchange,提问作者Henri
相关产品推荐
相关产品推荐

