PySpark DataFrame如何获取满足指定条件的首个前置行数据
解决PySpark DataFrame获取最近前置符合条件行的值的问题
这个需求确实没法用普通的lag()函数搞定——毕竟lag()只能取固定偏移量的行,而我们需要动态找到当前行之前最近的Flag=1的行的X值。这里有个简洁的实现方案,用窗口函数结合last()就能完美解决:
核心思路
- 先创建一个辅助列,只在
Flag=1的行保留X的值,其他行设为Null; - 定义一个覆盖从数据集开头到当前行的前一行的窗口;
- 在这个窗口里用
last()函数(忽略Null值),就能自动抓取到最近的那个Flag=1行的X值,没有前置符合条件行时就返回Null。
完整代码示例
from pyspark.sql import SparkSession from pyspark.sql import functions as F from pyspark.sql.window import Window # 初始化SparkSession spark = SparkSession.builder.appName("nearest_flag_example").getOrCreate() # 创建示例DataFrame data = [(1,1), (2,0), (3,0), (4,0), (5,1), (6,0), (7,0), (8,0), (9,1), (10,0)] df = spark.createDataFrame(data, ["X", "Flag"]) # 步骤1:创建辅助列,仅Flag=1时保留X值 df_with_temp = df.withColumn("flagged_X", F.when(F.col("Flag") == 1, F.col("X")).otherwise(F.lit(None))) # 步骤2:定义窗口范围:从起始行到当前行的前一行 window_spec = Window.orderBy("X").rowsBetween(Window.unboundedPreceding, -1) # 步骤3:计算Lag_X,取窗口内最后一个非Null的flagged_X result_df = df_with_temp.withColumn("Lag_X", F.last("flagged_X", ignorenulls=True).over(window_spec)) \ .drop("flagged_X") # 查看结果 result_df.show()
运行结果
+---+----+-----+ | X|Flag|Lag_X| +---+----+-----+ | 1| 1| null| | 2| 0| 1| | 3| 0| 1| | 4| 0| 1| | 5| 1| 1| | 6| 0| 5| | 7| 0| 5| | 8| 0| 5| | 9| 1| 5| | 10| 0| 9| +---+----+-----+
扩展:如果需要按分组计算
如果你的数据需要按某列分组(比如每个分组内单独计算最近的Flag=1行),只需要在窗口里加上partitionBy()即可:
# 假设按group_col分组 window_spec = Window.partitionBy("group_col").orderBy("X").rowsBetween(Window.unboundedPreceding, -1)
这样每个分组内会独立计算,互不干扰。
内容的提问来源于stack exchange,提问作者NME IX
相关产品推荐
相关产品推荐

