如何在PySpark的agg函数中使用lag窗口函数?
问题解答
不能在agg函数里直接使用lag函数,Spark的语法规则明确禁止窗口函数嵌套在聚合函数内部,报错信息已经提示你得用子查询先处理窗口逻辑。
要实现你要的计算,正确的步骤是先通过窗口函数算出每行的start - 前一个end差值,再对每个id的差值求和,具体代码如下:
from pyspark.sql.functions import * from pyspark.sql import SparkSession from pyspark.sql import Window # 初始化SparkSession(你原代码遗漏了这一步) spark = SparkSession.builder.appName("test").getOrCreate() data = [ ('a', 10, 15), ('a', 40, 60), ('a', 70, 100), ('b', 10, 20), ('b', 30, 50), ('b', 60, 80) ] schema = ['id', 'start', 'end'] df = spark.createDataFrame(data, schema=schema) # 1. 定义窗口:按id分区,按start排序(保证取到同一id的前一行end) window_spec = Window.partitionBy("id").orderBy("start") # 2. 计算每行的差值:当前start - 上一行的end,第一行没有上一行,结果为null df_diff = df.withColumn("diff", col("start") - lag(col("end"), 1).over(window_spec)) # 3. 按id分组,对差值求和(sum会自动忽略null值) result_df = df_diff.groupBy("id").agg(sum("diff").alias("total_diff")) result_df.show()
执行后会得到预期结果:
+---+----------+ | id|total_diff| +---+----------+ | a| 35| | b| 20| +---+----------+
原理很简单:窗口函数是对每行数据做行级计算,聚合函数是对分组数据做汇总,两者的执行阶段不同,必须先完成窗口函数的计算生成中间列,再基于中间列做聚合操作,不能直接嵌套调用。
内容的提问来源于stack exchange,提问作者wong.lok.yin
相关产品推荐
相关产品推荐

