Spark Scala如何按日期范围分组计算近12个月销售累计值?
问题:计算Spark DataFrame中近12个月滚动销售累计值
场景与数据
在Spark Scala环境下,现有DF_City数据集包含city、state、year_month、saleCount字段,原始数据如下:
+-------------+-------------+----------+------------+ | city| state |year_month|saleCount | +-------------+-------------+----------+------------+ | Bangalore | Karnataka | 2020-01| 10| | Bangalore | Karnataka | 2020-02| 10| | Bangalore | Karnataka | 2021-03| 10| | Bangalore | Karnataka | 2021-04| 10| | Bangalore | Karnataka | 2021-05| 10| | Bangalore | Karnataka | 2021-06| 10| | Bangalore | Karnataka | 2021-07| 10| | Bangalore | Karnataka | 2021-08| 10| | Bangalore | Karnataka | 2021-09| 10| | Bangalore | Karnataka | 2021-10| 10| | Chennai | Tamil Nadu| 2020-05| 20| | Chennai | Tamil Nadu| 2020-06| 20| | Chennai | Tamil Nadu| 2020-07| 20| | Chennai | Tamil Nadu| 2020-08| 20| | Chennai | Tamil Nadu| 2020-09| 20| | Chennai | Tamil Nadu| 2020-10| 20| | Chennai | Tamil Nadu| 2020-11| 20| +-------------+-------------+----------+------------+
需求
需生成新增last12MonthsSellCount字段的目标DataFrame,该字段为对应记录所在city和state分组内、当前year_month及过去12个月内的saleCount累计值,目标数据示例如下:
| city| state |year_month|saleCount | last12MonthsSellCount | +-------------+-------------+----------+------------+------------------------+ | Bangalore | Karnataka | 2020-01| 10| 10 | | Bangalore | Karnataka | 2020-02| 10| 20 | | Bangalore | Karnataka | 2021-03| 10| 30 | | Bangalore | Karnataka | 2021-04| 10| 40 | | Bangalore | Karnataka | 2021-05| 10| 50 | | Bangalore | Karnataka | 2021-06| 10| 60 | | Bangalore | Karnataka | 2021-07| 10| 70 | | Bangalore | Karnataka | 2021-08| 10| 80 | | Bangalore | Karnataka | 2021-09| 10| 90 | | Bangalore | Karnataka | 2021-10| 10|100 | | Chennai | Tamil Nadu| 2020-05| 20| 20 | | Chennai | Tamil Nadu| 2020-06| 20| 40 | | Chennai | Tamil Nadu| 2020-07| 20| 60 | | Chennai | Tamil Nadu| 2020-08| 20| 80 | | Chennai | Tamil Nadu| 2020-09| 20|100 | | Chennai | Tamil Nadu| 2020-10| 20|120 | | Chennai | Tamil Nadu| 2020-11| 20|140 | +-------------+-------------+----------+------------+------------------------+
原代码问题分析
用户尝试的代码未达预期,代码如下:
val cityStateMonthlyYearlycount = lastCityCountDF.withColumn("yearMonth", col("year_month")).groupBy(col("city"), col("state"), col("year_month")).agg(sum(when(datediff(col("year_month"), col("yearMonth")).leq(12), col("monthlyCount")).otherwise(0)).as("lastOneYearCount")).filter(col("lastOneYearCount") === 0).select("city", "state", "year_month", "lastOneYearCount")
存在的核心问题:
- 日期计算错误:
datediff用于计算天数差,不是月份差,无法正确判断"过去12个月"的范围;且year_month是字符串类型,不能直接传入datediff。 - 分组逻辑错误:按
city、state、year_month分组后,同一组内的year_month完全相同,导致datediff结果恒为0,只能统计当月数据,无法实现滚动累计。 - 字段名错误:原数据集字段为
saleCount,代码中误用了不存在的monthlyCount。 - 过滤逻辑错误:最后过滤
lastOneYearCount === 0的记录,直接丢弃了所有有效累计数据。
正确实现方案(DataFrame方式)
采用窗口函数的滚动范围窗口实现,步骤如下:
import org.apache.spark.sql.functions._ import org.apache.spark.sql.expressions.Window // 1. 将year_month转换为日期类型,再转为yyyyMM格式的整数,方便计算月份范围 val dfWithMonthNum = DF_City .withColumn("year_month_date", to_date(col("year_month"), "yyyy-MM")) .withColumn("month_num", year(col("year_month_date")) * 100 + month(col("year_month_date"))) // 2. 定义滚动窗口:按city、state分组,按month_num排序,范围覆盖当前月份及过去12个月 val rolling12MonthsWindow = Window .partitionBy("city", "state") .orderBy("month_num") .rangeBetween(-12, Window.currentRow) // 3. 计算近12个月累计销售值,保留目标字段 val resultDF = dfWithMonthNum .withColumn("last12MonthsSellCount", sum("saleCount").over(rolling12MonthsWindow)) .select("city", "state", "year_month", "saleCount", "last12MonthsSellCount") .orderBy("city", "year_month") // 查看结果 resultDF.show()
代码说明
- 将
year_month转为month_num(如2020-01转为202001),通过数值范围rangeBetween(-12, currentRow)精准覆盖"当前月份及过去12个月"的所有记录。 - 窗口函数
sum("saleCount").over(rolling12MonthsWindow)会在每个city+state分组内,对符合月份范围的记录自动求和,生成滚动累计值。 - 最后筛选并排序目标字段,得到符合需求的结果。
内容的提问来源于stack exchange,提问作者Pelab
相关产品推荐
相关产品推荐

