PySpark Pandas UDF处理数组列求最小值报错排查
Spark数组列提取最小值的问题与修复
问题场景
我有如下Spark DataFrame:
data_df = spark.createDataFrame([([1,2,3],'val1'),([4,5,6],'val2')],['col1','col2'])
数据集展示:
| Col1 | Col2 |
|---|---|
| [1,2,3] | val1 |
| [4,5,6] | val2 |
目标是提取col1数组列中的最小值,预期结果:
| Col1 |
|---|
| 1 |
| 4 |
报错情况
使用Pandas UDF实现时触发错误:
An exception was thrown from a UDF: 'AssertionError: Pandas SCALAR_ITER UDF outputted more rows than input rows.
出错代码
def generate_min(batch_iter: Iterator[pd.Series]) -> Iterator[pd.Series]: for x in batch_iter: yield min(x) generate__udf = pandas_udf(generate_min, returnType=IntegerType()) data_df.select(generate_min(F.col('col1'))
错误原因
你写的是迭代器式标量Pandas UDF,这类UDF要求每个输入批次的输出行数必须和输入行数完全一致。但代码中yield min(x)返回的是单个数值(整个批次的最小值),而不是和输入批次每行对应的Pandas Series,导致输出行数和输入不匹配,触发断言错误。
解决方案
方案1:修复Pandas UDF
修改逻辑,对输入Series中的每个数组单独取最小值,返回行数匹配的Series:
from pyspark.sql import functions as F from pyspark.sql.types import IntegerType from pyspark.sql.functions import pandas_udf import pandas as pd from typing import Iterator def generate_min(batch_iter: Iterator[pd.Series]) -> Iterator[pd.Series]: for x in batch_iter: # 遍历每个数组元素,取最小值后生成新Series yield x.apply(lambda arr: min(arr)) generate_udf = pandas_udf(generate_min, returnType=IntegerType()) # 执行查询 data_df.select(generate_udf(F.col('col1')).alias('Col1')).show()
方案2:使用Spark内置函数(推荐)
Spark提供了array_min内置函数,无需自定义UDF就能直接提取数组最小值,性能更优:
data_df.select(F.array_min(F.col('col1')).alias('Col1')).show()
内容的提问来源于stack exchange,提问作者lserlohn
相关产品推荐
相关产品推荐

