PySpark实现列值匹配生成连续日期区间(遇0拆分行)
拆分PySpark中被0中断的连续1值年份区间
原本可以通过PySpark原生函数提取值为1的年份列,得到年份数组后取整体的最小和最大值,但现在需要实现:当1出现在0之后时,拆分生成新行,得到连续的日期区间。
输入数据示例
# +---+-----+---+-----+-----+-----+-----+-----+-----+ # | a| b| id|m2000|m2001|m2002|m2003|m2004|m2005| # +---+-----+---+-----+-----+-----+-----+-----+-----+ # | a|world| 1| 0| 1| 1| 0| 0| 1| # | b|world| 2| 0| 1| 1| 1| 1| 1| # | c|world| 3| 1| 1| 0| 0| 1| 1| # +---+-----+---+-----+-----+-----+-----+-----+-----+
期望输出示例
# +---+-----+---+--------+--------+ # | a| b| id|startdate|enddate| # +---+-----+---+--------+--------- # | a|world| 1| 2001| 2002| # | a|world| 1| 2005| 2005| # | b|world| 2| 2001| 2005| # | c|world| 3| 2000| 2001| # | c|world| 3| 2004| 2005| # +---+-----+---+-----+-----+-----+
现有代码局限性
当前代码只能获取所有值为1的年份的整体最小/最大值,无法拆分被0中断的连续区间:
from pyspark.sql import functions as func data_ls = [ ("a", "world", "1", 0, 0, 1, 0, 0, 1), ("b", "world", "2", 0, 1, 0, 1, 0, 1), ("c", "world", "3", 0, 0, 0, 0, 0, 0) ] data_sdf = spark.sparkContext.parallelize(data_ls). \ toDF(['a', 'b', 'id', 'm2000', 'm2001', 'm2002', 'm2003', 'm2004', 'm2005']) yearcols = [k for k in data_sdf.columns if k.startswith('m20')] data_sdf. \ withColumn('yearcol_structs', func.array(*[func.struct(func.lit(int(c[-4:])).alias('year'), func.col(c).alias('value')) for c in yearcols] ) ). \ withColumn('yearcol_1s', func.expr('transform(filter(yearcol_structs, x -> x.value = 1), f -> f.year)') ). \ filter(func.size('yearcol_1s') >= 1). \ withColumn('year_start', func.concat(func.lit('10/10/'), func.array_min('yearcol_1s'))). \ withColumn('year_end', func.concat(func.lit('10/10/'), func.array_max('yearcol_1s'))). \ show(truncate=False)
优化方案:拆分连续区间
要实现拆分被0中断的连续1值区间,核心思路是:
- 将年份列转为包含年份和对应值的数组结构,保证年份按顺序排列
- 对每个元素计算分组标识:当当前值为1且前一个值为0时,分组标识递增,以此区分不同的连续区间
- 按分组标识聚合,得到每个区间的起始和结束年份
优化后的代码如下:
from pyspark.sql import functions as func from pyspark.sql.window import Window data_ls = [ ("a", "world", "1", 0, 1, 1, 0, 0, 1), ("b", "world", "2", 0, 1, 1, 1, 1, 1), ("c", "world", "3", 1, 1, 0, 0, 1, 1), ("d", "world", "4", 0, 0, 0, 0, 0, 0) ] data_sdf = spark.createDataFrame(data_ls, schema=['a', 'b', 'id', 'm2000', 'm2001', 'm2002', 'm2003', 'm2004', 'm2005']) # 提取年份列并按年份排序 yearcols = sorted([k for k in data_sdf.columns if k.startswith('m20')]) # 1. 生成包含年份、值的数组,并展开为行 expanded_sdf = data_sdf. \ withColumn('year_value', func.explode( func.array(*[func.struct(func.lit(int(c[-4:])).alias('year'), func.col(c).alias('value')) for c in yearcols]) )). \ select('a', 'b', 'id', func.col('year_value.year').cast('int'), func.col('year_value.value')) # 2. 过滤值为1的行,计算分组标识:当当前值为1且前一个值为0时,分组+1 window_spec = Window.partitionBy('a', 'b', 'id').orderBy('year') grouped_sdf = expanded_sdf. \ filter(func.col('value') == 1). \ withColumn('prev_value', func.lag('value', 1, 0).over(window_spec)). \ withColumn('group_id', func.sum(func.when((func.col('value') == 1) & (func.col('prev_value') == 0), 1)).over(window_spec)) # 3. 按分组聚合,得到每个区间的起始和结束年份 result_sdf = grouped_sdf. \ groupBy('a', 'b', 'id', 'group_id'). \ agg(func.min('year').alias('startdate'), func.max('year').alias('enddate')). \ drop('group_id'). \ orderBy('id', 'startdate') result_sdf.show(truncate=False)
代码说明
- 展开行:将每个年份列的键值对转为数组后展开,让每个年份成为单独一行,方便后续处理连续值
- 分组标识计算:用窗口函数
lag获取前一个年份的值,当当前值为1且前一个为0时,说明进入新的连续区间,分组ID递增 - 聚合区间:按分组ID聚合,取每个组的最小年份作为起始,最大年份作为结束,得到连续的日期区间
运行后即可得到符合期望的输出。
内容的提问来源于stack exchange,提问作者lunbox
相关产品推荐
相关产品推荐

