如何在PySpark SQL的when()子句中使用聚合值进行条件判断
PySpark 聚合值传入when()子句实现分类的解决方法
你遇到的报错本质是agg()返回的是DataFrame结构,而非可直接参与计算的数值标量,无法直接用在逻辑比较表达式中,以下是两种可行的实现方案:
方案1:提取聚合结果为Python标量(适合全局聚合场景)
聚合后通过collect()方法拿到结果行,再提取行内的数值即可得到float类型的标量,可直接传入when条件,修正后的可运行代码如下:
from pyspark.sql import SparkSession from pyspark.sql.functions import when import pyspark.sql.functions as f # 已有SparkSession可跳过初始化步骤 spark = SparkSession.builder.appName("quant_classify").getOrCreate() samp = spark.createDataFrame( [("A","A1",4,1.25),("B","B3",3,2.14),("C","C2",7,4.24),("A","A3",4,1.25),("B","B1",3,2.14),("C","C1",7,4.24)], ["Category","Sub-cat","quantity","cost"] ) # 核心修改:提取聚合结果的实际数值 psMean = samp.agg({'quantity':'mean'}).collect()[0][0] psStDev = samp.agg({'quantity':'stddev'}).collect()[0][0] # 执行分类逻辑 psCatVect = samp.withColumn('quant_category', when(samp['quantity'] <= (psMean - psStDev), 'small').otherwise('not small')) # 查看结果 psCatVect.show()
方案2:使用窗口函数(适合分组聚合场景,无需单独提取标量)
如果分类逻辑需要按分组计算组内均值/标准差,或者数据量较大不适合将聚合值拉取到Driver端,用窗口函数更高效,不需要单独提取每个组的聚合值,直接对每行计算对应分组的聚合结果:
from pyspark.sql.window import Window # 定义窗口:全局计算不加partitionBy,按Category分组就填partitionBy("Category") w = Window.partitionBy() psCatVect = samp.withColumn('mean_quant', f.mean('quantity').over(w))\ .withColumn('std_quant', f.stddev('quantity').over(w))\ .withColumn('quant_category', when(f.col('quantity') <= (f.col('mean_quant') - f.col('std_quant')), 'small').otherwise('not small'))\ .drop('mean_quant', 'std_quant') # 清除临时辅助列 psCatVect.show()
内容的提问来源于stack exchange,提问作者DeCodened
相关产品推荐
相关产品推荐

