如何在PySpark中调用返回复杂类型的Scala自定义UDF
首先,你的Scala UDF写法存在核心错误:继承UDF2时,泛型参数的第三个类型应该是实际返回值的类型,而不是UserDefinedFunction。你现在的代码在call方法里又定义了一个UDF并返回,这完全不符合UDF的实现逻辑——UDF2的call方法本身就是用来实现业务逻辑、返回计算结果的地方。
步骤1:修正Scala UDF代码
根据你想要的输出格式(数组的数组,而非结构体数组),我们调整代码返回Seq[Seq[String]],这样在PySpark中就能得到你期望的[["a","A"], ["a","B"], ...]格式:
import org.apache.spark.sql.api.java.UDF2 import scala.collection.Seq class DualArrayExplode extends UDF2[Seq[String], Seq[String], Seq[Seq[String]]] { override def call(x: Seq[String], y: Seq[String]): Seq[Seq[String]] = { // 生成笛卡尔积并转换为嵌套字符串序列 for (a <- x; b <- y) yield Seq(a, b) } } object DualArrayExplode { def apply(): DualArrayExplode = new DualArrayExplode() }
如果你需要保留原Scala中的元组格式(对应Spark的结构体类型),可以把返回类型改为Seq[(String, String)],后续PySpark中对应结构体数组类型即可。
步骤2:编译打包Jar包
重新编译代码生成Jar包,确保没有编译错误。
步骤3:在PySpark中注册并使用UDF
定义正确的返回类型
针对我们修正后的返回值(嵌套字符串数组),在PySpark中需要定义返回类型为ArrayType(ArrayType(StringType()))。如果是结构体类型,则定义为ArrayType(StructType([StructField("_1", StringType()), StructField("_2", StringType())]))。
注册并调用UDF
from pyspark.sql import SparkSession from pyspark.sql.types import ArrayType, StringType spark = SparkSession.builder \ .appName("ScalaUDFTest") \ .config("spark.jars", "/path/to/your/dual-array-explode.jar") # 替换为你的Jar包路径 .getOrCreate() # 注册Scala UDF spark.udf.registerJavaFunction( "dual_array_explode", "com.your.package.DualArrayExplode", # 替换为你的类全限定名 ArrayType(ArrayType(StringType())) ) # 测试使用 test_df = spark.createDataFrame([(["a", "b"], ["A", "B"])], ["x", "y"]) test_df.selectExpr("dual_array_explode(x, y) as result").show(truncate=False)
执行后会得到你期望的输出:
+--------------------------------+ |result | +--------------------------------+ |[[a, A], [a, B], [b, A], [b, B]]| +--------------------------------+
为什么原代码在PySpark中返回空列表?
原Scala代码的call方法返回的是UserDefinedFunction类型,这完全不符合Spark UDF的执行逻辑——Spark期望UDF直接返回计算结果,而不是返回另一个UDF对象,所以PySpark无法正确解析这个返回值,最终只能得到空列表。修正Scala UDF的返回类型和实现逻辑后,问题就解决了。
内容的提问来源于stack exchange,提问作者mamonu

