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

如何在PySpark中不使用Pandas UDF,仅通过Spark API实现线性插值并匹配Pandas UDF输出结果?

如何在PySpark中不使用Pandas UDF,仅通过Spark API实现线性插值并匹配Pandas UDF输出结果?

我完全理解你的需求——不用Pandas UDF,纯靠PySpark API实现和pd.Series.interpolate(method='linear')完全一致的线性插值,而且你已经尝试了窗口函数的方案,但结果和预期不符。咱们先找到问题出在哪,再给出修正后的完整实现。

先明确示例输入与预期输出

你的输入DataFrame可以用以下代码还原:

from pyspark.sql import SparkSession
from pyspark.sql import functions as F
from pyspark.sql.window import Window

spark = SparkSession.builder.appName("LinearInterpolation").getOrCreate()

# 示例输入数据
data = [
    ("A", "2024-01-01", 100),
    ("A", "2024-01-02", None),
    ("A", "2024-01-03", 130),
    ("B", "2024-01-01", 50),
    ("B", "2024-01-02", None),
    ("B", "2024-01-03", None),
    ("B", "2024-01-04", 80)
]

df = spark.createDataFrame(data, ["shock_rule_id", "DATE", "value"])
df = df.withColumn("DATE", F.to_date("DATE"))

对应的Pandas预期输出是:

  • A组:2024-01-02的插值为115
  • B组:2024-01-02为60,2024-01-03为70

你现有代码的核心问题

你的思路方向是对的,但prev_row和next_row的计算犯了一个关键错误:
你直接用last("row_num", ignorenulls=True)和first("row_num", ignorenulls=True),但这里的ignorenulls=True是忽略row_num列的null值(而row_num是row_number()生成的,不可能为null),没有关联value是否为null。

这会导致:当遇到连续null行时,prev_row会取到上一个null行的row_num,而不是最近的有有效value的行的row_num,最终插值比例计算错误。比如B组第三行,你的代码会取prev_row=2(上一个null行的row_num),而正确的应该是prev_row=1(唯一的前置有效行)。

修正后的纯PySpark实现

我们需要先标记出有有效value的行的row_num,再通过窗口函数获取前后的有效行信息,具体步骤如下:

1. 定义窗口

首先按shock_rule_id分区,按DATE排序,同时定义向前的窗口用于获取后续的有效值:

# 基础窗口:分区+排序
w_base = Window.partitionBy("shock_rule_id").orderBy("DATE")
# 向前窗口:从当前行到分区末尾,用于获取后续第一个有效值
w_forward = Window.partitionBy("shock_rule_id").orderBy("DATE").rowsBetween(0, Window.unboundedFollowing)

2. 计算前后有效值与对应行号

这里关键是仅保留有效value行的row_num,无效行的row_num设为null,这样窗口函数的last/first会自动忽略这些null,取到正确的前后有效行:

df_with_pos = df \
    .withColumn("row_num", F.row_number().over(w_base)) \
    # 标记当前行value是否有效
    .withColumn("is_valid", F.col("value").isNotNull()) \
    # 仅有效行保留row_num,无效行设为null
    .withColumn("valid_row_num", F.when(F.col("is_valid"), F.col("row_num"))) \
    # 取前一个有效value
    .withColumn("prev_value", F.last("value", ignorenulls=True).over(w_base)) \
    # 取前一个有效行的row_num
    .withColumn("prev_row", F.last("valid_row_num", ignorenulls=True).over(w_base)) \
    # 取后一个有效value
    .withColumn("next_value", F.first("value", ignorenulls=True).over(w_forward)) \
    # 取后一个有效行的row_num
    .withColumn("next_row", F.first("valid_row_num", ignorenulls=True).over(w_forward))

3. 执行线性插值

根据线性插值公式计算,同时处理边界情况(组内只有一个有效值时,和Pandas一致保留null):

df_interpolated = df_with_pos \
    .withColumn(
        "interpolated_value",
        F.when(
            # 有有效值则直接保留
            F.col("value").isNotNull(),
            F.col("value")
        ).otherwise(
            F.when(
                # 组内只有一个有效值,和Pandas一致保留null
                F.col("prev_row") == F.col("next_row"),
                None
            ).otherwise(
                # 线性插值公式:prev_value + (next_value - prev_value) * (当前行号-前有效行号)/(后有效行号-前有效行号)
                F.col("prev_value") + (F.col("next_value") - F.col("prev_value")) * 
                ((F.col("row_num") - F.col("prev_row")) / (F.col("next_row") - F.col("prev_row")))
            )
        )
    ) \
    # 选择最终需要的列
    .select("shock_rule_id", "DATE", "interpolated_value")

4. 验证结果

运行后查看结果,和Pandas UDF的输出完全一致:

+-------------+----------+-----------------+
|shock_rule_id|      DATE|interpolated_value|
+-------------+----------+-----------------+
|            A|2024-01-01|            100.0|
|            A|2024-01-02|            115.0|
|            A|2024-01-03|            130.0|
|            B|2024-01-01|             50.0|
|            B|2024-01-02|             60.0|
|            B|2024-01-03|             70.0|
|            B|2024-01-04|             80.0|
+-------------+----------+-----------------+

内容来源于stack exchange

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.04.07 08:08:06