如何在PySpark中按指定规则转换DataFrame?
PySpark实现按id分组的条件过滤
问题背景
给定一个包含id、event_name、event_value三列的PySpark DataFrame,每个id对应event_name为A、B、C的行,event_value为数值型数据。需要按照以下规则过滤数据:
- 若某个
id下event_name=A的行event_value为0,则删除该id下所有event_value=0的行,仅保留event_name=A或event_value>0的行; - 若某个
id下event_name=A的行event_value>0,则保留该id下的所有行。
示例输入DataFrame:
data = [ (1, "A", 0), (1, "B", 2), (1, "C", 0), (2, "A", 5), (2, "B", 0), (2, "C", 2), (3, "A", 7), (3, "B", 8), (3, "C", 9), ] columns = ["id", "event_name", "event_value"] df = spark.createDataFrame(data, columns)
期望输出:
+---+----------+-----------+ | id|event_name|event_value| +---+----------+-----------+ | 1| A| 0| | 1| B| 2| | 2| A| 5| | 2| B| 0| | 2| C| 2| | 3| A| 7| | 3| B| 8| | 3| C| 9| +---+----------+-----------+
实现方案
1. 导入依赖函数
首先导入PySpark所需的窗口函数和列操作工具:
from pyspark.sql import Window from pyspark.sql.functions import col, when, max
2. 添加辅助列标记A的value值
通过窗口函数,按id分组,提取每个分组中event_name=A对应的event_value,作为辅助列a_value:
window = Window.partitionBy("id") df_with_a_value = df.withColumn( "a_value", max(when(col("event_name") == "A", col("event_value"))).over(window) )
由于每个id只有一条event_name=A的记录,用max或first都能准确获取到对应值。
3. 按条件过滤数据
根据a_value的值编写过滤逻辑:
filtered_df = df_with_a_value.filter( # 情况1:A的value>0,保留所有行 (col("a_value") > 0) | # 情况2:A的value=0,仅保留A行或value>0的行 ((col("a_value") == 0) & ((col("event_name") == "A") | (col("event_value") > 0))) ).drop("a_value") # 清理辅助列
4. 验证结果
执行filtered_df.show()即可得到符合要求的输出。
逻辑说明
- 辅助列
a_value的作用是让每一行都能获取到对应id下A的value值,避免多次分组操作; - 过滤条件通过逻辑运算符组合,确保优先级正确:先判断A的value是否大于0,不满足时再筛选符合要求的行;
- 最后移除辅助列,得到干净的结果DataFrame。
内容的提问来源于stack exchange,提问作者invoro
相关产品推荐
相关产品推荐

