PySpark中创建滞后列并合并多个DataFrame的正确方法
问题描述
我尝试为多个DataFrame分别创建滞后列,再将它们合并为单个DataFrame。但由于PySpark是惰性求值的,实际是在合并DataFrame后才计算滞后列,导致结果不符合预期。
我的代码如下:
from pyspark.sql.functions import lag, col from pyspark.sql.window import Window from pyspark.sql.types import StructType, StructField, StringType, IntegerType, FloatType # Define the schema for the DataFrame schema = StructType([ StructField("YEARWEEK", StringType(), True), StructField("CATEGORY", StringType(), True), StructField("vendor_nbr", StringType(), True), StructField("dc_nbr", StringType(), True), StructField("rejection_rate", FloatType(), True), StructField("rating", IntegerType(), True), StructField("ts_id", StringType(), True) , StructField("rating_lag1", IntegerType(), True), StructField("rating_lag2", IntegerType(), True), StructField("rating_lag3", IntegerType(), True), StructField("rating_lag4", IntegerType(), True) ]) # Create an empty DataFrame with the specified schema df_comb = spark.createDataFrame([], schema) window_spec = Window.orderBy("ts_id", "YEARWEEK") dfs = [] for id in df.select(col("ts_id")).distinct(): dfs.append(df.filter(df.ts_id == id).sort('ts_id', 'YEARWEEK').withColumn("rating_lag1", lag("rating", 1).over(window_spec))\ .withColumn("rating_lag2", lag("rating", 2).over(window_spec))\ .withColumn("rating_lag3", lag("rating", 3).over(window_spec))\ .withColumn("rating_lag4", lag("rating", 4).over(window_spec))\ .na.fill({'rating_lag1': 1, 'rating_lag2': 1, 'rating_lag3': 1, 'rating_lag4': 1})) for dfx in dfs: df_comb = df_comb.union(dfx)
为了说明问题,我执行了以下代码:
dfp = df_comb.toPandas() dfp.iloc[1305:1315]
我的数据从2022年第一周开始,理想情况下lag3应为null,我将null值替换为1后,lag3应该是1,但实际显示为5,结果不正确。请问在PySpark中,合并带有滞后列的DataFrame的正确方法是什么?
解决方案
核心问题是窗口函数未按ts_id分区,且循环拆分合并的操作触发了惰性求值的执行计划合并,导致滞后值跨组计算。
问题根源
- 定义的
window_spec仅做全局排序,未通过partitionBy指定按ts_id分组,即使先过滤单个ts_id,PySpark惰性求值会将整个执行计划合并,最终窗口函数仍在全局数据上计算,取到其他ts_id的滞后值。 - 循环拆分再合并的方式既低效,又会引发执行计划的意外合并,放大惰性求值的问题。
正确实现
直接在原DataFrame上使用带分区的窗口函数,一次性完成所有计算,无需拆分合并:
from pyspark.sql.functions import lag, col from pyspark.sql.window import Window # 按ts_id分区,同一组内按YEARWEEK排序 window_spec = Window.partitionBy("ts_id").orderBy("YEARWEEK") # 计算滞后列并填充空值 df_comb = df.withColumn("rating_lag1", lag("rating", 1).over(window_spec))\ .withColumn("rating_lag2", lag("rating", 2).over(window_spec))\ .withColumn("rating_lag3", lag("rating", 3).over(window_spec))\ .withColumn("rating_lag4", lag("rating", 4).over(window_spec))\ .na.fill({'rating_lag1': 1, 'rating_lag2': 1, 'rating_lag3': 1, 'rating_lag4': 1}) # 可选:按ts_id和YEARWEEK排序结果 df_comb = df_comb.orderBy("ts_id", "YEARWEEK")
效果说明
partitionBy("ts_id")保证每个ts_id组内独立计算滞后值,不会跨组取值,符合业务逻辑。- 无需拆分DataFrame,直接在原数据上操作,避免了惰性求值导致的执行计划合并问题。
- 执行效率远高于循环拆分合并,减少了多次过滤、排序和union的开销。
内容的提问来源于stack exchange,提问作者Vinay
相关产品推荐
相关产品推荐

