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

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. 将年份列转为包含年份和对应值的数组结构,保证年份按顺序排列
  2. 对每个元素计算分组标识:当当前值为1且前一个值为0时,分组标识递增,以此区分不同的连续区间
  3. 按分组标识聚合,得到每个区间的起始和结束年份

优化后的代码如下:

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.25 05:23:24