PySpark调用Scala编写的Java UDF触发类型转换异常排查
问题原因与解决方案
核心问题
在Spark 2.3.3版本中,PySpark使用spark.udf.registerJavaFunction注册Scala实现的Java UDF时,默认的返回类型反射推断存在bug,无法正确识别UDF的Double返回类型,错误地将其推断为struct<>,从而触发类型转换异常。
这是因为Spark 2.x对跨语言UDF的类型推断逻辑不完善,尤其是当UDF基于Scala的UDF1接口实现时,PySpark侧无法准确通过反射获取返回类型的元数据,导致类型匹配失败。
解决方案
注册UDF时显式指定返回类型,避免依赖自动推断。修改PySpark中的注册代码,在registerJavaFunction中添加第三个参数,指定返回类型为DoubleType:
from pyspark.sql.types import DoubleType import pyspark.sql.functions as F def call_mid_val(): # 显式指定返回类型为DoubleType spark.udf.registerJavaFunction("getMidVal", "org.spark.udf.GetMidVal", DoubleType()) data = [ ("001", 3.0), ("001", 2.3), ("001", 1.5), ("001", 4.2), ("001", 9.6), ("001", 7.3) ] df = spark.createDataFrame(data, ['id', 'trx_amt']) \ .groupBy("id").agg(F.collect_list("trx_amt").alias("trx_amt_seq")) df.show(truncate=False) df.printSchema() print(df.dtypes) print(df.schema) spark.createDataFrame(data, ['id', 'trx_amt'])\ .groupBy("id").agg(F.collect_list("trx_amt").alias("trx_amt_seq"))\ .select(F.expr("getMidVal(trx_amt_seq)"))\ .show()
额外说明
- 输入类型兼容性:PySpark中
collect_list("trx_amt")返回的Array[Double]与Scala UDF接收的mutable.WrappedArray[Double]是兼容的,Spark会自动完成类型转换,无需修改Scala UDF的输入参数类型。 - 中位数逻辑修正:原Scala UDF中的中位数计算逻辑不符合标准规则,若需要正确计算中位数,建议调整为:
override def call(arr: mutable.WrappedArray[Double]): Double = { val n = arr.length val arr_sorted = arr.sorted if (n % 2 == 1) { arr_sorted(n / 2) } else { (arr_sorted(n/2 -1) + arr_sorted(n/2)) / 2.0 } }
内容的提问来源于stack exchange,提问作者jz w
相关产品推荐
相关产品推荐

