Spark Scala:Window函数计算移动平均出现空值问题求助
问题描述
输入DataFrame
+---+----------+----------+--------+-----+-------------------+ | id|product_id|sales_date|quantity|price| timestampCol| +---+----------+----------+--------+-----+-------------------+ | 1| 1|2022-12-31| 10| 10.0|2022-12-31 00:00:00| | 2| 1|2023-01-01| 10| 10.0|2023-01-01 00:00:00| | 3| 1|2023-01-02| 12| 12.0|2023-01-02 00:00:00| | 4| 1|2023-01-03| 15| 15.0|2023-01-03 00:00:00| | 5| 2|2023-01-01| 8| 8.0|2023-01-01 00:00:00| | 6| 2|2023-01-02| 10| 10.0|2023-01-02 00:00:00| | 7| 2|2023-01-03| 12| 12.0|2023-01-03 00:00:00| +---+----------+----------+--------+-----+-------------------+
任务要求
按product_id分区,计算价格的2天移动平均,窗口需包含当前时间戳和前一个时间戳。例如:
- id=2时,平均值应为
(10.0 + 10.0)/2 - id=3时,平均值应为
(12.0 + 10.0)/2
尝试的代码
val productWindow = Window .partitionBy(countriesWithTS("product_id")).orderBy(countriesWithTS("timestampCol")) .rowsBetween(2, Window.currentRow) countriesWithTS .withColumn("moved_avg", round(avg(countriesWithTS("price")).over(productWindow), 2)) .show()
错误结果
运行后moved_avg列全部为null:
+---+----------+----------+--------+-----+-------------------+---------+ | id|product_id|sales_date|quantity|price| timestampCol|moved_avg| +---+----------+----------+--------+-----+-------------------+---------+ | 1| 1|2022-12-31| 10| 10.0|2022-12-31 00:00:00| null| | 2| 1|2023-01-01| 10| 10.0|2023-01-01 00:00:00| null| | 3| 1|2023-01-02| 12| 12.0|2023-01-02 00:00:00| null| | 4| 1|2023-01-03| 15| 15.0|2023-01-03 00:00:00| null| | 5| 2|2023-01-01| 8| 8.0|2023-01-01 00:00:00| null| | 6| 2|2023-01-02| 10| 10.0|2023-01-02 00:00:00| null| | 7| 2|2023-01-03| 12| 12.0|2023-01-03 00:00:00| null| +---+----------+----------+--------+-----+-------------------+---------+
已知price字段类型为StructField("price", FloatType, nullable = true),注释掉rowsBetween部分后平均值可正常计算,但不是所需的移动平均。请问哪里出错了?
解决方案
错误原因
你写的rowsBetween(2, Window.currentRow)完全搞反了窗口范围:
rowsBetween的参数规则是**(起始位置, 结束位置)**,起始位置是相对于当前行的偏移量:负数表示当前行之前的行,正数表示当前行之后的行,0代表当前行。- 你传的
2是指当前行往后数第2行(数据按时间升序排列,后面的行是更晚的时间),但你的每个分区里最多只有4条或3条数据,根本不存在当前行往后第2行的记录,窗口内没有数据,计算出来的平均值自然是null。
正确代码
要实现包含当前行和前一行的2行移动平均,窗口范围应该设置为rowsBetween(-1, Window.currentRow):
val productWindow = Window .partitionBy(countriesWithTS("product_id")) .orderBy(countriesWithTS("timestampCol")) .rowsBetween(-1, Window.currentRow) // 关键修正:起始位置改为-1 countriesWithTS .withColumn("moved_avg", round(avg(countriesWithTS("price")).over(productWindow), 2)) .show()
结果说明
运行修正后的代码会得到符合要求的结果:
- 对于每个分区的第一行(比如id=1、id=5),因为没有前一行,窗口内只有当前行,所以移动平均等于当前价格
- id=2的移动平均为
(10.0+10.0)/2=10.0 - id=3的移动平均为
(12.0+10.0)/2=11.0,以此类推
如果希望第一行的移动平均显示null(因为没有前一行数据),可以额外加判断逻辑:
countriesWithTS .withColumn("moved_avg", when(row_number().over(productWindow) > 1, round(avg(countriesWithTS("price")).over(productWindow), 2)) .otherwise(null)) .show()
内容的提问来源于stack exchange,提问作者Jelly
相关产品推荐
相关产品推荐

