如何在PySpark DataFrame中按条件添加Flag与Part列
PySpark DataFrame 分组计算与列生成实现方案
原始数据
id vehicle production asIs EU EU_variant status 1 A3345 PQ1298 FV1 FV1_variant OK 2 A3346 A3346 PQ1287 FV2 FV2_variant NOT_OK 3 A3346 A3346 PQ1207 FV2 FV2_variant NOT_OK 4 A3347 QP9 QP9_variant OK 5 A3347 QP9 QP9_variant NOT_OK 6 A3347 QP3 QP3_variant OK 7 A3348 MP6553 YR34 YR34_variant NOT_OK 8 A3348 MP6554 YR35 YR35_variant NOT_OK 9 A3348 MP6554 YR35 YR35_variant NOT_OK
需求说明
- 分组规则:按
vehicle与EU分组;若vehicle为空,则按production与EU分组 - 生成
Flag列:分组内同时存在OK和NOT_OK状态时,Flag为0;仅存在NOT_OK时,Flag为1 - 生成
Part列:将分组内的asIs值去重后用逗号拼接
期望输出
id vehicle production asIs EU EU_variant status Flag Part 1 A3345 PQ1298 FV1 FV1_variant OK 0 PQ1298 2 A3346 A3346 PQ1287 FV2 FV2_variant NOT_OK 1 PQ1287,PQ1207 3 A3346 A3346 PQ1207 FV2 FV2_variant NOT_OK 1 PQ1287,PQ1207 4 A3347 QP9 QP9_variant OK 0 5 A3347 QP9 QP9_variant NOT_OK 0 6 A3347 QP3 QP3_variant OK 0 7 A3348 MP6553 YR34 YR34_variant NOT_OK 1 MP6553 8 A3348 MP6554 YR35 YR35_variant NOT_OK 1 MP6554 9 A3348 MP6554 YR35 YR35_variant NOT_OK 1 MP6554
实现代码
通过PySpark窗口函数与分组聚合即可完成需求,具体代码如下:
from pyspark.sql import SparkSession from pyspark.sql import functions as F from pyspark.sql.window import Window # 初始化SparkSession spark = SparkSession.builder.appName("VehicleGrouping").getOrCreate() # 构建原始DataFrame(实际场景可替换为数据源读取逻辑) data = [ (1, "A3345", None, "PQ1298", "FV1", "FV1_variant", "OK"), (2, "A3346", "A3346", "PQ1287", "FV2", "FV2_variant", "NOT_OK"), (3, "A3346", "A3346", "PQ1207", "FV2", "FV2_variant", "NOT_OK"), (4, None, "A3347", None, "QP9", "QP9_variant", "OK"), (5, None, "A3347", None, "QP9", "QP9_variant", "NOT_OK"), (6, None, "A3347", None, "QP3", "QP3_variant", "OK"), (7, None, "A3348", "MP6553", "YR34", "YR34_variant", "NOT_OK"), (8, None, "A3348", "MP6554", "YR35", "YR35_variant", "NOT_OK"), (9, None, "A3348", "MP6554", "YR35", "YR35_variant", "NOT_OK") ] schema = ["id", "vehicle", "production", "asIs", "EU", "EU_variant", "status"] df = spark.createDataFrame(data, schema) # 定义动态分组键:vehicle非空时用vehicle+EU,否则用production+EU group_key = F.when(F.col("vehicle").isNotNull(), F.col("vehicle")).otherwise(F.col("production")) window_spec = Window.partitionBy(group_key, F.col("EU")) # 计算Flag列:判断分组内状态组合 df = df.withColumn( "status_set", F.collect_set(F.col("status")).over(window_spec) ).withColumn( "Flag", F.when( F.array_contains(F.col("status_set"), "OK") & F.array_contains(F.col("status_set"), "NOT_OK"), 0 ).when( F.size(F.col("status_set")) == 1 & F.array_contains(F.col("status_set"), "NOT_OK"), 1 ).otherwise(0) ) # 计算Part列:分组内asIs去重后拼接 df = df.withColumn( "Part", F.concat_ws(",", F.collect_set(F.col("asIs")).over(window_spec)) ) # 清理中间列并调整列顺序 final_df = df.drop("status_set").select( "id", "vehicle", "production", "asIs", "EU", "EU_variant", "status", "Flag", "Part" ) # 查看结果 final_df.show(truncate=False)
代码说明
- 动态分组:用
when函数根据vehicle是否为空切换分组字段,满足需求中的分组规则 - Flag列生成:通过
collect_set收集分组内所有状态,再用array_contains判断状态组合,生成对应Flag值 - Part列生成:利用
collect_set对asIs去重,再通过concat_ws将集合元素拼接为字符串 - 窗口函数优势:无需额外关联操作,直接在原表上完成分组聚合,保证每一行都能匹配到分组计算结果
内容的提问来源于stack exchange,提问作者karthik
相关产品推荐
相关产品推荐

