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

如何在PySpark中实现SAS Retain语句的功能?

SAS Retain语句迁移到PySpark的解决方案

原SAS代码

data ds;
set  ds;
 
    by group date;
   
    retain Target; 
 
    if first.group                             then Target = Orig;
    if first.group and ( Orig in (1,2,3,4,5) ) then Target = 6;
 
    if not first.group and   Target  = 6 and (Orig in (1,2,3,4,5) )   then Target = 6 ;
    if not first.group and ~(Target  = 6 and (Orig in (1,2,3,4,5) ) ) then Target = Orig ;
   
    run;

核心逻辑规则

  • 组内第一条记录:
    • 若Orig值属于{1,2,3,4,5},Target设为6
    • 否则Target等于当前Orig值
  • 非组内第一条记录:
    • 如果前一行Target为6,且当前Orig属于{1,2,3,4,5},保持Target为6
    • 其他情况(前一行Target非6,或当前Orig不在{1,2,3,4,5}内),Target设为当前Orig值

PySpark实现方案

由于SAS的retain是逐行维护状态的逻辑,PySpark中可以通过groupBy+applyInPandas实现分组内的逐行迭代处理,逻辑和SAS完全对齐。

步骤1:导入依赖

from pyspark.sql import SparkSession
from pyspark.sql.types import StructType, StructField, IntegerType
import pandas as pd

步骤2:定义分组处理函数

该函数负责按组遍历每一行数据,维护Target的状态:

def process_group(pdf: pd.DataFrame) -> pd.DataFrame:
    # 严格按date排序,对应SAS的by group date逻辑
    pdf = pdf.sort_values("date")
    target_list = []
    prev_target = None
    
    for idx, row in pdf.iterrows():
        if idx == 0:  # 处理组内第一条记录
            if row["Orig"] in {1, 2, 3, 4, 5}:
                current_target = 6
            else:
                current_target = row["Orig"]
        else:  # 处理组内后续记录
            if prev_target == 6 and row["Orig"] in {1, 2, 3, 4, 5}:
                current_target = 6
            else:
                current_target = row["Orig"]
        
        target_list.append(current_target)
        prev_target = current_target
    
    pdf["Target"] = target_list
    return pdf

步骤3:定义输出Schema

指定applyInPandas的返回结构:

output_schema = StructType([
    StructField("Group", IntegerType(), nullable=False),
    StructField("date", IntegerType(), nullable=False),
    StructField("Orig", IntegerType(), nullable=False),
    StructField("Target", IntegerType(), nullable=False)
])

步骤4:执行分组处理

# 初始化SparkSession
spark = SparkSession.builder.appName("SASRetainMigration").getOrCreate()

# 假设原始DataFrame为df,包含Group、date、Orig三列
result_df = df.groupBy("Group").applyInPandas(process_group, schema=output_schema)

示例验证

使用提供的示例数据(补充date列以保证排序),可以验证结果与预期一致:

# 构造带date的示例数据
sample_data = [
(999, 1,5,6) ,(999, 2,6,6) ,(999, 3,4,6) ,(999, 4,6,6) ,(999, 5,3,6) ,
(999, 6,5,6) ,(999, 7,4,6) ,(999, 8,6,6) ,(999, 9,6,6) ,(999,10,6,6) ,
(999,11,6,6) ,(999,12,5,6) ,(999,13,3,6) ,(999,14,2,6) ,(999,15,2,6) ,
(999,16,2,6) ,(999,17,2,6) ,(999,18,2,6) ,(999,19,2,6) ,(999,20,2,6) ,
(999,21,2,6) ,(999,22,1,6) ,(999,23,0,0) ,(999,24,0,0) ,(999,25,0,0) ,
(999,26,0,0) ,(999,27,1,1) ,(999,28,1,1) ,(999,29,2,2) ,(999,30,2,2) ,
(999,31,3,3) ,(999,32,2,2) ,(999,33,3,3) ,(999,34,4,4) ,(999,35,5,5) ,
(999,36,6,6) ,(999,37,6,6) ,(999,38,6,6) ,(999,39,0,0) ,(999,40,1,1) ,
(999,41,0,0) ,(999,42,1,1) ,(999,43,2,2) ,(999,44,3,3) ,(999,45,4,4) ,
(999,46,5,5) ,(999,47,6,6) ,(999,48,6,6) ,(999,49,6,6) ,(999,50,6,6) ,
(999,51,4,6) ,(999,52,3,6) ,(999,53,2,6) ,(999,54,3,6) ,(999,55,4,6) ,
(999,56,5,6) ,(999,57,6,6) ,(999,58,6,6)
]

# 创建示例DataFrame
df = spark.createDataFrame(sample_data, ['Group', 'date', 'Orig', 'Expected_Target'])

# 执行处理
result_df = df.groupBy("Group").applyInPandas(process_group, schema=output_schema)

# 对比结果
result_df.join(df.select("Group", "date", "Expected_Target"), on=["Group", "date"]).show()

注意事项

  • 必须保证分组内按date排序,SAS的by group date会自动排序,PySpark中需手动执行排序逻辑,否则结果会出错
  • applyInPandas要求Spark版本≥3.0,若使用低版本Spark,可改用RDD的mapPartitions实现,但逻辑复杂度会更高
  • 逻辑严格对齐SAS的执行顺序:组内第一条记录的两个if语句是顺序执行的,最终Target以第二个if的结果为准

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.25 21:37:05