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

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分区,且循环拆分合并的操作触发了惰性求值的执行计划合并,导致滞后值跨组计算。

问题根源

  1. 定义的window_spec仅做全局排序,未通过partitionBy指定按ts_id分组,即使先过滤单个ts_id,PySpark惰性求值会将整个执行计划合并,最终窗口函数仍在全局数据上计算,取到其他ts_id的滞后值。
  2. 循环拆分再合并的方式既低效,又会引发执行计划的意外合并,放大惰性求值的问题。

正确实现

直接在原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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.30 10:03:09