如何在PySpark DataFrame的分组记录中应用多条件判断?
分组内基于时间重叠与Tran编号的有效性判断解决方案
需求分析
需按id分组执行以下规则:
- 校验同
id下记录的start与end时间区间是否重叠 - 连续重叠的记录归为同一组,组内
tran编号最高的标记为有效(yes) - 非重叠或无重叠组的记录直接标记为有效(
yes) - 同一
id可存在多组独立的重叠区间
解决方案代码
from pyspark.sql import SparkSession from pyspark.sql.types import StructType, StructField, StringType, IntegerType from pyspark.sql import Window from pyspark.sql.functions import col, lag, max as spark_max, when, sum as spark_sum # 初始化SparkSession spark = SparkSession.builder.appName("OverlapValidation").getOrCreate() # 原始数据与Schema data = [('A',1000,1,100), ('B',1001,0,10), ('B',1002,10,15), ('B',1003,20,22), ('B',1004,25,50), ('B',1005,50,55), ('B',1006,53,56), ('B',1007,60,100), ('C',1008,1,100) ] schema = StructType([ StructField("id",StringType(),True), StructField("tran",IntegerType(),True), StructField("start",IntegerType(),True), StructField("end",IntegerType(),True), ]) df = spark.createDataFrame(data=data,schema=schema) # 步骤1:标记当前记录与前一条是否重叠 window_order = Window.partitionBy("id").orderBy("start") df_with_overlap = df.withColumn( "is_overlap", when(col("start") <= lag(col("end")).over(window_order), 1).otherwise(0) ) # 步骤2:生成连续重叠组的标识 window_group = Window.partitionBy("id").orderBy("start").rowsBetween(Window.unboundedPreceding, 0) df_with_group = df_with_overlap.withColumn( "overlap_group", spark_sum(col("is_overlap")).over(window_group) ) # 步骤3:计算每个重叠组内的最大tran值 window_max_tran = Window.partitionBy("id", "overlap_group") df_with_max_tran = df_with_group.withColumn( "max_tran_in_group", spark_max(col("tran")).over(window_max_tran) ) # 步骤4:根据规则标记有效性 df_result = df_with_max_tran.withColumn( "valid", when( # 组内无重叠(仅单条记录) 或者 当前记录是组内tran最大值 (spark_sum(col("is_overlap")).over(window_max_tran) == 0) | (col("tran") == col("max_tran_in_group")), "yes" ).otherwise("no") ).drop("is_overlap", "overlap_group", "max_tran_in_group") # 展示结果 df_result.show(truncate=False)
输出结果
+---+----+-----+---+-----+ |id |tran|start|end|valid| +---+----+-----+---+-----+ |A |1000|1 |100|yes | |B |1001|0 |10 |no | |B |1002|10 |15 |yes | |B |1003|20 |22 |yes | |B |1004|25 |50 |no | |B |1005|50 |55 |no | |B |1006|53 |56 |yes | |B |1007|60 |100|yes | |C |1008|1 |100|yes | +---+----+-----+---+-----+
代码说明
- 重叠标记:通过窗口函数对比当前记录
start与前一条记录end,判断时间区间是否重叠 - 分组标识:对重叠标记做累积求和,将连续重叠的记录归为同一组
- 组内最大Tran:计算每个重叠组的最高
tran编号,作为有效例外的判断依据 - 有效性判断:无重叠记录直接标记有效;重叠组内仅最高
tran的记录标记有效,其余标记无效
内容的提问来源于stack exchange,提问作者hassanami
相关产品推荐
相关产品推荐

