Pyspark如何按分组基于上一行值迭代计算填充指定列?
问题分析
普通窗口函数的lag只能读取源表中已存储的上一行原始值,无法读取前面行动态计算生成的c值,因此不能直接用lag实现该递归累积计算逻辑,以下是两种可行方案:
方案1:Spark高阶函数实现(Spark 2.4+ 推荐,性能最优)
利用collect_list+aggregate高阶函数,在窗口内直接做递归累积计算:
from pyspark.sql import Window from pyspark.sql import functions as F # 定义窗口:按group分区,date升序排序,窗口覆盖分组内从首行到当前行的所有数据 w = Window.partitionBy("group").orderBy("date").rowsBetween(Window.unboundedPreceding, Window.currentRow) # 1. 收集窗口内有序的(a,b)结构数组,同时取分组首行的原始c作为计算初始值 df_with_aux = df.withColumn("ab_list", F.collect_list(F.struct("a", "b")).over(w)) \ .withColumn("init_c", F.first("c", ignorenulls=True).over(Window.partitionBy("group").orderBy("date"))) # 2. 递归计算当前行c值:从第二行开始遍历,每轮用上次累积值 * 当前行b + 当前行a result = df_with_aux.withColumn("c", F.expr(""" aggregate( slice(ab_list, 2, size(ab_list) - 1), init_c, (acc, cur_row) -> acc * cur_row.b + cur_row.a ) """)).drop("ab_list", "init_c") result.show()
如果计算逻辑需要取上一行的a/b值,只需调整cur_row为对应行的取值即可,适配性很强。
方案2:分组逐行计算(兼容所有Spark版本,适合单组数据量不大的场景)
如果Spark版本较低不支持高阶函数,可以用applyInPandas对每个分组排序后逐行计算:
def calc_group_c(iterator): # 同组数据按date升序排序 rows = sorted(list(iterator), key=lambda x: x["date"]) prev_c = None for row in rows: if prev_c is None: # 首行保留原始c值 current_c = row["c"] else: current_c = prev_c * row["b"] + row["a"] row["c"] = current_c prev_c = current_c yield row # 定义输出schema和原表一致 result_schema = df.schema result = df.groupBy("group").applyInPandas(lambda key, it: calc_group_c(it), schema=result_schema) result.show()
用你提供的示例数据执行上述代码,输出结果和预期完全一致。
内容的提问来源于stack exchange,提问作者Ben Bar Ami
相关产品推荐
相关产品推荐

