PySpark中基于历史值的累积积和计算问题求助
PySpark实现基于历史值的累积积和计算
原始DataFrame
dfz = spark.createDataFrame( [ (202401, 20240101, 20, 0, 0.1, 20.0), (202401, 20240102, 50, 20, 0.2, 50.0), (202401, 20240103, 20, 50, 0.2, 20.0), (202401, 20240104, None, 20, 0.2, None), (202401, 20240105, None, 0, 0.3, None), (202401, 20240106, None, 0, 0.1, None)], ["year_month", "date", "amount", "amount_lag", "perc_prev", "amount_prev"] )
计算规则
- 当
amount不为空时,amount_prev直接取amount的值; - 当
amount为空时,若amount_lag不为空,则amount_prev = amount_lag * (1 + perc_prev); - 当
amount和amount_lag都为空时,使用上一行计算得到的amount_prev值,按amount_prev = 上一行amount_prev * (1 + perc_prev)计算。
预期结果
dfz = spark.createDataFrame( [ (202401, 20240101, 20, 0, 0.1, 20.0), (202401, 20240102, 50, 20, 0.2, 50.0), (202401, 20240103, 20, 50, 0.2, 20.0), (202401, 20240104, None, 20, 0.2, 24.0), (202401, 20240105, None, 0, 0.3, 31.2), (202401, 20240106, None, 0, 0.1, 34.32)], ["year_month", "date", "amount", "amount_lag", "perc_prev", "amount_prev"] )
尝试的错误代码
w1 = ( Window.partitionBy("year_month", "date") .orderBy('date').rangeBetween(Window.unboundedPreceding, 0)) test = ( dbz .withColumn('amount_prev', when((~col('amount').isNull()), col('amount')) .otherwise((col('amount_lag')*col('perc_prev'))+col('amount_lag')) ) )
错误原因
- 窗口分区错误:
partitionBy("year_month", "date")会把每个日期单独作为一个分区,无法实现跨日期的累积计算; - 普通
when语句无法处理依赖上一行计算结果的递归逻辑,只能处理当前行的字段值,无法引用上一行生成的amount_prev。
正确实现方案
方法一:基于分组重置+累积乘积(推荐,性能更优)
通过标记重置点(amount非空的行)将数据分组,每组内计算累积乘积因子,结合初始值得到结果:
from pyspark.sql import functions as F from pyspark.sql import Window # 1. 定义分区窗口:按月份分组,日期排序 w_month = Window.partitionBy("year_month").orderBy("date") # 2. 标记重置点:amount非空时标记为1,否则为0,累积求和得到分组ID df = dfz.withColumn("reset", F.when(F.col("amount").isNotNull(), 1).otherwise(0)) df = df.withColumn("group_id", F.sum("reset").over(w_month)) # 3. 定义分组内窗口:按月份+分组ID分区,日期排序 w_group = Window.partitionBy("year_month", "group_id").orderBy("date") # 4. 计算每个分组内的累积乘积因子:重置行因子为1,非重置行因子为(1+perc_prev) df = df.withColumn("factor", F.when(F.col("reset") == 1, 1.0).otherwise(1 + F.col("perc_prev"))) df = df.withColumn("cum_factor", F.product("factor").over(w_group)) # 5. 获取每个分组的初始值(即重置行的amount) df = df.withColumn("initial_amount", F.first(F.col("amount"), ignorenulls=True).over(w_group)) # 6. 计算最终的amount_prev,保留两位小数 df = df.withColumn("amount_prev", F.round(F.col("initial_amount") * F.col("cum_factor"), 2)) # 7. 清理临时列并排序 result = df.drop("reset", "group_id", "factor", "cum_factor", "initial_amount") \ .orderBy("year_month", "date") result.show()
方法二:递归CTE(适合理解逻辑,小数据量场景)
逐行递归计算,依赖上一行的结果:
from pyspark.sql import functions as F from pyspark.sql import Window # 1. 为每个月份内的行添加序号 w = Window.partitionBy("year_month").orderBy("date") df = dfz.withColumn("row_num", F.row_number().over(w)) # 2. 初始化递归CTE:第一行直接取amount作为初始值 with_recursion = df.filter(F.col("row_num") == 1) \ .withColumn("amount_prev", F.col("amount")) \ .alias("base") # 3. 递归计算后续每一行 max_row = df.select(F.max("row_num")).first()[0] for i in range(2, max_row + 1): # 获取上一行的计算结果 prev_row = with_recursion.filter(F.col("row_num") == i - 1) \ .select("year_month", F.col("amount_prev").alias("prev_amount")) # 获取当前行数据 current_row = df.filter(F.col("row_num") == i) # 计算当前行的amount_prev updated_row = current_row.join(prev_row, on="year_month", how="left") \ .withColumn("amount_prev", F.when(F.col("amount").isNotNull(), F.col("amount")) .when(F.col("amount_lag").isNotNull(), F.col("amount_lag") * (1 + F.col("perc_prev"))) .otherwise(F.round(F.col("prev_amount") * (1 + F.col("perc_prev")), 2)) ) \ .drop("prev_amount") # 合并结果 with_recursion = with_recursion.union(updated_row) # 4. 清理临时列并排序 result = with_recursion.drop("row_num").orderBy("year_month", "date") result.show()
结果验证
两种方法都能得到符合预期的结果,方法一利用Spark内置窗口函数,性能更优,适合大数据量场景;方法二更直观,适合理解递归逻辑。
内容的提问来源于stack exchange,提问作者MrData
相关产品推荐
相关产品推荐

