PySpark实现年份范围拆分展开的函数需求
问题描述
我在Spark中有如下DataFrame:
from pyspark.sql.types import StructType,StructField, StringType, IntegerType data2 = [("a","2010 - 2012"), ("b","from 2020",) ] schema = StructType([ \ StructField("product",StringType(),True), \ StructField("reportingYears",StringType(),True) ]) df = spark.createDataFrame(data=data2,schema=schema) df.printSchema() df.display()
当前输出:
+-------+--------------+ |product|reportingYears| +-------+--------------+ | a| 2010 - 2012| | b| from 2020| | c| 2010| +-------+--------------+
需要编写函数将年份范围展开为如下形式:
+-------+--------------+ |product|reportingYears| +-------+--------------+ | a| 2010| | a| 2011| | a| 2012| | b| 2020| | b| 2021| | b| 2022| | c| 2010| +-------+--------------+
不确定PySpark是否有原生功能实现,希望获取类似Python函数的解决方案。补充:数据集中reportingYears字段还存在无范围的单个年份值。
解决方案
可以通过自定义UDF或Spark内置函数实现需求,以下提供两种可行方案:
方案一:自定义UDF实现
适合逻辑复杂、需要灵活扩展格式处理的场景:
from pyspark.sql import SparkSession from pyspark.sql.functions import udf, explode from pyspark.sql.types import ArrayType, IntegerType from datetime import datetime # 初始化SparkSession spark = SparkSession.builder.appName("YearExpand").getOrCreate() # 补充完整原始数据 data2 = [("a","2010 - 2012"), ("b","from 2020"), ("c","2010") ] schema = StructType([ StructField("product",StringType(),True), StructField("reportingYears",StringType(),True) ]) df = spark.createDataFrame(data=data2,schema=schema) # 定义年份解析函数 def parse_years(year_str): # 处理单个年份 if year_str.isdigit(): return [int(year_str)] # 处理"X - Y"格式的年份范围 elif " - " in year_str: start, end = year_str.split(" - ") return list(range(int(start), int(end)+1)) # 处理"from X"格式,示例中到2022,可替换为datetime.now().year获取当前年份 elif year_str.startswith("from "): start_year = int(year_str.split("from ")[1]) end_year = 2022 # 替换为datetime.now().year可动态取当前年份 return list(range(start_year, end_year+1)) # 处理未知格式 else: return [] # 注册UDF,指定返回类型为整数数组 parse_years_udf = udf(parse_years, ArrayType(IntegerType())) # 应用UDF并展开数组为多行 result_df = df.withColumn("year_list", parse_years_udf(df.reportingYears)) \ .select("product", explode("year_list").alias("reportingYears")) # 查看结果 result_df.show()
方案二:Spark内置函数实现
无需UDF,性能更优,适合大规模数据场景:
from pyspark.sql.functions import regexp_extract, expr, explode, sequence, col # 使用内置函数解析年份范围 result_df = df.withColumn( "start_year", expr(""" CASE WHEN reportingYears RLIKE '^\\d+$' THEN cast(reportingYears as int) WHEN reportingYears RLIKE '^\\d+ - \\d+$' THEN cast(split(reportingYears, ' - ')[0] as int) WHEN reportingYears RLIKE '^from \\d+$' THEN cast(regexp_extract(reportingYears, 'from (\\d+)', 1) as int) ELSE NULL END """) ).withColumn( "end_year", expr(""" CASE WHEN reportingYears RLIKE '^\\d+$' THEN cast(reportingYears as int) WHEN reportingYears RLIKE '^\\d+ - \\d+$' THEN cast(split(reportingYears, ' - ')[1] as int) WHEN reportingYears RLIKE '^from \\d+$' THEN 2022 # 替换为year(current_date())可动态取当前年份 ELSE NULL END """) ).withColumn( "year_list", sequence(col("start_year"), col("end_year")) ).select("product", explode("year_list").alias("reportingYears")) # 查看结果 result_df.show()
两种方案执行后均可得到目标输出。
内容的提问来源于stack exchange,提问作者Amdbi
相关产品推荐
相关产品推荐

