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()
步骤解释
- 标记非空行:创建
flag列,非空行标记为1,空行标记为0,用于区分空值与非空行。 - 识别集群起始点:创建
group_start列,当当前行是非空行(flag=1)且前一行是空行(或为分区第一行)时,标记为1,表示这是一个新集群的起始。 - 生成集群编号:在每个
id分区内,对group_start列做累加求和,非空行的累加结果就是对应的集群编号,空行则设为null。 - 清理中间列:删除用于计算的
flag和group_start临时列,得到最终结果。
说明
- 若数据本身有明确的排序字段(如时间戳、序列ID),可将
orderBy(F.monotonically_increasing_id())替换为该字段,保证排序逻辑符合业务需求。 - 该方案支持任意数量的连续非空集群,无需提前指定集群数。
内容的提问来源于stack exchange,提问作者jj_coder
相关产品推荐
相关产品推荐

