You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

PySpark Pandas UDF处理数组列求最小值报错排查

Spark数组列提取最小值的问题与修复

问题场景

我有如下Spark DataFrame:

data_df = spark.createDataFrame([([1,2,3],'val1'),([4,5,6],'val2')],['col1','col2'])

数据集展示:

Col1Col2
[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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.27 15:45:16