You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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'))
               )
)

错误原因

  1. 窗口分区错误:partitionBy("year_month", "date")会把每个日期单独作为一个分区,无法实现跨日期的累积计算;
  2. 普通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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.23 22:26:01