PySpark实现分组递增计数器:生成符合cond1规则的ExpectedGroup列
PySpark实现ExpectedGroup列生成需求
需求:生成ExpectedGroup列,规则为:
- 当
cond1(df.FromState == 'O'且df.ToState == 'O')为True时,该列值保持一致; - 每当遇到
cond1为False的行时,后续符合cond1的分组值递增1。
样例DataFrame
df = spark.createDataFrame(sc.parallelize([ ['A', '2019-01-01', 'P', 'O', 2, None], ['A', '2019-01-02', 'O', 'O', 5, 1], ['A', '2019-01-03', 'O', 'O', 10, 1], ['A', '2019-01-04', 'O', 'P', 4, None], ['A', '2019-01-05', 'P', 'P', 300, None], ['A', '2019-01-06', 'P', 'O', 2, None], ['A', '2019-01-07', 'O', 'O', 5, 2], ['A', '2019-01-08', 'O', 'O', 10, 2], ['A', '2019-01-09', 'O', 'P', 4, None], ['A', '2019-01-10', 'P', 'P', 300, None], ['B', '2019-01-01', 'P', 'O', 2, None], ['B', '2019-01-02', 'O', 'O', 5, 3], ['B', '2019-01-03', 'O', 'O', 10, 3], ['B', '2019-01-04', 'O', 'P', 4, None], ['B', '2019-01-05', 'P', 'P', 300, None], ]), ['ID', 'Time', 'FromState', 'ToState', 'Hours', 'ExpectedGroup'])
已尝试的代码
# condition statement cond1 = (df.FromState == 'O') & (df.ToState == 'O') df = df.withColumn('condition', cond1.cast("int")) df = df.withColumn('conditionLead', F.lead('condition').over(Window.orderBy('ID', 'Time'))) df = df.na.fill(value=0, subset=["conditionLead"]) df = df.withColumn('finalCondition', ( (F.col('condition') == 1) & (F.col('conditionLead') == 1)).cast('int'))
Pandas可行实现方案
# working pandas option: # cond1 = ( (df.FromState == 'O') & (df.ToState == 'O') ) # df['ExpectedGroup'] = (cond1.shift(-1) & cond1).cumsum().mask(~cond1) # other working option: # cond1 = ( (df.FromState == 'O') & (df.ToState == 'O') ) # df['ExpectedGroup'] = (cond1.diff()&cond1).cumsum().where(cond1)
失败的PySpark代码
# failing here windowval = (Window.partitionBy('ID').orderBy('Time').rowsBetween(Window.unboundedPreceding, 0)) df = df.withColumn('ExpectedGroup2', F.sum(F.when(cond1, F.col('finalCondition'))).over(windowval))
正确的PySpark解决方案
核心思路是:先标记每个cond1组的起始点,再通过累加起始点的数量生成分组ID,最后只保留cond1为True的行的分组值,其余设为None。
from pyspark.sql import functions as F from pyspark.sql.window import Window # 定义核心条件 cond1 = (F.col("FromState") == "O") & (F.col("ToState") == "O") # 1. 标记当前行是否是新组的起始:当前行满足cond1,且上一行不满足cond1 window_order = Window.partitionBy("ID").orderBy("Time") df = df.withColumn( "is_new_group", F.when(cond1 & (F.lag(cond1, default=False).over(window_order) == False), 1).otherwise(0) ) # 2. 累加新组标记,生成全局分组ID(按ID分区) df = df.withColumn( "group_id", F.sum("is_new_group").over(window_order.rowsBetween(Window.unboundedPreceding, 0)) ) # 3. 只保留cond1为True的行的group_id,其余设为None df = df.withColumn( "ExpectedGroup", F.when(cond1, F.col("group_id")).otherwise(None) ) # 可选:清理临时列 df = df.drop("is_new_group", "group_id") df.show()
代码说明
- 标记新组起始:使用
lag函数获取上一行的cond1状态,当当前行满足cond1且上一行不满足时,标记为新组起始(值为1)。 - 生成分组ID:按
ID分区、Time排序,累加新组标记值,得到每个行的分组ID。 - 过滤非目标行:仅在
cond1为True时保留分组ID,其余行设为None,与样例预期一致。
内容的提问来源于stack exchange,提问作者John Stud
相关产品推荐
相关产品推荐

