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

PySpark按分区统计非空集群数并生成集群标识列

为PySpark DataFrame中连续非空行集群生成分区内标识列

问题场景

给定如下PySpark DataFrame:

df1_l = [
  (0, 1),
  (0, 2),
  (0, 3),
  (0, 4),
  (0, None),
  (0, None),
  (0, None),
  (0, 801),
  (0, 802),
  (0, 803),
  (0, None),
  (0, None),
  (1, 1),
  (1, 2),
  (1, 3),
  (1, 4),
  (1, None),
  (1, None),
  (1, None),
  (1, 801),
  (1, 802),
  (1, 803),
  (1, None),
  (1, None)
]

df1 = spark.createDataFrame(df1_l, schema = ["id", "val"])
df1.show()

需求:在每个id分区内,为val列的连续非空行集群生成统一整数标识列n_cluster,空行对应null;连续非空行(含单个非空行)构成一个集群,预期输出如题目所示。

解决方案

通过PySpark窗口函数实现,核心思路是:标记集群起始点,再对起始点累加得到集群编号。

完整代码实现

from pyspark.sql import functions as F
from pyspark.sql.window import Window

# 定义窗口:按id分区,用monotonically_increasing_id保证原输入顺序
window = Window.partitionBy("id").orderBy(F.monotonically_increasing_id())

# 生成n_cluster列
result_df = df1.withColumn("flag", F.when(F.col("val").isNotNull(), 1).otherwise(0)) \
               .withColumn("group_start", 
                           F.when((F.col("flag") == 1) & (F.lag("flag", 1, 0).over(window) == 0), 1).otherwise(0)) \
               .withColumn("n_cluster", 
                           F.when(F.col("flag") == 1, F.sum("group_start").over(window.rangeBetween(Window.unboundedPreceding, 0)))) \
               .drop("flag", "group_start")

result_df.show()

步骤解释

  1. 标记非空行:创建flag列,非空行标记为1,空行标记为0,用于区分空值与非空行。
  2. 识别集群起始点:创建group_start列,当当前行是非空行(flag=1)且前一行是空行(或为分区第一行)时,标记为1,表示这是一个新集群的起始。
  3. 生成集群编号:在每个id分区内,对group_start列做累加求和,非空行的累加结果就是对应的集群编号,空行则设为null。
  4. 清理中间列:删除用于计算的flag和group_start临时列,得到最终结果。

说明

  • 若数据本身有明确的排序字段(如时间戳、序列ID),可将orderBy(F.monotonically_increasing_id())替换为该字段,保证排序逻辑符合业务需求。
  • 该方案支持任意数量的连续非空集群,无需提前指定集群数。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.25 14:15:41