如何使用PySpark为特定值'c'实现分段递增频率列?
PySpark实现分段递增列需求
需求说明
现有数据集(列名假设为col):
a c c d b a a d d c c b a b
需要新增一列new,规则如下:
- 当
col取值为'c'时,new设为0 - 连续的非
'c'行组成一个分段,各分段按出现顺序依次递增编号(示例中第一段非'c'编号1,第二段2,第三段3)
期望输出示例:
a 1 c 0 c 0 d 2 b 2 a 2 a 2 d 2 d 2 c 0 c 0 b 3 a 3 b 3
问题代码(未生效)
用户尝试的代码如下,但未达到预期效果:
from pyspark.sql.functions import col, when, lag, sum s = df.filter(col("col") == 'c') df = df.withColumn("new", when(s.neq(lag("s", 1).over()), sum("s").over(Window.orderBy("index"))).otherwise(0))
解决方案
要实现这个需求,核心是识别非'c'分段的起始点,再对分段进行编号。具体步骤如下:
1. 准备工作:添加行号(若已有可跳过)
PySpark需要明确的排序依据,先给数据集添加index列作为行号:
from pyspark.sql import SparkSession spark_session = SparkSession.builder.getOrCreate() df = df.withColumn("index", spark_session.sparkContext.range(df.count()).toDF("index"))
2. 标记非'c'分段的起始行
通过lag函数判断当前行是否是新分段的开始:当前行不是'c',且前一行是'c',或是第一行且非'c',标记为1,否则为0:
from pyspark.sql.functions import col, when, lag, sum as spark_sum from pyspark.sql.window import Window window_order = Window.orderBy("index") df = df.withColumn( "new_block", when( (col("col") != "c") & (lag(col("col"), 1).over(window_order).isNull() | (lag(col("col"), 1).over(window_order) == "c")), 1 ).otherwise(0) )
3. 计算分段编号
对new_block列进行累加,得到每个非'c'分段的编号:
df = df.withColumn("block_id", spark_sum(col("new_block")).over(window_order))
4. 生成最终的new列
根据规则设置new列:'c'行取0,非'c'行取对应的分段编号:
df = df.withColumn( "new", when(col("col") == "c", 0).otherwise(col("block_id")) ) # 可选:删除中间生成的辅助列 df = df.drop("index", "new_block", "block_id")
执行以上代码后,就能得到符合要求的结果。
内容的提问来源于stack exchange,提问作者Akshat Srivastav
相关产品推荐
相关产品推荐

