PySpark DataFrame实现获取上一个不同值的PREVIOUSPRICE列
PySpark实现获取上一个不同的PRICE值(按ID分组、排序)
需求说明
按ID分组,按DateCOL和PRICE排序,新增PREVIOUSPRICE列:
- 取当前行上一个不同的PRICE值
- 若当前PRICE与前一行相同,则继续向上查找直到找到不同值
- 组内第一行的
PREVIOUSPRICE为null
解决方案代码
from pyspark.sql import SparkSession from pyspark.sql import functions as F from pyspark.sql.window import Window # 初始化SparkSession spark = SparkSession.builder.appName("previous_price_calculation").getOrCreate() # 构建测试DataFrame data = [ (1, "20240301", 10), (1, "20240301", 10), (1, "20240302", 20), (1, "20240303", 30), (1, "20240304", 40), (1, "20240305", 50), (1, "20240305", 50), (1, "20240305", 60), (1, "20240305", 60), (1, "20240306", 70), (2, "20240306", 33), (2, "20240307", 44) ] df = spark.createDataFrame(data, ["ID", "DateCOL", "PRICE"]) # 步骤1:划分连续相同PRICE的分组 # 定义排序窗口:按ID分区,DateCOL、PRICE排序 order_window = Window.partitionBy("ID").orderBy("DateCOL", "PRICE") # 标记PRICE变化的位置:当前行与前一行PRICE不同则标记为1,否则0 df_with_flag = df.withColumn( "price_change", F.when(F.lag("PRICE").over(order_window) != F.col("PRICE"), 1).otherwise(0) ) # 累积求和生成分组ID,连续相同PRICE的行归为同一组 df_with_group = df_with_flag.withColumn( "group_id", F.sum("price_change").over(order_window.rowsBetween(Window.unboundedPreceding, 0)) ) # 步骤2:获取每个分组的上一个不同PRICE # 定义分组窗口:按ID分区,group_id排序 group_window = Window.partitionBy("ID").orderBy("group_id") # 取前一个分组的PRICE作为当前分组的PREVIOUSPRICE df_with_prev = df_with_group.withColumn( "PREVIOUSPRICE", F.lag("PRICE").over(group_window) ) # 将分组内的PREVIOUSPRICE统一填充(确保同一组所有行值一致) final_df = df_with_prev.withColumn( "PREVIOUSPRICE", F.first("PREVIOUSPRICE").over(Window.partitionBy("ID", "group_id")) ).select("ID", "DateCOL", "PRICE", "PREVIOUSPRICE") # 查看结果 final_df.orderBy("ID", "DateCOL", "PRICE").show()
结果展示
+---+--------+-----+-------------+ | ID| DateCOL|PRICE|PREVIOUSPRICE| +---+--------+-----+-------------+ | 1|20240301| 10| null| | 1|20240301| 10| null| | 1|20240302| 20| 10| | 1|20240303| 30| 20| | 1|20240304| 40| 30| | 1|20240305| 50| 40| | 1|20240305| 50| 40| | 1|20240305| 60| 50| | 1|20240305| 60| 50| | 1|20240306| 70| 60| | 2|20240306| 33| null| | 2|20240307| 44| 33| +---+--------+-----+-------------+
逻辑说明
- 分组标记:通过
lag判断PRICE是否变化,结合累积求和生成group_id,把连续相同的PRICE行归为同一组。 - 获取上一个不同值:在分组维度上用
lag取前一组的PRICE,再通过first函数将该值广播到当前组的所有行,确保同一组内的PREVIOUSPRICE一致。
内容的提问来源于stack exchange,提问作者user28823504
相关产品推荐
相关产品推荐

