PySpark中非递归实现依赖自身历史值的累计计算
问题:计算依赖历史值的Required列(PySpark非递归实现)
我需要计算依赖自身历史值的Required列,逻辑如下:
- 当Col D=1时,取值为Col E;
- 当Col D≠1时,取值为Col E减去此前所有
Required列值的总和。
我尝试用窗口函数编写了以下SQL代码,但仅能正确计算第二行的值,第三行及以后无法累加已计算出的Required值,结果不符合预期,想知道能否在PySpark中通过非递归方式实现该需求:
SELECT A,B,C,D,E,F, E - SUM(F) OVER(PARTITION BY A,B,C ORDER BY D) AS F1 FROM CTE
输入表
| Col A | Col B | Col C | Col D | Col E |
|---|---|---|---|---|
| Base | 3 | 1 | 1 | 0.80022360 |
| Base | 3 | 1 | 2 | 0.87889069 |
| Base | 3 | 1 | 3 | 0.89611630 |
| Base | 3 | 1 | 4 | 0.91105699 |
| Base | 3 | 1 | 5 | 0.92868688 |
预期输出表
| Col A | Col B | Col C | Col D | Col E | Required |
|---|---|---|---|---|---|
| Base | 3 | 1 | 1 | 0.80022360 | 0.80022360 |
| Base | 3 | 1 | 2 | 0.87889069 | 0.07866708 |
| Base | 3 | 1 | 3 | 0.89611630 | 0.01722561 |
| Base | 3 | 1 | 4 | 0.91105699 | 0.01494068 |
| Base | 3 | 1 | 5 | 0.92868688 | 0.01762989 |
计算示例
第二行:0.87889069 (Col E) - 0.80022360(前一行Required) = 0.07866708
第三行:0.89611630(Col E) - (0.80022360+0.07866708)(前两行Required总和) = 0.01722561
非递归实现方案(PySpark)
通过拆解需求逻辑和预期结果,可以发现核心规律:每行的Required值等于当前行的Col E减去上一行的Col E(第一行无历史数据,直接取Col E)。推导过程如下:
- 第1行:
Required₁ = E₁,此时历史Required总和S₁ = E₁ - 第2行:
Required₂ = E₂ - S₁ = E₂ - E₁,总和S₂ = S₁ + Required₂ = E₂ - 第3行:
Required₃ = E₃ - S₂ = E₃ - E₂,总和S₃ = E₃ - 以此类推,每行的历史
Required总和等于上一行的Col E,因此无需递归,直接用LAG窗口函数即可实现。
PySpark代码实现
from pyspark.sql import SparkSession from pyspark.sql.window import Window from pyspark.sql.functions import lag, col, when # 初始化Spark会话 spark = SparkSession.builder.appName("RequiredColumnCalc").getOrCreate() # 构造输入DataFrame data = [ ("Base", 3, 1, 1, 0.80022360), ("Base", 3, 1, 2, 0.87889069), ("Base", 3, 1, 3, 0.89611630), ("Base", 3, 1, 4, 0.91105699), ("Base", 3, 1, 5, 0.92868688) ] df = spark.createDataFrame(data, ["Col A", "Col B", "Col C", "Col D", "Col E"]) # 定义窗口:按Col A/B/C分区,Col D排序 window_spec = Window.partitionBy("Col A", "Col B", "Col C").orderBy("Col D") # 计算Required列 df_result = df.withColumn( "Required", when( col("Col D") == 1, col("Col E") ).otherwise( col("Col E") - lag(col("Col E"), 1).over(window_spec) ) ).select("Col A", "Col B", "Col C", "Col D", "Col E", "Required") # 打印结果 df_result.show(truncate=False)
运行代码后输出的结果与预期一致(浮点精度差异属于正常现象)。
内容的提问来源于stack exchange,提问作者Liz
相关产品推荐
相关产品推荐

