diff --git a/tools/codegen_jvm.py b/tools/codegen_jvm.py index 8841cb6ae..e6fb45d29 100644 --- a/tools/codegen_jvm.py +++ b/tools/codegen_jvm.py @@ -2178,12 +2178,13 @@ def _spark_planning(arity): _SPARK_AGGREGATES = '''package {pkg}; import functions.GeneratedFunctions; -import org.apache.spark.sql.Column; import org.apache.spark.sql.Encoder; import org.apache.spark.sql.Encoders; import org.apache.spark.sql.SparkSession; import org.apache.spark.sql.catalyst.FunctionIdentifier; +import org.apache.spark.sql.catalyst.encoders.package$; import org.apache.spark.sql.catalyst.expressions.Expression; +import org.apache.spark.sql.execution.aggregate.ScalaAggregator; import org.apache.spark.sql.expressions.Aggregator; import org.apache.spark.sql.types.DataType; import org.mobilitydb.spark.generated.TemporalAggregate; @@ -2247,8 +2248,11 @@ def _spark_planning(arity): for (Overload o : overloads) {{ if (o.arg.equals(args.head().dataType())) {{ Aggregator typed = new Typed(o, Encoders.bean(o.out)); - return org.apache.spark.sql.functions.udaf(typed, Encoders.bean(o.in)) - .apply(new Column(args.head())).expr(); + return new ScalaAggregator(args, typed, + package$.MODULE$.encoderFor(Encoders.bean(o.in)), + package$.MODULE$.encoderFor(typed.bufferEncoder()), + true, true, 0, 0, Option.empty()) + .toAggregateExpression(); }} }} }}