Spark Aggregator接收Array[String]输入时触发空指针异常问题
我在将Scala Spark中的UDAF从UserDefinedAggregateFunction迁移到Aggregator时,碰到一个奇怪问题:当Aggregator以Array[String]作为输入类型时,本地测试触发空指针异常。即便把代码简化到最基础的版本,读取Array时还是报错,但其他输入类型完全正常。
问题代码
Aggregator实现
class ArrayInputAggregator extends Aggregator[Array[String], Int, Int] with Serializable { override def zero = {0} override def reduce(buffer: Int, newItem: Array[String]): Int = { buffer + newItem.length } override def merge(b1: Int, b2: Int): Int = { b1 + b2 } override def finish(reduction: Int): Int = reduction def bufferEncoder: Encoder[Int] = Encoders.scalaInt def outputEncoder: Encoder[Int] = Encoders.scalaInt }
测试代码
val test = udaf(new ArrayInputAggregator()) val d = spark .sql("select array('asd','tre','asd') arr") .groupBy() .agg(test($"arr").as("cnt")) d.show
触发的异常信息
2023-12-24 12:06:24,678 ERROR spark.executor.Executor - Exception in
task 0.0 in stage 0.0 (TID 0) java.lang.NullPointerException: null at
org.apache.spark.sql.catalyst.expressions.objects.MapObjects$.apply(objects.scala:682)
~[spark-catalyst_2.12-3.0.0.jar:3.0.0] at
org.apache.spark.sql.catalyst.analysis.Analyzer$ResolveDeserializer$$anonfun$apply$31$$anonfun$applyOrElse$172$$anonfun$10.applyOrElse(Analyzer.scala:3033)
~[spark-catalyst_2.12-3.0.0.jar:3.0.0] at
org.apache.spark.sql.catalyst.analysis.Analyzer$ResolveDeserializer$$anonfun$apply$31$$anonfun$applyOrElse$172$$anonfun$10.applyOrElse(Analyzer.scala:3029)
~[spark-catalyst_2.12-3.0.0.jar:3.0.0] at
org.apache.spark.sql.catalyst.trees.TreeNode.$anonfun$transformDown$1(TreeNode.scala:309)
~[spark-catalyst_2.12-3.0.0.jar:3.0.0] at
org.apache.spark.sql.catalyst.trees.CurrentOrigin$.withOrigin(TreeNode.scala:72)
~[spark-catalyst_2.12-3.0.0.jar:3.0.0] at
(异常信息后续内容省略)
解决办法
问题根源在于Spark对原生Array类型的Encoder支持不足,直接用Array[String]作为Aggregator输入类型时,Spark无法自动生成正确的序列化/反序列化逻辑,导致空指针异常。
解决方法很简单:把输入类型换成Spark内部的数组包装类型WrappedArray[String],它能被Spark正确识别并处理。
修正后的Aggregator代码
import org.apache.spark.sql.catalyst.util.WrappedArray import org.apache.spark.sql.{Encoder, Encoders} import org.apache.spark.sql.expressions.Aggregator class ArrayInputAggregator extends Aggregator[WrappedArray[String], Int, Int] with Serializable { override def zero: Int = 0 override def reduce(buffer: Int, newItem: WrappedArray[String]): Int = { buffer + newItem.length } override def merge(b1: Int, b2: Int): Int = { b1 + b2 } override def finish(reduction: Int): Int = reduction override def bufferEncoder: Encoder[Int] = Encoders.scalaInt override def outputEncoder: Encoder[Int] = Encoders.scalaInt }
测试代码不需要任何修改,Spark会自动将SQL中的array类型转换为WrappedArray[String]传入Aggregator,运行后就能正常统计数组长度总和了。
内容的提问来源于stack exchange,提问作者Eyal

