PySpark:基于条件重置累积求和(cumsum)列的实现
实现PySpark中基于reset标记的cumsum重置逻辑
现有DataFrame
+----+----------+-----+------+ | id| date|reset|cumsum| +----+----------+-----+------+ |1001|2023-04-01|false| 0| |1001|2023-04-02|false| 0| |1001|2023-04-03|false| 1| |1001|2023-04-04|false| 1| |1001|2023-04-05| true| 4| |1001|2023-04-06|false| 4| |1001|2023-04-07|false| 4| |1001|2023-04-08|false| 10| |1001|2023-04-09| true| 10| |1001|2023-04-10|false| 12| |1001|2023-04-11|false| 13| +----+----------+-----+------+
预期输出
需要新增new_cumsum列,按特定重置逻辑计算,结果如下:
+----+----------+-----+------+----------+ | id| date|reset|cumsum|new_cumsum| +----+----------+-----+------+----------+ |1001|2023-04-01|false| 0| 0| |1001|2023-04-02|false| 0| 0| |1001|2023-04-03|false| 1| 1| |1001|2023-04-04|false| 1| 1| |1001|2023-04-05| true| 4| 3| |1001|2023-04-06|false| 4| 3| |1001|2023-04-07|false| 4| 3| |1001|2023-04-08|false| 10| 6| |1001|2023-04-09| true| 10| 0| |1001|2023-04-10|false| 12| 2| |1001|2023-04-11|false| 13| 3| +----+----------+-----+------+----------+
具体逻辑
- 4/01至4/04无重置标记,
new_cumsum与原cumsum一致; - 4/05首次触发重置,因cumsum从1增至4,
new_cumsum取差值3; - 4/05至4/07 cumsum无变化,
new_cumsum保持3; - 4/08 cumsum从4增至10,
new_cumsum取差值6; - 4/09再次触发重置,cumsum无变化,
new_cumsum设为0; - 4/10 cumsum从10增至12,
new_cumsum取差值2; - 4/11 cumsum从12增至13,
new_cumsum取与最近重置日(4/09)的差值3。
PySpark解决方案
以下是实现该逻辑的代码:
from pyspark.sql import SparkSession from pyspark.sql import functions as F from pyspark.sql.window import Window # 初始化SparkSession(如果未初始化) spark = SparkSession.builder.appName("ResetCumsum").getOrCreate() # 创建原始DataFrame(如果已存在可跳过) df = spark.createDataFrame( [ (1001, "2023-04-01", False, 0), (1001, "2023-04-02", False, 0), (1001, "2023-04-03", False, 1), (1001, "2023-04-04", False, 1), (1001, "2023-04-05", True, 4), (1001, "2023-04-06", False, 4), (1001, "2023-04-07", False, 4), (1001, "2023-04-08", False, 10), (1001, "2023-04-09", True, 10), (1001, "2023-04-10", False, 12), (1001, "2023-04-11", False, 13), ], ["id", "date", "reset", "cumsum"], ) # 定义窗口:按id分区,date升序排序 window_partition = Window.partitionBy("id").orderBy("date") # 步骤1:标记重置分组 - 累计reset的次数,每次reset后开启新分组 df = df.withColumn("reset_group", F.sum(F.col("reset").cast("int")).over(window_partition)) # 步骤2:获取每个分组的基准值 window_group = Window.partitionBy("id", "reset_group").orderBy("date") df = df.withColumn( "prev_cumsum", F.lag("cumsum", 1).over(window_partition) ).withColumn( "group_first_reset", F.first(F.when(F.col("reset"), F.struct("prev_cumsum", "cumsum"))).over(window_group) ).withColumn( "base_value", F.when( F.col("reset_group") == 0, F.first("cumsum").over(window_group) ).when( F.col("group_first_reset.prev_cumsum") == F.col("group_first_reset.cumsum"), F.col("group_first_reset.cumsum") ).otherwise( F.col("group_first_reset.prev_cumsum") ) ) # 步骤3:计算new_cumsum df = df.withColumn( "new_cumsum", F.when( F.col("reset") & (F.col("cumsum") == F.col("prev_cumsum")), 0 ).otherwise( F.col("cumsum") - F.col("base_value") ) ).drop("reset_group", "prev_cumsum", "group_first_reset", "base_value") # 查看结果 df.show()
代码说明
- 重置分组标记:通过累计
reset的次数,将数据划分为不同的重置区间,每个区间对应一次reset后的计算周期。 - 基准值计算:针对每个区间确定计算
new_cumsum的基准值:- 初始区间直接用第一个cumsum值作为基准;
- 后续区间根据reset点的cumsum变化情况,选择前一条记录的cumsum或当前reset点的cumsum作为基准。
- new_cumsum计算:根据规则,当reset触发且cumsum无变化时设为0,否则用当前cumsum减去基准值得到结果。
内容的提问来源于stack exchange,提问作者MS25
相关产品推荐
相关产品推荐

