PySpark中基于多列移除DataFrame连续重复行的方法
在PySpark中按Unit分组移除状态连续重复行的实现方法
问题描述
现有一张事件表,包含timestamp、unit、state 1、state n等列,示例数据如下:
- 01:00时,unit a的state 1为x
- 02:00时,unit a的state 1仍为x
- 03:00时,unit a的state 1变为y
- 04:00时,unit a的state 1又变回x
- 05:00时,unit b的state 1为x
需求:按unit分组,基于指定状态列(单个如state 1或多个状态列)移除连续重复行,仅保留状态变更后的行(如04:00的行),避免普通去重丢失状态变更的历史信息。
实现方案
核心逻辑是利用PySpark的窗口函数和lag函数,按unit分组并按timestamp排序,对比当前行与前一行的指定状态值,仅保留状态变化的行(含每组首行)。
1. 单状态列场景(以state 1为例)
from pyspark.sql import SparkSession from pyspark.sql.window import Window from pyspark.sql.functions import lag, col # 初始化SparkSession spark = SparkSession.builder.appName("RemoveConsecutiveDuplicates").getOrCreate() # 构建示例数据 data = [ ("01:00", "a", "x"), ("02:00", "a", "x"), ("03:00", "a", "y"), ("04:00", "a", "x"), ("05:00", "b", "x") ] df = spark.createDataFrame(data, ["timestamp", "unit", "state 1"]) # 定义窗口:按unit分组,按时间升序排序 window_spec = Window.partitionBy("unit").orderBy("timestamp") # 添加前一行的state值,过滤连续重复行 df_with_prev = df.withColumn("prev_state", lag(col("state 1")).over(window_spec)) result_df = df_with_prev.filter( col("prev_state").isNull() | (col("state 1") != col("prev_state")) ).drop("prev_state") result_df.show()
执行结果:
+---------+----+-------+ |timestamp|unit|state 1| +---------+----+-------+ | 01:00| a| x| | 03:00| a| y| | 04:00| a| x| | 05:00| b| x| +---------+----+-------+
2. 多状态列场景(如state 1+state 2)
若需基于多个状态列判断连续重复,可将目标列打包为结构体,对比结构体是否相等:
# 扩展示例数据,新增state 2列 data_multi = [ ("01:00", "a", "x", "m"), ("02:00", "a", "x", "m"), ("03:00", "a", "y", "m"), ("04:00", "a", "x", "n"), ("05:00", "b", "x", "m") ] df_multi = spark.createDataFrame(data_multi, ["timestamp", "unit", "state 1", "state 2"]) # 指定需要判断的状态列集合 state_cols = ["state 1", "state 2"] # 定义窗口 window_spec_multi = Window.partitionBy("unit").orderBy("timestamp") # 组合状态列为结构体,获取前一行的结构体并过滤重复 df_with_prev_multi = df_multi.withColumn( "current_state", struct(*state_cols) ).withColumn( "prev_state", lag(col("current_state")).over(window_spec_multi) ) result_multi_df = df_with_prev_multi.filter( col("prev_state").isNull() | (col("current_state") != col("prev_state")) ).drop("current_state", "prev_state") result_multi_df.show()
执行结果:
+---------+----+-------+-------+ |timestamp|unit|state 1|state 2| +---------+----+-------+-------+ | 01:00| a| x| m| | 03:00| a| y| m| | 04:00| a| x| n| | 05:00| b| x| m| +---------+----+-------+-------+
核心说明
- 窗口函数:
partitionBy("unit")确保同设备的行被分组,orderBy("timestamp")保证行按时间顺序排列,让lag能正确获取上一条记录的状态。 - lag函数:默认偏移量为1,用于获取分组内前一行的指定列值,是判断连续重复的核心。
- 多列处理:通过
struct(*state_cols)将多个状态列打包成整体,简化多列对比逻辑,只要结构体内容不同,即判定为状态变更。
内容的提问来源于stack exchange,提问作者Adam Andersson
相关产品推荐
相关产品推荐

