如何在PySpark中高效创建t-1至t-30的时序滞后特征
PySpark高效生成多时序滞后特征方案
核心思路:用Spark原生窗口函数替代Pandas UDF
Spark原生窗口函数经过引擎级优化,比Pandas UDF更适配大数据场景,能避免序列化/反序列化的额外开销,实现逻辑也更简洁直接。
步骤实现
1. 预处理:补全日历(可选但关键)
如果分组内存在日期断层(比如某组缺少中间某一天的数据),直接生成滞后会导致取值错误,建议先补全所有分组的每日数据:
from pyspark.sql.functions import sequence, explode, col # 获取全局日期范围 date_range = df.agg({"date_col": "min", "date_col": "max"}).collect()[0] min_date, max_date = date_range["min(date_col)"], date_range["max(date_col)"] # 生成完整日期序列 date_df = spark.createDataFrame([(min_date, max_date)], ["min_dt", "max_dt"]) date_df = date_df.withColumn("date_col", explode(sequence(col("min_dt"), col("max_dt")))).drop("min_dt", "max_dt") # 生成所有分组-日期组合,左关联原表补全缺失值 full_df = df.select("COL_A").distinct().crossJoin(date_df) full_df = full_df.join(df, on=["COL_A", "date_col"], how="left").fillna({"Count_n": 0})
2. 定义窗口并生成滞后特征
通过row_number标记分组内的初始日期,结合lag函数生成t-1至t-30的特征,初始日期的滞后值统一设为0:
from pyspark.sql import Window from pyspark.sql.functions import row_number, lag, when # 按分组和日期排序的窗口 window = Window.partitionBy("COL_A").orderBy("date_col") # 添加行号标记分组内的初始日期 full_df = full_df.withColumn("group_rn", row_number().over(window)) # 循环生成所有滞后特征 for lag_day in range(1, 31): full_df = full_df.withColumn( f"Count_n_t-{lag_day}", when(col("group_rn") == 1, 0) .otherwise(lag(col("Count_n"), lag_day).over(window)) ) # 清理临时列 final_df = full_df.drop("group_rn")
方案优势
- 性能更优:原生窗口函数由Spark Catalyst优化,避免了Pandas UDF的序列化开销
- 逻辑清晰:直接通过SQL风格的窗口操作实现,无需编写复杂的UDF逻辑
- 鲁棒性强:补全日历的步骤确保了日期断层场景下的滞后值准确性
内容的提问来源于stack exchange,提问作者Neethu Paul
相关产品推荐
相关产品推荐

