You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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()

代码说明

  1. 标记新组起始:使用lag函数获取上一行的cond1状态,当当前行满足cond1且上一行不满足时,标记为新组起始(值为1)。
  2. 生成分组ID:按ID分区、Time排序,累加新组标记值,得到每个行的分组ID。
  3. 过滤非目标行:仅在cond1为True时保留分组ID,其余行设为None,与样例预期一致。

内容的提问来源于stack exchange,提问作者John Stud

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.15 10:40:30