如何为返回Double数组的PandasUDFType.GROUPED_AGG设置正确类型提示
Pandas UDF分组聚合数组列元素级平均的类型提示修正
问题场景
使用分组聚合(Grouped Agg)Pandas UDF对数组列做元素级平均(均值池化)时,触发弃用警告,且找不到适配返回ArrayType(DoubleType())的GROUPED_AGG类型的正确类型提示。
警告信息
UserWarning: In Python 3.6+ and Spark 3.0+, it is preferred to specify type hints for pandas UDF instead of specifying pandas UDF type which will be deprecated in the future releases. See SPARK-28264 for more details.
原实现代码
from pyspark.sql.types import ArrayType, DoubleType import pyspark.sql.functions as F import numpy as np import pandas as pd # 构造测试数据 pdf = pd.DataFrame( {"id": [1, 1, 2, 2], "x": [[1.0, 1.0], [2.0, 2.0], [3.0, 3.0], [4.0, 4.0]]} ) df = spark.createDataFrame(pdf) df.show() # 定义均值池化UDF @F.pandas_udf( returnType=ArrayType(DoubleType()), functionType=F.PandasUDFType.GROUPED_AGG ) def mean_pooling_udf(x: pd.Series) -> pd.Series: return np.mean(x, axis=0) # 应用UDF并输出结果 df.groupby("id").agg(mean_pooling_udf(df["x"])).show()
问题根源
- Spark 3.0+ 已弃用显式指定
functionType的方式,要求改用类型提示定义Pandas UDF的输入输出类型。 - 分组聚合类型的Pandas UDF要求返回单个标量值(此处指单个数组对象),原代码中类型提示标注返回
pd.Series不符合该要求。
修正后的代码
from pyspark.sql.types import ArrayType, DoubleType import pyspark.sql.functions as F import numpy as np import pandas as pd # 构造测试数据 pdf = pd.DataFrame( {"id": [1, 1, 2, 2], "x": [[1.0, 1.0], [2.0, 2.0], [3.0, 3.0], [4.0, 4.0]]} ) df = spark.createDataFrame(pdf) df.show() # 定义修正后的均值池化UDF @F.pandas_udf(ArrayType(DoubleType())) def mean_pooling_udf(x: pd.Series) -> np.ndarray: # 将Series中的数组元素堆叠为二维数组,按行计算均值 stacked_arrays = np.stack(x.values) return np.mean(stacked_arrays, axis=0) # 应用UDF并输出结果 df.groupby("id").agg(mean_pooling_udf(F.col("x")).alias("mean_x")).show()
输出结果
+---+----------+ | id| x| +---+----------+ | 1|[1.0, 1.0]| | 1|[2.0, 2.0]| | 2|[3.0, 3.0]| | 2|[4.0, 4.0]| +---+----------+ +---+----------+ | id| mean_x| +---+----------+ | 1|[1.5, 1.5]| | 2|[3.5, 3.5]| +---+----------+
内容的提问来源于stack exchange,提问作者David
相关产品推荐
相关产品推荐

