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

如何在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的计算模型:

  1. 为两个DataFrame添加年份标识列,合并为一个长表。
  2. 按数据的粒度维度(如业务主键、分组列)分组,基于窗口函数获取每组内的前后年份及对应数值。
  3. 根据目标年份计算线性插值系数,对数值列进行插值计算。
  4. 筛选或生成目标年份的结果数据,保持与原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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.25 16:05:19