如何在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

