PySpark:按需重置累积求和(cumsum)列的实现问题
PySpark实现按新id_B引入重置对应cumsum值的需求
需求说明
现有PySpark DataFrame,需将cumsum列转换为new_cumsum列,规则如下:
- 当新的
id_B被引入(即reset=True的行),该id_B的new_cumsum重置为对应初始值; - 已存在的
id_B,保留原有累积求和值,不受新id_B引入的影响。
示例场景
- 2023-04-05引入
id_B=2002时,其new_cumsum为3,id_B=2001的cumsum保持为4; - 2023-04-09引入
id_B=2003时,其new_cumsum为0,id_B=2001和id_B=2002的累积值不受影响。
原始DataFrame代码
import pyspark.sql.functions as F from pyspark.sql import Window df = spark.createDataFrame( [ (1001, 2001, "2023-04-01", "2023-04-01", False, 0, 0), (1001, 2001, "2023-04-02", "2023-04-01", False, 0, 0), (1001, 2001, "2023-04-03", "2023-04-01", False, 1, 1), (1001, 2001, "2023-04-04", "2023-04-01", False, 1, 1), (1001, 2002, "2023-04-05", "2023-04-05", True, 4, 3), (1001, 2001, "2023-04-05", "2023-04-01", False, 4, 4), (1001, 2001, "2023-04-06", "2023-04-01", False, 4, 4), (1001, 2002, "2023-04-06", "2023-04-05", False, 4, 3), (1001, 2001, "2023-04-07", "2023-04-01", False, 4, 4), (1001, 2002, "2023-04-07", "2023-04-05", False, 4, 3), (1001, 2001, "2023-04-08", "2023-04-01", False, 10, 10), (1001, 2002, "2023-04-08", "2023-04-05", False, 10, 9), (1001, 2003, "2023-04-09", "2023-04-09", True, 10, 0), (1001, 2001, "2023-04-09", "2023-04-01", False, 10, 10), (1001, 2002, "2023-04-09", "2023-04-05", False, 10, 9), (1001, 2001, "2023-04-10", "2023-04-01", False, 12, 12), (1001, 2002, "2023-04-10", "2023-04-05", False, 12, 11), (1001, 2003, "2023-04-10", "2023-04-09", False, 12, 2), (1001, 2001, "2023-04-11", "2023-04-01", False, 13, 13), (1001, 2002, "2023-04-11", "2023-04-05", False, 13, 12), (1001, 2003, "2023-04-11", "2023-04-09", False, 13, 3), ], ["id_A", "id_B", "date", "id_B_entry_date", "reset", "cumsum", "new_cumsum"], ) df.show()
尝试的错误代码
w1 = Window.partitionBy("id_A").orderBy("date") w2 = Window.partitionBy("id_A", "id_B_entry_date").orderBy("date") w3 = Window.partitionBy("partition2", "id_A", "id_B_entry_date").orderBy("date") df2 = ( df .withColumn("diff", F.col("cumsum") - F.lag("cumsum", default=0).over(w2)) .withColumn("partition", F.when(F.col("reset"), 1).otherwise(0)) .withColumn("partition2", F.sum("partition").over(w1)) .withColumn("new_cumsum_attempt", F.sum(F.col("diff")).over(w3)) .drop("diff", "partition", "partition2") ) df2.orderBy('date').show()
正确实现方案
思路
核心逻辑是:每个id_B的累积值仅从其加入日期(id_B_entry_date)开始,累加全局每日的cumsum增量。这样新id_B加入时会以当天的增量作为初始值,后续跟随全局增量累积;旧id_B则继续累积所有增量,不受新id引入影响。
代码实现
import pyspark.sql.functions as F from pyspark.sql import Window # 步骤1:计算全局(按id_A分组)每日的cumsum增量 w_global = Window.partitionBy("id_A").orderBy("date") df_with_increment = df.withColumn( "daily_increment", F.col("cumsum") - F.lag("cumsum", default=0).over(w_global) ) # 步骤2:对每个id_A+id_B分组,从entry_date开始累加每日增量 w_idb = Window.partitionBy("id_A", "id_B").orderBy("date") df_result = df_with_increment.withColumn( "new_cumsum_correct", F.sum( F.when(F.col("date") >= F.col("id_B_entry_date"), F.col("daily_increment")).otherwise(0) ).over(w_idb) ).drop("daily_increment") # 查看结果 df_result.orderBy("date", "id_B").show()
验证结果
运行后new_cumsum_correct列将与示例中的new_cumsum列完全匹配,满足需求:
id_B=2002在2023-04-05的new_cumsum_correct为3;id_B=2003在2023-04-09的new_cumsum_correct为0;- 旧id_B的累积值保持连续不变。
内容的提问来源于stack exchange,提问作者MS25
相关产品推荐
相关产品推荐

