PySpark同一列递归操作:实现带条件重置规则的drift_MS列计算
PySpark 实现重置式累加drift_MS列方案
核心思路
你需要的带重置条件的逐行累加,可以通过「标记重置点+分组内累加」的标准Spark SQL模式实现,不需要逐行迭代也能完成逻辑:
- 先计算每一行对应的增量值(符合条件加1/100,不符合则为0,0即为重置标记)
- 对重置点做累计计数生成组ID,连续未触发重置的行会被划分到同一个组
- 每个组内对增量值做累加,就是最终的
drift_MS结果
完整实现代码
from pyspark.sql import functions as f from pyspark.sql.window import Window from pyspark.sql.functions import col # 按SEQ_ID分区、时间戳排序的基础窗口 w_base = Window.partitionBy('SEQ_ID').orderBy(col('TIME_STAMP').asc()) # 分组累加用的窗口,分区维度是SEQ_ID+重置组ID w_group = Window.partitionBy('SEQ_ID', 'reset_group').orderBy(col('TIME_STAMP').asc()) df = df.withColumn('prev_MS', f.lag(col('_MS'), 1).over(w_base)) \ # 计算当前行增量 .withColumn('incr', f.when((col('_MS') < 3) & (col('prev_MS') < col('_MS')), 1) .when((col('_MS') >= 3) & (col('prev_MS') < col('_MS')), 100) .otherwise(0) ) \ # 标记重置组:遇到增量为0时组ID加1,同一个组内是连续需要累加的行 .withColumn('reset_flag', f.when(col('incr') == 0, 1).otherwise(0)) \ .withColumn('reset_group', f.sum(col('reset_flag')).over(w_base)) \ # 组内累加得到最终结果 .withColumn('drift_MS', f.sum(col('incr')).over(w_group)) \ # 清理中间临时列 .drop('prev_MS', 'incr', 'reset_flag', 'reset_group')
效果说明
上述代码完全基于Spark原生列操作实现,支持分布式大数据量运行,输出的drift_MS列完全匹配你给出的预期结果。
内容的提问来源于stack exchange,提问作者thentangler
相关产品推荐
相关产品推荐

