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

PySpark实现SAS Retain逻辑:计算estimate_day_to_sustain列

PySpark实现SAS RETAIN风格的estimate_day_to_sustain计算

需求说明

需计算补货前的estimate_day_to_sustain字段,核心规则:

  • 按日期排序后的第一行(第1天):estimate_day_to_sustain = 当前日期
  • 后续行:
    • 若上一行estimate_day_to_sustain + 上一行supply ≤ 当前日期,则当前值为当前日期
    • 否则当前值 = 上一行estimate_day_to_sustain + 上一行supply

实现方案

由于该计算依赖前一行的输出结果,PySpark中推荐用递归CTE模拟SAS的retain迭代逻辑,步骤如下:

1. 预处理数据

先对数据按日期排序并添加行号,用于递归遍历:

from pyspark.sql import SparkSession, functions as F
from pyspark.sql.window import Window

spark = SparkSession.builder.appName("SustainDayCalc").getOrCreate()

# 假设输入DataFrame为df,包含date(日期)、supply(供应量)字段
sort_window = Window.orderBy("date")
df_rn = df.withColumn("row_num", F.row_number().over(sort_window))

2. 递归CTE计算

通过递归CTE逐行迭代计算目标字段:

# 注册临时表供SQL使用
df_rn.createOrReplaceTempView("sustain_data")

# 递归CTE逻辑
result_df = spark.sql("""
WITH RECURSIVE calc_cte AS (
    -- 基准行:初始化第一行的estimate值
    SELECT date, supply, row_num, date AS estimate_day_to_sustain
    FROM sustain_data
    WHERE row_num = 1
    
    UNION ALL
    
    -- 递归迭代后续行
    SELECT 
        curr.date,
        curr.supply,
        curr.row_num,
        CASE
            WHEN prev.estimate_day_to_sustain + prev.supply <= curr.date THEN curr.date
            ELSE prev.estimate_day_to_sustain + prev.supply
        END AS estimate_day_to_sustain
    FROM calc_cte prev
    JOIN sustain_data curr ON prev.row_num + 1 = curr.row_num
)
SELECT date, supply, estimate_day_to_sustain FROM calc_cte ORDER BY date
""")

# 查看结果
result_df.show()

3. 关键场景验证

对应给出的算法说明:

  • 第3天:前一行结果1+3=4,当前日期3,4>3,故estimate_day_to_sustain=4
  • 第9天:前一行结果4+1=5,当前日期9,5≤9,故estimate_day_to_sustain=9
  • 后续行按规则迭代即可得到正确输出

性能优化方案(可选)

如果数据量极大,递归CTE性能不足,可尝试用窗口函数结合lag函数(仅适用于连续无跳行的日期场景):

window_spec = Window.orderBy("row_num")

df_result = df_rn.withColumn(
    "prev_estimate", F.lag("estimate_day_to_sustain", 1).over(window_spec)
).withColumn(
    "prev_supply", F.lag("supply", 1).over(window_spec)
).withColumn(
    "estimate_day_to_sustain",
    F.when(F.col("row_num") == 1, F.col("date"))
     .when(F.col("prev_estimate") + F.col("prev_supply") <= F.col("date"), F.col("date"))
     .otherwise(F.col("prev_estimate") + F.col("prev_supply"))
)

注意:该窗口函数方法无法处理非连续的行依赖,递归CTE是通用可靠方案。


内容的提问来源于stack exchange,提问作者Bobby

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.26 03:45:49