如何在PySpark pandas API的groupby操作中使用UDF?
错误原因
- 你使用的是PySpark的Pandas API(即pyspark.pandas,原Koalas项目),它的
groupby.agg方法仅支持用字符串形式调用内置聚合函数(如min、max),无法识别你手动注册到Spark会话的自定义UDF,会将传入的UDF名称识别为普通列名,因此抛出列不在分组键中也不是聚合函数的AnalysisException。 - 你定义的UDF没有指定聚合类型:普通
udf是标量函数,输入为单行单个值、输出为单个值,无法用于分组聚合;普通定义的pandas_udf默认也是标量类型,同样不能直接用于分组聚合场景,因此会触发未实现或语法错误。 - 提示
PandasNotImplementedError是因为PySpark Pandas API的分组聚合方法确实没有适配「传入UDF名称字符串调用自定义UDF」的场景,仅适配了内置函数的字符串调用逻辑,并非报错描述不准确。
正确实现方式
首先需要定义分组聚合类型的Pandas UDF,再直接传入UDF对象到agg方法中即可,示例代码如下:
import pyspark from pyspark.sql import SparkSession from pyspark.sql.functions import pandas_udf, PandasUDFType from pyspark import pandas as ps spark = SparkSession.builder.getOrCreate() df = ps.DataFrame({'A': 'a a b'.split(), 'B': [1, 2, 3], 'C': [4, 6, 5]}, columns=['A', 'B', 'C']) # 定义分组聚合类型的pandas UDF,指定functionType为GROUPED_AGG @pandas_udf('float', functionType=PandasUDFType.GROUPED_AGG) def agg_a(x): return (x**2).mean() @pandas_udf('float', functionType=PandasUDFType.GROUPED_AGG) def agg_b(x): return x.mean() # 聚合时直接传入UDF对象并指定作用的列 dfout = df.groupby('A').agg( B=agg_a(df['B']), C=agg_b(df['C']) )
如果更习惯用SQL语法操作,也可以将PySpark Pandas DataFrame转为原生Spark DataFrame处理后再转回:
# 转为原生Spark DataFrame并注册临时表 spark_df = df.to_spark() spark_df.createOrReplaceTempView('tmp_df') # 注册UDF到Spark会话 spark.udf.register('agg_a_', agg_a) spark.udf.register('agg_b_', agg_b) # 用SQL执行分组聚合 result_spark_df = spark.sql(""" SELECT A, agg_a_(B) AS B, agg_b_(C) AS C FROM tmp_df GROUP BY A """) # 转回PySpark Pandas DataFrame dfout = result_spark_df.to_pandas_on_spark()
内容的提问来源于stack exchange,提问作者figs_and_nuts
相关产品推荐
相关产品推荐

