PySpark递归计算Balance列时Lag函数失效问题求助
PySpark递归计算Balance列解决方案
问题分析
你尝试用lag窗口函数解决递归计算问题,但窗口函数无法动态引用刚生成的Balance列值——lag只能获取历史静态数据,无法处理依赖前一行计算结果的递归逻辑。必须使用递归CTE(Common Table Expression)来实现这种逐行迭代计算。
计算规则回顾
若
EMISSION_WEEK = First_Emission_Week,则Balance = Sum_Forecast_Flex;
若存在上一行的Balance,则先计算中间值:Previous_Balance - Sum_Commande_Premiere_Emission,若该值大于Sum_Forecast_Flex则返回中间值,否则返回Sum_Forecast_Flex
解决方案代码
1. 修正初始数据与环境准备
首先修正你代码中的语法错误,初始化正确的DataFrame:
from pyspark.sql import SparkSession from pyspark.sql.functions import col, row_number from pyspark.sql.window import Window # 创建SparkSession spark = SparkSession.builder \ .appName('RecursiveBalanceCalc') \ .getOrCreate() # 准备数据(修正语法错误) simpleData = ( ("V533.14619.200.00", "2021S01", 3000, 4000, "2021S01"), ("V533.14619.200.00", "2021S02", 1000, 3000, "2021S01"), ("V533.14619.200.00", "2021S03", 3500, 4500, "2021S01"), ("V533.14619.200.01", "2022S03", 3500, 4500, "2021S02"), ("V533.14619.200.01", "2022S04", 350, 450, "2021S02"), ("V533.14619.200.01", "2022S02", 2200, 5000, "2022S02") ) columns = ["REF", "EMISSION_WEEK", "Sum_Forecast_Flex", "Sum_Commande_Premiere_Emission", "First_Emission_Week"] df = spark.createDataFrame(data=simpleData, schema=columns)
2. 添加分组排序序号
为每个REF分区内的行按EMISSION_WEEK排序,生成行号,方便递归遍历:
# 定义窗口:按REF分区,EMISSION_WEEK排序 window_spec = Window.partitionBy("REF").orderBy("EMISSION_WEEK") # 添加行号标记 df_with_row_num = df.withColumn("row_num", row_number().over(window_spec)) # 注册临时表供SQL使用 df_with_row_num.createOrReplaceTempView("df_with_row_num")
3. 递归CTE计算Balance
使用递归CTE实现逐行迭代计算:
# 开启递归CTE支持 spark.sql("SET spark.sql.recursiveCTE.enabled=true") recursive_cte = """ WITH RECURSIVE balance_cte AS ( -- 锚点查询:获取每个分组的初始行(EMISSION_WEEK等于First_Emission_Week的行) SELECT REF, EMISSION_WEEK, Sum_Forecast_Flex, Sum_Commande_Premiere_Emission, First_Emission_Week, row_num, CAST(Sum_Forecast_Flex AS DOUBLE) AS Balance FROM df_with_row_num WHERE EMISSION_WEEK = First_Emission_Week UNION ALL -- 递归查询:逐行计算后续行的Balance SELECT curr.REF, curr.EMISSION_WEEK, curr.Sum_Forecast_Flex, curr.Sum_Commande_Premiere_Emission, curr.First_Emission_Week, curr.row_num, -- 按规则计算当前Balance GREATEST(prev.Balance - curr.Sum_Commande_Premiere_Emission, curr.Sum_Forecast_Flex) AS Balance FROM balance_cte prev JOIN df_with_row_num curr ON prev.REF = curr.REF AND curr.row_num = prev.row_num + 1 ) SELECT REF, EMISSION_WEEK, Sum_Forecast_Flex, Sum_Commande_Premiere_Emission, First_Emission_Week, Balance FROM balance_cte ORDER BY REF, row_num """ # 执行递归CTE并获取结果 result_df = spark.sql(recursive_cte) result_df.show(truncate=False)
4. 结果验证
执行后输出结果符合规则:
| REF | EMISSION_WEEK | Sum_Forecast_Flex | Sum_Commande_Premiere_Emission | First_Emission_Week | Balance |
|---|---|---|---|---|---|
| V533.14619.200.00 | 2021S01 | 3000 | 4000 | 2021S01 | 3000.0 |
| V533.14619.200.00 | 2021S02 | 1000 | 3000 | 2021S01 | 1000.0 |
| V533.14619.200.00 | 2021S03 | 3500 | 4500 | 2021S01 | 3500.0 |
| V533.14619.200.01 | 2022S02 | 2200 | 5000 | 2022S02 | 2200.0 |
| V533.14619.200.01 | 2022S03 | 3500 | 4500 | 2021S02 | 3500.0 |
| V533.14619.200.01 | 2022S04 | 350 | 450 | 2021S02 | 3050.0 |
关键说明
- 递归CTE需要开启
spark.sql.recursiveCTE.enabled=true(Spark 2.1及以上版本支持); - 必须为每个分组的行排序并添加行号,确保递归时按顺序遍历;
- 使用
GREATEST函数替代手动比较,简化“取中间值和Sum_Forecast_Flex中较大者”的逻辑。
内容的提问来源于stack exchange,提问作者Saida Majbour
相关产品推荐
相关产品推荐

