如何在PySpark中实现两个年份DataFrame的线性插值生成中间数据?
PySpark 跨年份DataFrame线性插值实现方案
问题说明
现有两个结构一致、存储数值型数据的PySpark DataFrame,分别对应不同年份(如2020、2030)且数据粒度相同,需要通过线性插值生成中间年份(如2025)的新DataFrame。同时询问是否推荐使用pyspark.pandas.DataFrame.interpolate方法,并提供了一段Pandas插值代码,需要迁移到PySpark环境。
关于pyspark.pandas.DataFrame.interpolate的建议
不推荐直接使用该方法,原因如下:
pyspark.pandas是Spark对Pandas API的封装,底层依赖分布式执行,但interpolate方法的实现逻辑更适配单机Pandas场景,在分布式环境下容易出现分区内插值而非全局插值的问题,结果不符合预期。- 该方法性能开销较高,调试和优化难度大于原生Spark API,无法充分利用Spark的分布式计算优势。
PySpark实现思路(对应原Pandas代码逻辑迁移)
原Pandas代码通过拼接宽表、转置重索引后插值的逻辑,在Spark分布式环境下需要调整为长表分组插值的方式,更符合Spark的计算模型:
- 为两个DataFrame添加年份标识列,合并为一个长表。
- 按数据的粒度维度(如业务主键、分组列)分组,基于窗口函数获取每组内的前后年份及对应数值。
- 根据目标年份计算线性插值系数,对数值列进行插值计算。
- 筛选或生成目标年份的结果数据,保持与原DataFrame一致的结构。
代码示例
1. 准备测试数据
from pyspark.sql import SparkSession from pyspark.sql import functions as F from pyspark.sql.window import Window # 初始化Spark会话 spark = SparkSession.builder.appName("yearly_interpolation").getOrCreate() # 模拟2020年数据(id为分组键,value1、value2为数值列) df_2020 = spark.createDataFrame([ ("A", 100, 200), ("B", 150, 250) ], ["id", "value1", "value2"]) df_2020 = df_2020.withColumn("year", F.lit(2020)) # 模拟2030年数据 df_2030 = spark.createDataFrame([ ("A", 300, 400), ("B", 450, 550) ], ["id", "value1", "value2"]) df_2030 = df_2030.withColumn("year", F.lit(2030))
2. 生成指定中间年份的插值数据
def interpolate_target_year(df1, df2, target_year): # 获取两个基准年份 year1 = df1.select("year").first()[0] year2 = df2.select("year").first()[0] # 合并两年数据 combined_df = df1.unionByName(df2) # 定义窗口:按分组键id分区,按年份排序 window = Window.partitionBy("id").orderBy("year") # 为每组添加前后年份、前后数值的辅助列 processed_df = combined_df value_cols = [col for col in df1.columns if col not in ["id", "year"]] processed_df = processed_df.withColumn("prev_year", F.lag("year").over(window)) \ .withColumn("next_year", F.lead("year").over(window)) for col in value_cols: processed_df = processed_df.withColumn(f"prev_{col}", F.lag(col).over(window)) \ .withColumn(f"next_{col}", F.lead(col).over(window)) # 计算插值系数与目标年份数值 coeff = (target_year - F.col("prev_year")) / (F.col("next_year") - F.col("prev_year")) for col in value_cols: processed_df = processed_df.withColumn( col, F.col(f"prev_{col}") + coeff * (F.col(f"next_{col}") - F.col(f"prev_{col}")) ) # 筛选并整理目标年份的结果 result = processed_df.filter(F.col("year").isNull()) \ .withColumn("year", F.lit(target_year)) \ .select(["id", "year"] + value_cols) return result # 生成2025年插值数据 df_2025 = interpolate_target_year(df_2020, df_2030, 2025) df_2025.show()
3. 生成所有中间年份的插值数据
如果需要生成两个基准年份之间的所有年份数据,可以使用以下方法:
def interpolate_all_years(df1, df2): year1 = df1.select("year").first()[0] year2 = df2.select("year").first()[0] all_years = list(range(year1, year2 + 1)) combined_df = df1.unionByName(df2) window = Window.partitionBy("id").orderBy("year") # 生成所有年份的行 exploded_df = combined_df.withColumn( "year", F.explode(F.array([F.lit(y) for y in all_years])) ) # 添加前后年份和数值的辅助列 value_cols = [col for col in df1.columns if col not in ["id", "year"]] exploded_df = exploded_df.withColumn("prev_year", F.lag("year").over(window)) \ .withColumn("next_year", F.lead("year").over(window)) for col in value_cols: exploded_df = exploded_df.withColumn(f"prev_{col}", F.lag(col).over(window)) \ .withColumn(f"next_{col}", F.lead(col).over(window)) # 对原始年份保留原值,中间年份计算插值 for col in value_cols: exploded_df = exploded_df.withColumn( col, F.when(F.col("year").isin([year1, year2]), F.col(col)) .otherwise( F.col(f"prev_{col}") + (F.col("year") - F.col("prev_year")) / (F.col("next_year") - F.col("prev_year")) * (F.col(f"next_{col}") - F.col(f"prev_{col}")) ) ) # 去重并整理列顺序 result = exploded_df.select(["id", "year"] + value_cols).distinct() return result # 生成2020-2030所有年份的数据 all_years_df = interpolate_all_years(df_2020, df_2030) all_years_df.orderBy("id", "year").show()
内容的提问来源于stack exchange,提问作者Pfinnn
相关产品推荐
相关产品推荐

