PySpark中移除DataFrame指定列的连续重复值
解决PySpark DataFrame多列连续重复值移除问题
要移除[id, st]列的连续重复并保留对应最早日期(同日期随机选)的记录,我们可以借助窗口函数来实现,核心思路是对比当前行与前一行的id和st组合,只保留组合不同的行(或第一行)。以下是具体实现步骤:
步骤1:导入所需函数
首先导入PySpark的窗口函数和相关操作函数:
from pyspark.sql import Window from pyspark.sql.functions import lag, col, struct, when, rand
步骤2:定义排序窗口
我们需要先按date排序(同日期时加入rand()保证随机顺序),确保连续重复的行是相邻的,并且最早/随机的行排在前面:
# 按date排序,同日期随机打乱顺序 window_spec = Window.orderBy("date", rand())
步骤3:添加前一行的对比列
使用lag函数获取前一行的id和st组合,然后和当前行的组合做对比:
# 生成前一行的(id, st)结构体,用于对比 df_with_prev = test_df.withColumn( "prev_id_st", lag(struct("id", "st")).over(window_spec) )
步骤4:过滤保留目标行
标记需要保留的行:第一行(prev_id_st为null),或者当前行与前一行的id+st组合不同的行,然后过滤出这些行:
df_result = df_with_prev.withColumn( "keep_row", when( col("prev_id_st").isNull() | (struct("id", "st") != col("prev_id_st")), True ).otherwise(False) ).filter(col("keep_row")).drop("prev_id_st", "keep_row")
验证结果
执行完上述代码后,df_result的输出就会符合你的预期:
| id | num | st | date |
|---|---|---|---|
| 2 | 3.0 | a | 2020-01-01 |
| 3 | 2.0 | a | 2020-01-02 |
| 4 | 1.0 | b | 2020-01-04 |
| 2 | 3.0 | a | 2020-01-08 |
| 4 | 7.0 | b | 2020-01-09 |
替代方案(不使用结构体)
如果你不想用结构体组合列,也可以分别获取前一行的id和st来对比:
df_with_prev = test_df.withColumn( "prev_id", lag("id").over(window_spec) ).withColumn( "prev_st", lag("st").over(window_spec) ) df_result = df_with_prev.withColumn( "keep_row", when( col("prev_id").isNull() | (col("id") != col("prev_id")) | (col("st") != col("prev_st")), True ).otherwise(False) ).filter(col("keep_row")).drop("prev_id", "prev_st")
这个方案和之前的逻辑完全一致,只是拆分了对比条件,适合更直观的理解。
内容的提问来源于stack exchange,提问作者eljiwo
相关产品推荐
相关产品推荐

