如何在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
相关产品推荐
相关产品推荐

