diff --git a/core/src/main/scala/org/apache/spark/api/python/PythonRunner.scala b/core/src/main/scala/org/apache/spark/api/python/PythonRunner.scala index d7800b1147e1a..015067766c9c6 100644 --- a/core/src/main/scala/org/apache/spark/api/python/PythonRunner.scala +++ b/core/src/main/scala/org/apache/spark/api/python/PythonRunner.scala @@ -87,6 +87,14 @@ private[spark] object PythonEvalType { val SQL_WINDOW_AGG_ARROW_UDF = 253 val SQL_GROUPED_AGG_ARROW_ITER_UDF = 254 + // Incremental (partial + final) Arrow aggregator. Unlike the whole-group grouped-agg UDFs + // above, these support true partial aggregation: the PARTIAL eval type folds input rows into a + // per-group buffer (via the aggregator's `reduce`) on the map side, and the FINAL eval type + // merges partial buffers across the shuffle (via `merge`) and produces the output (via `finish`). + // See PythonIncrementalAggregateExec and the Python `Aggregator` API. + val SQL_GROUPED_AGG_ARROW_INCREMENTAL_PARTIAL_UDF = 255 + val SQL_GROUPED_AGG_ARROW_INCREMENTAL_FINAL_UDF = 256 + val SQL_TABLE_UDF = 300 val SQL_ARROW_TABLE_UDF = 301 val SQL_ARROW_UDTF = 302 @@ -130,6 +138,10 @@ private[spark] object PythonEvalType { case SQL_GROUPED_AGG_ARROW_UDF => "SQL_GROUPED_AGG_ARROW_UDF" case SQL_WINDOW_AGG_ARROW_UDF => "SQL_WINDOW_AGG_ARROW_UDF" case SQL_GROUPED_AGG_ARROW_ITER_UDF => "SQL_GROUPED_AGG_ARROW_ITER_UDF" + case SQL_GROUPED_AGG_ARROW_INCREMENTAL_PARTIAL_UDF => + "SQL_GROUPED_AGG_ARROW_INCREMENTAL_PARTIAL_UDF" + case SQL_GROUPED_AGG_ARROW_INCREMENTAL_FINAL_UDF => + "SQL_GROUPED_AGG_ARROW_INCREMENTAL_FINAL_UDF" } // The eval types produced by ExtractPythonUDFFromLambda: a scalar UDF lifted out of a diff --git a/dev/sparktestsupport/modules.py b/dev/sparktestsupport/modules.py index f272075525efa..e1daa401c1dd0 100644 --- a/dev/sparktestsupport/modules.py +++ b/dev/sparktestsupport/modules.py @@ -610,6 +610,7 @@ def __hash__(self): "pyspark.sql.tests.arrow.test_arrow_cogrouped_map", "pyspark.sql.tests.arrow.test_arrow_cogrouped_map_misc", "pyspark.sql.tests.arrow.test_arrow_grouped_map", + "pyspark.sql.tests.arrow.test_arrow_python_aggregator", "pyspark.sql.tests.arrow.test_arrow_python_udf", "pyspark.sql.tests.arrow.test_arrow_python_udf_cached", "pyspark.sql.tests.arrow.test_arrow_udf", @@ -1253,6 +1254,7 @@ def __hash__(self): "pyspark.sql.tests.connect.arrow.test_parity_arrow_grouped_map", "pyspark.sql.tests.connect.arrow.test_parity_arrow_cogrouped_map", "pyspark.sql.tests.connect.arrow.test_parity_arrow_cogrouped_map_misc", + "pyspark.sql.tests.connect.arrow.test_parity_arrow_python_aggregator", "pyspark.sql.tests.connect.arrow.test_parity_arrow_python_udf", "pyspark.sql.tests.connect.arrow.test_parity_arrow_udf", "pyspark.sql.tests.connect.arrow.test_parity_arrow_udf_scalar", diff --git a/python/pyspark/sql/aggregator.py b/python/pyspark/sql/aggregator.py new file mode 100644 index 0000000000000..b94593c4f8542 --- /dev/null +++ b/python/pyspark/sql/aggregator.py @@ -0,0 +1,202 @@ +# +# Licensed to the Apache Software Foundation (ASF) under one or more +# contributor license agreements. See the NOTICE file distributed with +# this work for additional information regarding copyright ownership. +# The ASF licenses this file to You under the Apache License, Version 2.0 +# (the "License"); you may not use this file except in compliance with +# the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# +""" +Incremental user-defined aggregators for PySpark, the Python analog of Scala's +``org.apache.spark.sql.expressions.Aggregator``. +""" + +from abc import ABC, abstractmethod +from typing import Any, Tuple + +from pyspark.errors import PySparkNotImplementedError, PySparkTypeError +from pyspark.sql.types import DataType, StructType +from pyspark.util import PythonEvalType + +__all__ = ["Aggregator", "udaf"] + + +class Aggregator(ABC): + """ + Base class for a user-defined *incremental* aggregator, the Python analog of Scala's + :class:`org.apache.spark.sql.expressions.Aggregator`. + + Unlike a grouped-aggregate ``pandas_udf`` (which materializes the whole group and is invoked + once), an :class:`Aggregator` is executed as a genuine two-stage aggregation with map-side + combine: :meth:`reduce` folds input rows into a per-group *buffer* on the map side, the buffers + are shuffled by the grouping key, :meth:`merge` combines the partial buffers of each group, and + :meth:`finish` produces the final output value. + + The buffer is represented as a Python :class:`tuple` whose elements correspond, in order, to the + fields of :attr:`bufferSchema`. An input row is likewise a tuple of the argument values passed + to the aggregator call. :meth:`merge` must be associative and commutative, since the framework + may combine partial buffers in any order. + + .. versionadded:: 4.4.0 + + Examples + -------- + A mean aggregator:: + + from pyspark.sql.aggregator import Aggregator, udaf + from pyspark.sql.types import StructType, StructField, DoubleType, LongType + + class Mean(Aggregator): + @property + def bufferSchema(self): + return StructType([ + StructField("sum", DoubleType()), + StructField("count", LongType()), + ]) + + @property + def outputType(self): + return DoubleType() + + def zero(self): + return (0.0, 0) + + def reduce(self, buffer, value): + (v,) = value + if v is None: # ignore null inputs, like SQL aggregates do + return buffer + return (buffer[0] + v, buffer[1] + 1) + + def merge(self, b1, b2): + return (b1[0] + b2[0], b1[1] + b2[1]) + + def finish(self, buffer): + return buffer[0] / buffer[1] if buffer[1] else None + + mean = udaf(Mean()) + df.groupBy("k").agg(mean(df.v)).show() + """ + + @property + @abstractmethod + def bufferSchema(self) -> StructType: + """The schema of the intermediate buffer that crosses the shuffle.""" + ... + + @property + @abstractmethod + def outputType(self) -> DataType: + """The data type of the aggregator's output value.""" + ... + + @abstractmethod + def zero(self) -> Tuple[Any, ...]: + """The initial (identity) buffer value, as a tuple matching :attr:`bufferSchema`.""" + ... + + @abstractmethod + def reduce(self, buffer: Tuple[Any, ...], value: Tuple[Any, ...]) -> Tuple[Any, ...]: + """Fold a single input row ``value`` into ``buffer`` and return the updated buffer.""" + ... + + @abstractmethod + def merge(self, buffer1: Tuple[Any, ...], buffer2: Tuple[Any, ...]) -> Tuple[Any, ...]: + """Merge two partial buffers into one. Must be associative and commutative.""" + ... + + @abstractmethod + def finish(self, buffer: Tuple[Any, ...]) -> Any: + """Produce the output value from the final merged buffer.""" + ... + + # The aggregator instance is shipped to the worker as the UDF "function"; making it callable + # lets it satisfy ``UserDefinedFunction``'s ``callable`` check. It is never actually invoked as + # a function -- the worker calls :meth:`zero`/:meth:`reduce`/:meth:`merge`/:meth:`finish`. + def __call__(self, *args: Any, **kwargs: Any) -> Any: + raise PySparkNotImplementedError( + errorClass="NOT_IMPLEMENTED", + messageParameters={"feature": "calling an Aggregator directly; wrap it with udaf(...)"}, + ) + + +def udaf(agg: "Aggregator") -> Any: + """ + Turn an :class:`Aggregator` instance into a callable usable in ``groupBy().agg(...)``, the + Python counterpart of Scala's ``functions.udaf``. + + The aggregator is executed with true incremental (partial) aggregation and transfers its + intermediate buffer as Arrow; PyArrow is therefore required. + + .. versionadded:: 4.4.0 + + Parameters + ---------- + agg : :class:`Aggregator` + The aggregator instance. + + Returns + ------- + function + A callable that, applied to input columns, produces an aggregate :class:`Column`. + + Raises + ------ + :class:`PySparkImportError` + If a supported version of PyArrow is not installed. + :class:`PySparkTypeError` + If ``agg`` is not an :class:`Aggregator`, or its ``bufferSchema`` is not a + :class:`StructType`. + """ + from pyspark.sql.pandas.utils import require_minimum_pyarrow_version + from pyspark.sql.utils import is_remote + + require_minimum_pyarrow_version() + + if is_remote(): + from pyspark.sql.connect.udf import UserDefinedFunction + else: + # The classic UserDefinedFunction is a distinct class from the Connect one above; + # both provide the same interface used below, so silence mypy's reassignment check. + from pyspark.sql.udf import UserDefinedFunction # type: ignore[assignment] + + if not isinstance(agg, Aggregator): + raise PySparkTypeError( + errorClass="NOT_EXPECTED_TYPE", + messageParameters={ + "arg_name": "agg", + "expected_type": "Aggregator", + "arg_type": type(agg).__name__, + }, + ) + if not isinstance(agg.bufferSchema, StructType): + raise PySparkTypeError( + errorClass="NOT_EXPECTED_TYPE", + messageParameters={ + "arg_name": "bufferSchema", + "expected_type": "StructType", + "arg_type": type(agg.bufferSchema).__name__, + }, + ) + + udf_obj = UserDefinedFunction( + agg, + returnType=agg.outputType, + name=agg.__class__.__name__, + evalType=PythonEvalType.SQL_GROUPED_AGG_ARROW_INCREMENTAL_FINAL_UDF, + deterministic=True, + ) + # Threaded to the JVM in UserDefinedFunction._create_judf so PythonAggregate can plan the + # two-stage aggregation. Set on both the UDF and its wrapper so it survives + # ``spark.udf.register`` (which reconstructs the UDF from the wrapper). + udf_obj.bufferSchema = agg.bufferSchema # type: ignore[attr-defined] + wrapped = udf_obj._wrapped() + wrapped.bufferSchema = agg.bufferSchema # type: ignore[attr-defined] + return wrapped diff --git a/python/pyspark/sql/connect/client/core.py b/python/pyspark/sql/connect/client/core.py index 0ff6f0ae90f94..462ad8fe2670c 100644 --- a/python/pyspark/sql/connect/client/core.py +++ b/python/pyspark/sql/connect/client/core.py @@ -1081,6 +1081,7 @@ def register_udf( name: Optional[str] = None, eval_type: int = PythonEvalType.SQL_BATCHED_UDF, deterministic: bool = True, + buffer_type: Optional["DataType"] = None, ) -> str: """ Create a temporary UDF in the session catalog on the other side. We generate a @@ -1096,6 +1097,8 @@ def register_udf( eval_type=eval_type, func=function, python_ver="%d.%d" % sys.version_info[:2], + # Set for the incremental aggregator (see pyspark.sql.aggregator). + buffer_type=buffer_type, ) # construct a CommonInlineUserDefinedFunction diff --git a/python/pyspark/sql/connect/expressions.py b/python/pyspark/sql/connect/expressions.py index 57270398118f7..bbe0dfa394783 100644 --- a/python/pyspark/sql/connect/expressions.py +++ b/python/pyspark/sql/connect/expressions.py @@ -741,6 +741,7 @@ def __init__( eval_type: int, func: Callable[..., Any], python_ver: str, + buffer_type: Optional[DataType] = None, ) -> None: self._output_type: DataType = ( UnparsedDataType(output_type) if isinstance(output_type, str) else output_type @@ -748,6 +749,8 @@ def __init__( self._eval_type = eval_type self._func = func self._python_ver = python_ver + # Intermediate buffer schema for an incremental Python aggregator; None otherwise. + self._buffer_type = buffer_type def to_plan(self, session: "SparkConnectClient") -> proto.PythonUDF: if isinstance(self._output_type, UnparsedDataType): @@ -763,6 +766,8 @@ def to_plan(self, session: "SparkConnectClient") -> proto.PythonUDF: expr.eval_type = self._eval_type expr.command = CloudPickleSerializer().dumps((self._func, output_type)) expr.python_ver = self._python_ver + if self._buffer_type is not None: + expr.buffer_type.CopyFrom(pyspark_types_to_proto_types(self._buffer_type)) return expr def __repr__(self) -> str: diff --git a/python/pyspark/sql/connect/proto/expressions_pb2.py b/python/pyspark/sql/connect/proto/expressions_pb2.py index aa51c393c043f..5bb6335cfe3b9 100644 --- a/python/pyspark/sql/connect/proto/expressions_pb2.py +++ b/python/pyspark/sql/connect/proto/expressions_pb2.py @@ -41,7 +41,7 @@ DESCRIPTOR = _descriptor_pool.Default().AddSerializedFile( - b'\n\x1fspark/connect/expressions.proto\x12\rspark.connect\x1a\x19google/protobuf/any.proto\x1a\x19spark/connect/types.proto\x1a\x1aspark/connect/common.proto"\x90<\n\nExpression\x12\x37\n\x06\x63ommon\x18\x12 \x01(\x0b\x32\x1f.spark.connect.ExpressionCommonR\x06\x63ommon\x12=\n\x07literal\x18\x01 \x01(\x0b\x32!.spark.connect.Expression.LiteralH\x00R\x07literal\x12\x62\n\x14unresolved_attribute\x18\x02 \x01(\x0b\x32-.spark.connect.Expression.UnresolvedAttributeH\x00R\x13unresolvedAttribute\x12_\n\x13unresolved_function\x18\x03 \x01(\x0b\x32,.spark.connect.Expression.UnresolvedFunctionH\x00R\x12unresolvedFunction\x12Y\n\x11\x65xpression_string\x18\x04 \x01(\x0b\x32*.spark.connect.Expression.ExpressionStringH\x00R\x10\x65xpressionString\x12S\n\x0funresolved_star\x18\x05 \x01(\x0b\x32(.spark.connect.Expression.UnresolvedStarH\x00R\x0eunresolvedStar\x12\x37\n\x05\x61lias\x18\x06 \x01(\x0b\x32\x1f.spark.connect.Expression.AliasH\x00R\x05\x61lias\x12\x34\n\x04\x63\x61st\x18\x07 \x01(\x0b\x32\x1e.spark.connect.Expression.CastH\x00R\x04\x63\x61st\x12V\n\x10unresolved_regex\x18\x08 \x01(\x0b\x32).spark.connect.Expression.UnresolvedRegexH\x00R\x0funresolvedRegex\x12\x44\n\nsort_order\x18\t \x01(\x0b\x32#.spark.connect.Expression.SortOrderH\x00R\tsortOrder\x12S\n\x0flambda_function\x18\n \x01(\x0b\x32(.spark.connect.Expression.LambdaFunctionH\x00R\x0elambdaFunction\x12:\n\x06window\x18\x0b \x01(\x0b\x32 .spark.connect.Expression.WindowH\x00R\x06window\x12l\n\x18unresolved_extract_value\x18\x0c \x01(\x0b\x32\x30.spark.connect.Expression.UnresolvedExtractValueH\x00R\x16unresolvedExtractValue\x12M\n\rupdate_fields\x18\r \x01(\x0b\x32&.spark.connect.Expression.UpdateFieldsH\x00R\x0cupdateFields\x12\x82\x01\n unresolved_named_lambda_variable\x18\x0e \x01(\x0b\x32\x37.spark.connect.Expression.UnresolvedNamedLambdaVariableH\x00R\x1dunresolvedNamedLambdaVariable\x12~\n#common_inline_user_defined_function\x18\x0f \x01(\x0b\x32..spark.connect.CommonInlineUserDefinedFunctionH\x00R\x1f\x63ommonInlineUserDefinedFunction\x12\x42\n\rcall_function\x18\x10 \x01(\x0b\x32\x1b.spark.connect.CallFunctionH\x00R\x0c\x63\x61llFunction\x12\x64\n\x19named_argument_expression\x18\x11 \x01(\x0b\x32&.spark.connect.NamedArgumentExpressionH\x00R\x17namedArgumentExpression\x12?\n\x0cmerge_action\x18\x13 \x01(\x0b\x32\x1a.spark.connect.MergeActionH\x00R\x0bmergeAction\x12g\n\x1atyped_aggregate_expression\x18\x14 \x01(\x0b\x32\'.spark.connect.TypedAggregateExpressionH\x00R\x18typedAggregateExpression\x12T\n\x13subquery_expression\x18\x15 \x01(\x0b\x32!.spark.connect.SubqueryExpressionH\x00R\x12subqueryExpression\x12s\n\x1b\x64irect_shuffle_partition_id\x18\x16 \x01(\x0b\x32\x32.spark.connect.Expression.DirectShufflePartitionIDH\x00R\x18\x64irectShufflePartitionId\x12\x35\n\textension\x18\xe7\x07 \x01(\x0b\x32\x14.google.protobuf.AnyH\x00R\textension\x1a\x8f\x06\n\x06Window\x12\x42\n\x0fwindow_function\x18\x01 \x01(\x0b\x32\x19.spark.connect.ExpressionR\x0ewindowFunction\x12@\n\x0epartition_spec\x18\x02 \x03(\x0b\x32\x19.spark.connect.ExpressionR\rpartitionSpec\x12\x42\n\norder_spec\x18\x03 \x03(\x0b\x32#.spark.connect.Expression.SortOrderR\torderSpec\x12K\n\nframe_spec\x18\x04 \x01(\x0b\x32,.spark.connect.Expression.Window.WindowFrameR\tframeSpec\x1a\xed\x03\n\x0bWindowFrame\x12U\n\nframe_type\x18\x01 \x01(\x0e\x32\x36.spark.connect.Expression.Window.WindowFrame.FrameTypeR\tframeType\x12P\n\x05lower\x18\x02 \x01(\x0b\x32:.spark.connect.Expression.Window.WindowFrame.FrameBoundaryR\x05lower\x12P\n\x05upper\x18\x03 \x01(\x0b\x32:.spark.connect.Expression.Window.WindowFrame.FrameBoundaryR\x05upper\x1a\x91\x01\n\rFrameBoundary\x12!\n\x0b\x63urrent_row\x18\x01 \x01(\x08H\x00R\ncurrentRow\x12\x1e\n\tunbounded\x18\x02 \x01(\x08H\x00R\tunbounded\x12\x31\n\x05value\x18\x03 \x01(\x0b\x32\x19.spark.connect.ExpressionH\x00R\x05valueB\n\n\x08\x62oundary"O\n\tFrameType\x12\x18\n\x14\x46RAME_TYPE_UNDEFINED\x10\x00\x12\x12\n\x0e\x46RAME_TYPE_ROW\x10\x01\x12\x14\n\x10\x46RAME_TYPE_RANGE\x10\x02\x1a\xa9\x03\n\tSortOrder\x12/\n\x05\x63hild\x18\x01 \x01(\x0b\x32\x19.spark.connect.ExpressionR\x05\x63hild\x12O\n\tdirection\x18\x02 \x01(\x0e\x32\x31.spark.connect.Expression.SortOrder.SortDirectionR\tdirection\x12U\n\rnull_ordering\x18\x03 \x01(\x0e\x32\x30.spark.connect.Expression.SortOrder.NullOrderingR\x0cnullOrdering"l\n\rSortDirection\x12\x1e\n\x1aSORT_DIRECTION_UNSPECIFIED\x10\x00\x12\x1c\n\x18SORT_DIRECTION_ASCENDING\x10\x01\x12\x1d\n\x19SORT_DIRECTION_DESCENDING\x10\x02"U\n\x0cNullOrdering\x12\x1a\n\x16SORT_NULLS_UNSPECIFIED\x10\x00\x12\x14\n\x10SORT_NULLS_FIRST\x10\x01\x12\x13\n\x0fSORT_NULLS_LAST\x10\x02\x1aK\n\x18\x44irectShufflePartitionID\x12/\n\x05\x63hild\x18\x01 \x01(\x0b\x32\x19.spark.connect.ExpressionR\x05\x63hild\x1a\xbb\x02\n\x04\x43\x61st\x12-\n\x04\x65xpr\x18\x01 \x01(\x0b\x32\x19.spark.connect.ExpressionR\x04\x65xpr\x12-\n\x04type\x18\x02 \x01(\x0b\x32\x17.spark.connect.DataTypeH\x00R\x04type\x12\x1b\n\x08type_str\x18\x03 \x01(\tH\x00R\x07typeStr\x12\x44\n\teval_mode\x18\x04 \x01(\x0e\x32\'.spark.connect.Expression.Cast.EvalModeR\x08\x65valMode"b\n\x08\x45valMode\x12\x19\n\x15\x45VAL_MODE_UNSPECIFIED\x10\x00\x12\x14\n\x10\x45VAL_MODE_LEGACY\x10\x01\x12\x12\n\x0e\x45VAL_MODE_ANSI\x10\x02\x12\x11\n\rEVAL_MODE_TRY\x10\x03\x42\x0e\n\x0c\x63\x61st_to_type\x1a\x9c\x15\n\x07Literal\x12-\n\x04null\x18\x01 \x01(\x0b\x32\x17.spark.connect.DataTypeH\x00R\x04null\x12\x18\n\x06\x62inary\x18\x02 \x01(\x0cH\x00R\x06\x62inary\x12\x1a\n\x07\x62oolean\x18\x03 \x01(\x08H\x00R\x07\x62oolean\x12\x14\n\x04\x62yte\x18\x04 \x01(\x05H\x00R\x04\x62yte\x12\x16\n\x05short\x18\x05 \x01(\x05H\x00R\x05short\x12\x1a\n\x07integer\x18\x06 \x01(\x05H\x00R\x07integer\x12\x14\n\x04long\x18\x07 \x01(\x03H\x00R\x04long\x12\x16\n\x05\x66loat\x18\n \x01(\x02H\x00R\x05\x66loat\x12\x18\n\x06\x64ouble\x18\x0b \x01(\x01H\x00R\x06\x64ouble\x12\x45\n\x07\x64\x65\x63imal\x18\x0c \x01(\x0b\x32).spark.connect.Expression.Literal.DecimalH\x00R\x07\x64\x65\x63imal\x12\x18\n\x06string\x18\r \x01(\tH\x00R\x06string\x12\x14\n\x04\x64\x61te\x18\x10 \x01(\x05H\x00R\x04\x64\x61te\x12\x1e\n\ttimestamp\x18\x11 \x01(\x03H\x00R\ttimestamp\x12%\n\rtimestamp_ntz\x18\x12 \x01(\x03H\x00R\x0ctimestampNtz\x12\x61\n\x11\x63\x61lendar_interval\x18\x13 \x01(\x0b\x32\x32.spark.connect.Expression.Literal.CalendarIntervalH\x00R\x10\x63\x61lendarInterval\x12\x30\n\x13year_month_interval\x18\x14 \x01(\x05H\x00R\x11yearMonthInterval\x12,\n\x11\x64\x61y_time_interval\x18\x15 \x01(\x03H\x00R\x0f\x64\x61yTimeInterval\x12?\n\x05\x61rray\x18\x16 \x01(\x0b\x32\'.spark.connect.Expression.Literal.ArrayH\x00R\x05\x61rray\x12\x39\n\x03map\x18\x17 \x01(\x0b\x32%.spark.connect.Expression.Literal.MapH\x00R\x03map\x12\x42\n\x06struct\x18\x18 \x01(\x0b\x32(.spark.connect.Expression.Literal.StructH\x00R\x06struct\x12\x61\n\x11specialized_array\x18\x19 \x01(\x0b\x32\x32.spark.connect.Expression.Literal.SpecializedArrayH\x00R\x10specializedArray\x12<\n\x04time\x18\x1a \x01(\x0b\x32&.spark.connect.Expression.Literal.TimeH\x00R\x04time\x12\x65\n\x13timestamp_ntz_nanos\x18\x1d \x01(\x0b\x32\x33.spark.connect.Expression.Literal.TimestampNTZNanosH\x00R\x11timestampNtzNanos\x12\x65\n\x13timestamp_ltz_nanos\x18\x1e \x01(\x0b\x32\x33.spark.connect.Expression.Literal.TimestampLTZNanosH\x00R\x11timestampLtzNanos\x12\x34\n\tdata_type\x18\x64 \x01(\x0b\x32\x17.spark.connect.DataTypeR\x08\x64\x61taType\x1au\n\x07\x44\x65\x63imal\x12\x14\n\x05value\x18\x01 \x01(\tR\x05value\x12!\n\tprecision\x18\x02 \x01(\x05H\x00R\tprecision\x88\x01\x01\x12\x19\n\x05scale\x18\x03 \x01(\x05H\x01R\x05scale\x88\x01\x01\x42\x0c\n\n_precisionB\x08\n\x06_scale\x1a\x62\n\x10\x43\x61lendarInterval\x12\x16\n\x06months\x18\x01 \x01(\x05R\x06months\x12\x12\n\x04\x64\x61ys\x18\x02 \x01(\x05R\x04\x64\x61ys\x12"\n\x0cmicroseconds\x18\x03 \x01(\x03R\x0cmicroseconds\x1a\x86\x01\n\x05\x41rray\x12>\n\x0c\x65lement_type\x18\x01 \x01(\x0b\x32\x17.spark.connect.DataTypeB\x02\x18\x01R\x0b\x65lementType\x12=\n\x08\x65lements\x18\x02 \x03(\x0b\x32!.spark.connect.Expression.LiteralR\x08\x65lements\x1a\xeb\x01\n\x03Map\x12\x36\n\x08key_type\x18\x01 \x01(\x0b\x32\x17.spark.connect.DataTypeB\x02\x18\x01R\x07keyType\x12:\n\nvalue_type\x18\x02 \x01(\x0b\x32\x17.spark.connect.DataTypeB\x02\x18\x01R\tvalueType\x12\x35\n\x04keys\x18\x03 \x03(\x0b\x32!.spark.connect.Expression.LiteralR\x04keys\x12\x39\n\x06values\x18\x04 \x03(\x0b\x32!.spark.connect.Expression.LiteralR\x06values\x1a\x85\x01\n\x06Struct\x12<\n\x0bstruct_type\x18\x01 \x01(\x0b\x32\x17.spark.connect.DataTypeB\x02\x18\x01R\nstructType\x12=\n\x08\x65lements\x18\x02 \x03(\x0b\x32!.spark.connect.Expression.LiteralR\x08\x65lements\x1a\xc0\x02\n\x10SpecializedArray\x12,\n\x05\x62ools\x18\x01 \x01(\x0b\x32\x14.spark.connect.BoolsH\x00R\x05\x62ools\x12)\n\x04ints\x18\x02 \x01(\x0b\x32\x13.spark.connect.IntsH\x00R\x04ints\x12,\n\x05longs\x18\x03 \x01(\x0b\x32\x14.spark.connect.LongsH\x00R\x05longs\x12/\n\x06\x66loats\x18\x04 \x01(\x0b\x32\x15.spark.connect.FloatsH\x00R\x06\x66loats\x12\x32\n\x07\x64oubles\x18\x05 \x01(\x0b\x32\x16.spark.connect.DoublesH\x00R\x07\x64oubles\x12\x32\n\x07strings\x18\x06 \x01(\x0b\x32\x16.spark.connect.StringsH\x00R\x07stringsB\x0c\n\nvalue_type\x1aK\n\x04Time\x12\x12\n\x04nano\x18\x01 \x01(\x03R\x04nano\x12!\n\tprecision\x18\x02 \x01(\x05H\x00R\tprecision\x88\x01\x01\x42\x0c\n\n_precision\x1a\x95\x01\n\x11TimestampNTZNanos\x12!\n\x0c\x65poch_micros\x18\x01 \x01(\x03R\x0b\x65pochMicros\x12,\n\x12nanos_within_micro\x18\x02 \x01(\x05R\x10nanosWithinMicro\x12!\n\tprecision\x18\x03 \x01(\x05H\x00R\tprecision\x88\x01\x01\x42\x0c\n\n_precision\x1a\x95\x01\n\x11TimestampLTZNanos\x12!\n\x0c\x65poch_micros\x18\x01 \x01(\x03R\x0b\x65pochMicros\x12,\n\x12nanos_within_micro\x18\x02 \x01(\x05R\x10nanosWithinMicro\x12!\n\tprecision\x18\x03 \x01(\x05H\x00R\tprecision\x88\x01\x01\x42\x0c\n\n_precisionB\x0e\n\x0cliteral_typeJ\x04\x08\x1b\x10\x1cJ\x04\x08\x1c\x10\x1d\x1a\xba\x01\n\x13UnresolvedAttribute\x12/\n\x13unparsed_identifier\x18\x01 \x01(\tR\x12unparsedIdentifier\x12\x1c\n\x07plan_id\x18\x02 \x01(\x03H\x00R\x06planId\x88\x01\x01\x12\x31\n\x12is_metadata_column\x18\x03 \x01(\x08H\x01R\x10isMetadataColumn\x88\x01\x01\x42\n\n\x08_plan_idB\x15\n\x13_is_metadata_column\x1a\x82\x02\n\x12UnresolvedFunction\x12#\n\rfunction_name\x18\x01 \x01(\tR\x0c\x66unctionName\x12\x37\n\targuments\x18\x02 \x03(\x0b\x32\x19.spark.connect.ExpressionR\targuments\x12\x1f\n\x0bis_distinct\x18\x03 \x01(\x08R\nisDistinct\x12\x37\n\x18is_user_defined_function\x18\x04 \x01(\x08R\x15isUserDefinedFunction\x12$\n\x0bis_internal\x18\x05 \x01(\x08H\x00R\nisInternal\x88\x01\x01\x42\x0e\n\x0c_is_internal\x1a\x32\n\x10\x45xpressionString\x12\x1e\n\nexpression\x18\x01 \x01(\tR\nexpression\x1a|\n\x0eUnresolvedStar\x12,\n\x0funparsed_target\x18\x01 \x01(\tH\x00R\x0eunparsedTarget\x88\x01\x01\x12\x1c\n\x07plan_id\x18\x02 \x01(\x03H\x01R\x06planId\x88\x01\x01\x42\x12\n\x10_unparsed_targetB\n\n\x08_plan_id\x1aV\n\x0fUnresolvedRegex\x12\x19\n\x08\x63ol_name\x18\x01 \x01(\tR\x07\x63olName\x12\x1c\n\x07plan_id\x18\x02 \x01(\x03H\x00R\x06planId\x88\x01\x01\x42\n\n\x08_plan_id\x1a\x84\x01\n\x16UnresolvedExtractValue\x12/\n\x05\x63hild\x18\x01 \x01(\x0b\x32\x19.spark.connect.ExpressionR\x05\x63hild\x12\x39\n\nextraction\x18\x02 \x01(\x0b\x32\x19.spark.connect.ExpressionR\nextraction\x1a\xbb\x01\n\x0cUpdateFields\x12\x46\n\x11struct_expression\x18\x01 \x01(\x0b\x32\x19.spark.connect.ExpressionR\x10structExpression\x12\x1d\n\nfield_name\x18\x02 \x01(\tR\tfieldName\x12\x44\n\x10value_expression\x18\x03 \x01(\x0b\x32\x19.spark.connect.ExpressionR\x0fvalueExpression\x1ax\n\x05\x41lias\x12-\n\x04\x65xpr\x18\x01 \x01(\x0b\x32\x19.spark.connect.ExpressionR\x04\x65xpr\x12\x12\n\x04name\x18\x02 \x03(\tR\x04name\x12\x1f\n\x08metadata\x18\x03 \x01(\tH\x00R\x08metadata\x88\x01\x01\x42\x0b\n\t_metadata\x1a\x9e\x01\n\x0eLambdaFunction\x12\x35\n\x08\x66unction\x18\x01 \x01(\x0b\x32\x19.spark.connect.ExpressionR\x08\x66unction\x12U\n\targuments\x18\x02 \x03(\x0b\x32\x37.spark.connect.Expression.UnresolvedNamedLambdaVariableR\targuments\x1a>\n\x1dUnresolvedNamedLambdaVariable\x12\x1d\n\nname_parts\x18\x01 \x03(\tR\tnamePartsB\x0b\n\texpr_type"A\n\x10\x45xpressionCommon\x12-\n\x06origin\x18\x01 \x01(\x0b\x32\x15.spark.connect.OriginR\x06origin"\x8d\x03\n\x1f\x43ommonInlineUserDefinedFunction\x12#\n\rfunction_name\x18\x01 \x01(\tR\x0c\x66unctionName\x12$\n\rdeterministic\x18\x02 \x01(\x08R\rdeterministic\x12\x37\n\targuments\x18\x03 \x03(\x0b\x32\x19.spark.connect.ExpressionR\targuments\x12\x39\n\npython_udf\x18\x04 \x01(\x0b\x32\x18.spark.connect.PythonUDFH\x00R\tpythonUdf\x12I\n\x10scalar_scala_udf\x18\x05 \x01(\x0b\x32\x1d.spark.connect.ScalarScalaUDFH\x00R\x0escalarScalaUdf\x12\x33\n\x08java_udf\x18\x06 \x01(\x0b\x32\x16.spark.connect.JavaUDFH\x00R\x07javaUdf\x12\x1f\n\x0bis_distinct\x18\x07 \x01(\x08R\nisDistinctB\n\n\x08\x66unction"\xcc\x01\n\tPythonUDF\x12\x38\n\x0boutput_type\x18\x01 \x01(\x0b\x32\x17.spark.connect.DataTypeR\noutputType\x12\x1b\n\teval_type\x18\x02 \x01(\x05R\x08\x65valType\x12\x18\n\x07\x63ommand\x18\x03 \x01(\x0cR\x07\x63ommand\x12\x1d\n\npython_ver\x18\x04 \x01(\tR\tpythonVer\x12/\n\x13\x61\x64\x64itional_includes\x18\x05 \x03(\tR\x12\x61\x64\x64itionalIncludes"\xd6\x01\n\x0eScalarScalaUDF\x12\x18\n\x07payload\x18\x01 \x01(\x0cR\x07payload\x12\x37\n\ninputTypes\x18\x02 \x03(\x0b\x32\x17.spark.connect.DataTypeR\ninputTypes\x12\x37\n\noutputType\x18\x03 \x01(\x0b\x32\x17.spark.connect.DataTypeR\noutputType\x12\x1a\n\x08nullable\x18\x04 \x01(\x08R\x08nullable\x12\x1c\n\taggregate\x18\x05 \x01(\x08R\taggregate"\x95\x01\n\x07JavaUDF\x12\x1d\n\nclass_name\x18\x01 \x01(\tR\tclassName\x12=\n\x0boutput_type\x18\x02 \x01(\x0b\x32\x17.spark.connect.DataTypeH\x00R\noutputType\x88\x01\x01\x12\x1c\n\taggregate\x18\x03 \x01(\x08R\taggregateB\x0e\n\x0c_output_type"c\n\x18TypedAggregateExpression\x12G\n\x10scalar_scala_udf\x18\x01 \x01(\x0b\x32\x1d.spark.connect.ScalarScalaUDFR\x0escalarScalaUdf"l\n\x0c\x43\x61llFunction\x12#\n\rfunction_name\x18\x01 \x01(\tR\x0c\x66unctionName\x12\x37\n\targuments\x18\x02 \x03(\x0b\x32\x19.spark.connect.ExpressionR\targuments"\\\n\x17NamedArgumentExpression\x12\x10\n\x03key\x18\x01 \x01(\tR\x03key\x12/\n\x05value\x18\x02 \x01(\x0b\x32\x19.spark.connect.ExpressionR\x05value"\x80\x04\n\x0bMergeAction\x12\x46\n\x0b\x61\x63tion_type\x18\x01 \x01(\x0e\x32%.spark.connect.MergeAction.ActionTypeR\nactionType\x12<\n\tcondition\x18\x02 \x01(\x0b\x32\x19.spark.connect.ExpressionH\x00R\tcondition\x88\x01\x01\x12G\n\x0b\x61ssignments\x18\x03 \x03(\x0b\x32%.spark.connect.MergeAction.AssignmentR\x0b\x61ssignments\x1aj\n\nAssignment\x12+\n\x03key\x18\x01 \x01(\x0b\x32\x19.spark.connect.ExpressionR\x03key\x12/\n\x05value\x18\x02 \x01(\x0b\x32\x19.spark.connect.ExpressionR\x05value"\xa7\x01\n\nActionType\x12\x17\n\x13\x41\x43TION_TYPE_INVALID\x10\x00\x12\x16\n\x12\x41\x43TION_TYPE_DELETE\x10\x01\x12\x16\n\x12\x41\x43TION_TYPE_INSERT\x10\x02\x12\x1b\n\x17\x41\x43TION_TYPE_INSERT_STAR\x10\x03\x12\x16\n\x12\x41\x43TION_TYPE_UPDATE\x10\x04\x12\x1b\n\x17\x41\x43TION_TYPE_UPDATE_STAR\x10\x05\x42\x0c\n\n_condition"\xc5\x05\n\x12SubqueryExpression\x12\x17\n\x07plan_id\x18\x01 \x01(\x03R\x06planId\x12S\n\rsubquery_type\x18\x02 \x01(\x0e\x32..spark.connect.SubqueryExpression.SubqueryTypeR\x0csubqueryType\x12\x62\n\x11table_arg_options\x18\x03 \x01(\x0b\x32\x31.spark.connect.SubqueryExpression.TableArgOptionsH\x00R\x0ftableArgOptions\x88\x01\x01\x12G\n\x12in_subquery_values\x18\x04 \x03(\x0b\x32\x19.spark.connect.ExpressionR\x10inSubqueryValues\x1a\xea\x01\n\x0fTableArgOptions\x12@\n\x0epartition_spec\x18\x01 \x03(\x0b\x32\x19.spark.connect.ExpressionR\rpartitionSpec\x12\x42\n\norder_spec\x18\x02 \x03(\x0b\x32#.spark.connect.Expression.SortOrderR\torderSpec\x12\x37\n\x15with_single_partition\x18\x03 \x01(\x08H\x00R\x13withSinglePartition\x88\x01\x01\x42\x18\n\x16_with_single_partition"\x90\x01\n\x0cSubqueryType\x12\x19\n\x15SUBQUERY_TYPE_UNKNOWN\x10\x00\x12\x18\n\x14SUBQUERY_TYPE_SCALAR\x10\x01\x12\x18\n\x14SUBQUERY_TYPE_EXISTS\x10\x02\x12\x1b\n\x17SUBQUERY_TYPE_TABLE_ARG\x10\x03\x12\x14\n\x10SUBQUERY_TYPE_IN\x10\x04\x42\x14\n\x12_table_arg_optionsB6\n\x1eorg.apache.spark.connect.protoP\x01Z\x12internal/generatedb\x06proto3' + b'\n\x1fspark/connect/expressions.proto\x12\rspark.connect\x1a\x19google/protobuf/any.proto\x1a\x19spark/connect/types.proto\x1a\x1aspark/connect/common.proto"\x90<\n\nExpression\x12\x37\n\x06\x63ommon\x18\x12 \x01(\x0b\x32\x1f.spark.connect.ExpressionCommonR\x06\x63ommon\x12=\n\x07literal\x18\x01 \x01(\x0b\x32!.spark.connect.Expression.LiteralH\x00R\x07literal\x12\x62\n\x14unresolved_attribute\x18\x02 \x01(\x0b\x32-.spark.connect.Expression.UnresolvedAttributeH\x00R\x13unresolvedAttribute\x12_\n\x13unresolved_function\x18\x03 \x01(\x0b\x32,.spark.connect.Expression.UnresolvedFunctionH\x00R\x12unresolvedFunction\x12Y\n\x11\x65xpression_string\x18\x04 \x01(\x0b\x32*.spark.connect.Expression.ExpressionStringH\x00R\x10\x65xpressionString\x12S\n\x0funresolved_star\x18\x05 \x01(\x0b\x32(.spark.connect.Expression.UnresolvedStarH\x00R\x0eunresolvedStar\x12\x37\n\x05\x61lias\x18\x06 \x01(\x0b\x32\x1f.spark.connect.Expression.AliasH\x00R\x05\x61lias\x12\x34\n\x04\x63\x61st\x18\x07 \x01(\x0b\x32\x1e.spark.connect.Expression.CastH\x00R\x04\x63\x61st\x12V\n\x10unresolved_regex\x18\x08 \x01(\x0b\x32).spark.connect.Expression.UnresolvedRegexH\x00R\x0funresolvedRegex\x12\x44\n\nsort_order\x18\t \x01(\x0b\x32#.spark.connect.Expression.SortOrderH\x00R\tsortOrder\x12S\n\x0flambda_function\x18\n \x01(\x0b\x32(.spark.connect.Expression.LambdaFunctionH\x00R\x0elambdaFunction\x12:\n\x06window\x18\x0b \x01(\x0b\x32 .spark.connect.Expression.WindowH\x00R\x06window\x12l\n\x18unresolved_extract_value\x18\x0c \x01(\x0b\x32\x30.spark.connect.Expression.UnresolvedExtractValueH\x00R\x16unresolvedExtractValue\x12M\n\rupdate_fields\x18\r \x01(\x0b\x32&.spark.connect.Expression.UpdateFieldsH\x00R\x0cupdateFields\x12\x82\x01\n unresolved_named_lambda_variable\x18\x0e \x01(\x0b\x32\x37.spark.connect.Expression.UnresolvedNamedLambdaVariableH\x00R\x1dunresolvedNamedLambdaVariable\x12~\n#common_inline_user_defined_function\x18\x0f \x01(\x0b\x32..spark.connect.CommonInlineUserDefinedFunctionH\x00R\x1f\x63ommonInlineUserDefinedFunction\x12\x42\n\rcall_function\x18\x10 \x01(\x0b\x32\x1b.spark.connect.CallFunctionH\x00R\x0c\x63\x61llFunction\x12\x64\n\x19named_argument_expression\x18\x11 \x01(\x0b\x32&.spark.connect.NamedArgumentExpressionH\x00R\x17namedArgumentExpression\x12?\n\x0cmerge_action\x18\x13 \x01(\x0b\x32\x1a.spark.connect.MergeActionH\x00R\x0bmergeAction\x12g\n\x1atyped_aggregate_expression\x18\x14 \x01(\x0b\x32\'.spark.connect.TypedAggregateExpressionH\x00R\x18typedAggregateExpression\x12T\n\x13subquery_expression\x18\x15 \x01(\x0b\x32!.spark.connect.SubqueryExpressionH\x00R\x12subqueryExpression\x12s\n\x1b\x64irect_shuffle_partition_id\x18\x16 \x01(\x0b\x32\x32.spark.connect.Expression.DirectShufflePartitionIDH\x00R\x18\x64irectShufflePartitionId\x12\x35\n\textension\x18\xe7\x07 \x01(\x0b\x32\x14.google.protobuf.AnyH\x00R\textension\x1a\x8f\x06\n\x06Window\x12\x42\n\x0fwindow_function\x18\x01 \x01(\x0b\x32\x19.spark.connect.ExpressionR\x0ewindowFunction\x12@\n\x0epartition_spec\x18\x02 \x03(\x0b\x32\x19.spark.connect.ExpressionR\rpartitionSpec\x12\x42\n\norder_spec\x18\x03 \x03(\x0b\x32#.spark.connect.Expression.SortOrderR\torderSpec\x12K\n\nframe_spec\x18\x04 \x01(\x0b\x32,.spark.connect.Expression.Window.WindowFrameR\tframeSpec\x1a\xed\x03\n\x0bWindowFrame\x12U\n\nframe_type\x18\x01 \x01(\x0e\x32\x36.spark.connect.Expression.Window.WindowFrame.FrameTypeR\tframeType\x12P\n\x05lower\x18\x02 \x01(\x0b\x32:.spark.connect.Expression.Window.WindowFrame.FrameBoundaryR\x05lower\x12P\n\x05upper\x18\x03 \x01(\x0b\x32:.spark.connect.Expression.Window.WindowFrame.FrameBoundaryR\x05upper\x1a\x91\x01\n\rFrameBoundary\x12!\n\x0b\x63urrent_row\x18\x01 \x01(\x08H\x00R\ncurrentRow\x12\x1e\n\tunbounded\x18\x02 \x01(\x08H\x00R\tunbounded\x12\x31\n\x05value\x18\x03 \x01(\x0b\x32\x19.spark.connect.ExpressionH\x00R\x05valueB\n\n\x08\x62oundary"O\n\tFrameType\x12\x18\n\x14\x46RAME_TYPE_UNDEFINED\x10\x00\x12\x12\n\x0e\x46RAME_TYPE_ROW\x10\x01\x12\x14\n\x10\x46RAME_TYPE_RANGE\x10\x02\x1a\xa9\x03\n\tSortOrder\x12/\n\x05\x63hild\x18\x01 \x01(\x0b\x32\x19.spark.connect.ExpressionR\x05\x63hild\x12O\n\tdirection\x18\x02 \x01(\x0e\x32\x31.spark.connect.Expression.SortOrder.SortDirectionR\tdirection\x12U\n\rnull_ordering\x18\x03 \x01(\x0e\x32\x30.spark.connect.Expression.SortOrder.NullOrderingR\x0cnullOrdering"l\n\rSortDirection\x12\x1e\n\x1aSORT_DIRECTION_UNSPECIFIED\x10\x00\x12\x1c\n\x18SORT_DIRECTION_ASCENDING\x10\x01\x12\x1d\n\x19SORT_DIRECTION_DESCENDING\x10\x02"U\n\x0cNullOrdering\x12\x1a\n\x16SORT_NULLS_UNSPECIFIED\x10\x00\x12\x14\n\x10SORT_NULLS_FIRST\x10\x01\x12\x13\n\x0fSORT_NULLS_LAST\x10\x02\x1aK\n\x18\x44irectShufflePartitionID\x12/\n\x05\x63hild\x18\x01 \x01(\x0b\x32\x19.spark.connect.ExpressionR\x05\x63hild\x1a\xbb\x02\n\x04\x43\x61st\x12-\n\x04\x65xpr\x18\x01 \x01(\x0b\x32\x19.spark.connect.ExpressionR\x04\x65xpr\x12-\n\x04type\x18\x02 \x01(\x0b\x32\x17.spark.connect.DataTypeH\x00R\x04type\x12\x1b\n\x08type_str\x18\x03 \x01(\tH\x00R\x07typeStr\x12\x44\n\teval_mode\x18\x04 \x01(\x0e\x32\'.spark.connect.Expression.Cast.EvalModeR\x08\x65valMode"b\n\x08\x45valMode\x12\x19\n\x15\x45VAL_MODE_UNSPECIFIED\x10\x00\x12\x14\n\x10\x45VAL_MODE_LEGACY\x10\x01\x12\x12\n\x0e\x45VAL_MODE_ANSI\x10\x02\x12\x11\n\rEVAL_MODE_TRY\x10\x03\x42\x0e\n\x0c\x63\x61st_to_type\x1a\x9c\x15\n\x07Literal\x12-\n\x04null\x18\x01 \x01(\x0b\x32\x17.spark.connect.DataTypeH\x00R\x04null\x12\x18\n\x06\x62inary\x18\x02 \x01(\x0cH\x00R\x06\x62inary\x12\x1a\n\x07\x62oolean\x18\x03 \x01(\x08H\x00R\x07\x62oolean\x12\x14\n\x04\x62yte\x18\x04 \x01(\x05H\x00R\x04\x62yte\x12\x16\n\x05short\x18\x05 \x01(\x05H\x00R\x05short\x12\x1a\n\x07integer\x18\x06 \x01(\x05H\x00R\x07integer\x12\x14\n\x04long\x18\x07 \x01(\x03H\x00R\x04long\x12\x16\n\x05\x66loat\x18\n \x01(\x02H\x00R\x05\x66loat\x12\x18\n\x06\x64ouble\x18\x0b \x01(\x01H\x00R\x06\x64ouble\x12\x45\n\x07\x64\x65\x63imal\x18\x0c \x01(\x0b\x32).spark.connect.Expression.Literal.DecimalH\x00R\x07\x64\x65\x63imal\x12\x18\n\x06string\x18\r \x01(\tH\x00R\x06string\x12\x14\n\x04\x64\x61te\x18\x10 \x01(\x05H\x00R\x04\x64\x61te\x12\x1e\n\ttimestamp\x18\x11 \x01(\x03H\x00R\ttimestamp\x12%\n\rtimestamp_ntz\x18\x12 \x01(\x03H\x00R\x0ctimestampNtz\x12\x61\n\x11\x63\x61lendar_interval\x18\x13 \x01(\x0b\x32\x32.spark.connect.Expression.Literal.CalendarIntervalH\x00R\x10\x63\x61lendarInterval\x12\x30\n\x13year_month_interval\x18\x14 \x01(\x05H\x00R\x11yearMonthInterval\x12,\n\x11\x64\x61y_time_interval\x18\x15 \x01(\x03H\x00R\x0f\x64\x61yTimeInterval\x12?\n\x05\x61rray\x18\x16 \x01(\x0b\x32\'.spark.connect.Expression.Literal.ArrayH\x00R\x05\x61rray\x12\x39\n\x03map\x18\x17 \x01(\x0b\x32%.spark.connect.Expression.Literal.MapH\x00R\x03map\x12\x42\n\x06struct\x18\x18 \x01(\x0b\x32(.spark.connect.Expression.Literal.StructH\x00R\x06struct\x12\x61\n\x11specialized_array\x18\x19 \x01(\x0b\x32\x32.spark.connect.Expression.Literal.SpecializedArrayH\x00R\x10specializedArray\x12<\n\x04time\x18\x1a \x01(\x0b\x32&.spark.connect.Expression.Literal.TimeH\x00R\x04time\x12\x65\n\x13timestamp_ntz_nanos\x18\x1d \x01(\x0b\x32\x33.spark.connect.Expression.Literal.TimestampNTZNanosH\x00R\x11timestampNtzNanos\x12\x65\n\x13timestamp_ltz_nanos\x18\x1e \x01(\x0b\x32\x33.spark.connect.Expression.Literal.TimestampLTZNanosH\x00R\x11timestampLtzNanos\x12\x34\n\tdata_type\x18\x64 \x01(\x0b\x32\x17.spark.connect.DataTypeR\x08\x64\x61taType\x1au\n\x07\x44\x65\x63imal\x12\x14\n\x05value\x18\x01 \x01(\tR\x05value\x12!\n\tprecision\x18\x02 \x01(\x05H\x00R\tprecision\x88\x01\x01\x12\x19\n\x05scale\x18\x03 \x01(\x05H\x01R\x05scale\x88\x01\x01\x42\x0c\n\n_precisionB\x08\n\x06_scale\x1a\x62\n\x10\x43\x61lendarInterval\x12\x16\n\x06months\x18\x01 \x01(\x05R\x06months\x12\x12\n\x04\x64\x61ys\x18\x02 \x01(\x05R\x04\x64\x61ys\x12"\n\x0cmicroseconds\x18\x03 \x01(\x03R\x0cmicroseconds\x1a\x86\x01\n\x05\x41rray\x12>\n\x0c\x65lement_type\x18\x01 \x01(\x0b\x32\x17.spark.connect.DataTypeB\x02\x18\x01R\x0b\x65lementType\x12=\n\x08\x65lements\x18\x02 \x03(\x0b\x32!.spark.connect.Expression.LiteralR\x08\x65lements\x1a\xeb\x01\n\x03Map\x12\x36\n\x08key_type\x18\x01 \x01(\x0b\x32\x17.spark.connect.DataTypeB\x02\x18\x01R\x07keyType\x12:\n\nvalue_type\x18\x02 \x01(\x0b\x32\x17.spark.connect.DataTypeB\x02\x18\x01R\tvalueType\x12\x35\n\x04keys\x18\x03 \x03(\x0b\x32!.spark.connect.Expression.LiteralR\x04keys\x12\x39\n\x06values\x18\x04 \x03(\x0b\x32!.spark.connect.Expression.LiteralR\x06values\x1a\x85\x01\n\x06Struct\x12<\n\x0bstruct_type\x18\x01 \x01(\x0b\x32\x17.spark.connect.DataTypeB\x02\x18\x01R\nstructType\x12=\n\x08\x65lements\x18\x02 \x03(\x0b\x32!.spark.connect.Expression.LiteralR\x08\x65lements\x1a\xc0\x02\n\x10SpecializedArray\x12,\n\x05\x62ools\x18\x01 \x01(\x0b\x32\x14.spark.connect.BoolsH\x00R\x05\x62ools\x12)\n\x04ints\x18\x02 \x01(\x0b\x32\x13.spark.connect.IntsH\x00R\x04ints\x12,\n\x05longs\x18\x03 \x01(\x0b\x32\x14.spark.connect.LongsH\x00R\x05longs\x12/\n\x06\x66loats\x18\x04 \x01(\x0b\x32\x15.spark.connect.FloatsH\x00R\x06\x66loats\x12\x32\n\x07\x64oubles\x18\x05 \x01(\x0b\x32\x16.spark.connect.DoublesH\x00R\x07\x64oubles\x12\x32\n\x07strings\x18\x06 \x01(\x0b\x32\x16.spark.connect.StringsH\x00R\x07stringsB\x0c\n\nvalue_type\x1aK\n\x04Time\x12\x12\n\x04nano\x18\x01 \x01(\x03R\x04nano\x12!\n\tprecision\x18\x02 \x01(\x05H\x00R\tprecision\x88\x01\x01\x42\x0c\n\n_precision\x1a\x95\x01\n\x11TimestampNTZNanos\x12!\n\x0c\x65poch_micros\x18\x01 \x01(\x03R\x0b\x65pochMicros\x12,\n\x12nanos_within_micro\x18\x02 \x01(\x05R\x10nanosWithinMicro\x12!\n\tprecision\x18\x03 \x01(\x05H\x00R\tprecision\x88\x01\x01\x42\x0c\n\n_precision\x1a\x95\x01\n\x11TimestampLTZNanos\x12!\n\x0c\x65poch_micros\x18\x01 \x01(\x03R\x0b\x65pochMicros\x12,\n\x12nanos_within_micro\x18\x02 \x01(\x05R\x10nanosWithinMicro\x12!\n\tprecision\x18\x03 \x01(\x05H\x00R\tprecision\x88\x01\x01\x42\x0c\n\n_precisionB\x0e\n\x0cliteral_typeJ\x04\x08\x1b\x10\x1cJ\x04\x08\x1c\x10\x1d\x1a\xba\x01\n\x13UnresolvedAttribute\x12/\n\x13unparsed_identifier\x18\x01 \x01(\tR\x12unparsedIdentifier\x12\x1c\n\x07plan_id\x18\x02 \x01(\x03H\x00R\x06planId\x88\x01\x01\x12\x31\n\x12is_metadata_column\x18\x03 \x01(\x08H\x01R\x10isMetadataColumn\x88\x01\x01\x42\n\n\x08_plan_idB\x15\n\x13_is_metadata_column\x1a\x82\x02\n\x12UnresolvedFunction\x12#\n\rfunction_name\x18\x01 \x01(\tR\x0c\x66unctionName\x12\x37\n\targuments\x18\x02 \x03(\x0b\x32\x19.spark.connect.ExpressionR\targuments\x12\x1f\n\x0bis_distinct\x18\x03 \x01(\x08R\nisDistinct\x12\x37\n\x18is_user_defined_function\x18\x04 \x01(\x08R\x15isUserDefinedFunction\x12$\n\x0bis_internal\x18\x05 \x01(\x08H\x00R\nisInternal\x88\x01\x01\x42\x0e\n\x0c_is_internal\x1a\x32\n\x10\x45xpressionString\x12\x1e\n\nexpression\x18\x01 \x01(\tR\nexpression\x1a|\n\x0eUnresolvedStar\x12,\n\x0funparsed_target\x18\x01 \x01(\tH\x00R\x0eunparsedTarget\x88\x01\x01\x12\x1c\n\x07plan_id\x18\x02 \x01(\x03H\x01R\x06planId\x88\x01\x01\x42\x12\n\x10_unparsed_targetB\n\n\x08_plan_id\x1aV\n\x0fUnresolvedRegex\x12\x19\n\x08\x63ol_name\x18\x01 \x01(\tR\x07\x63olName\x12\x1c\n\x07plan_id\x18\x02 \x01(\x03H\x00R\x06planId\x88\x01\x01\x42\n\n\x08_plan_id\x1a\x84\x01\n\x16UnresolvedExtractValue\x12/\n\x05\x63hild\x18\x01 \x01(\x0b\x32\x19.spark.connect.ExpressionR\x05\x63hild\x12\x39\n\nextraction\x18\x02 \x01(\x0b\x32\x19.spark.connect.ExpressionR\nextraction\x1a\xbb\x01\n\x0cUpdateFields\x12\x46\n\x11struct_expression\x18\x01 \x01(\x0b\x32\x19.spark.connect.ExpressionR\x10structExpression\x12\x1d\n\nfield_name\x18\x02 \x01(\tR\tfieldName\x12\x44\n\x10value_expression\x18\x03 \x01(\x0b\x32\x19.spark.connect.ExpressionR\x0fvalueExpression\x1ax\n\x05\x41lias\x12-\n\x04\x65xpr\x18\x01 \x01(\x0b\x32\x19.spark.connect.ExpressionR\x04\x65xpr\x12\x12\n\x04name\x18\x02 \x03(\tR\x04name\x12\x1f\n\x08metadata\x18\x03 \x01(\tH\x00R\x08metadata\x88\x01\x01\x42\x0b\n\t_metadata\x1a\x9e\x01\n\x0eLambdaFunction\x12\x35\n\x08\x66unction\x18\x01 \x01(\x0b\x32\x19.spark.connect.ExpressionR\x08\x66unction\x12U\n\targuments\x18\x02 \x03(\x0b\x32\x37.spark.connect.Expression.UnresolvedNamedLambdaVariableR\targuments\x1a>\n\x1dUnresolvedNamedLambdaVariable\x12\x1d\n\nname_parts\x18\x01 \x03(\tR\tnamePartsB\x0b\n\texpr_type"A\n\x10\x45xpressionCommon\x12-\n\x06origin\x18\x01 \x01(\x0b\x32\x15.spark.connect.OriginR\x06origin"\x8d\x03\n\x1f\x43ommonInlineUserDefinedFunction\x12#\n\rfunction_name\x18\x01 \x01(\tR\x0c\x66unctionName\x12$\n\rdeterministic\x18\x02 \x01(\x08R\rdeterministic\x12\x37\n\targuments\x18\x03 \x03(\x0b\x32\x19.spark.connect.ExpressionR\targuments\x12\x39\n\npython_udf\x18\x04 \x01(\x0b\x32\x18.spark.connect.PythonUDFH\x00R\tpythonUdf\x12I\n\x10scalar_scala_udf\x18\x05 \x01(\x0b\x32\x1d.spark.connect.ScalarScalaUDFH\x00R\x0escalarScalaUdf\x12\x33\n\x08java_udf\x18\x06 \x01(\x0b\x32\x16.spark.connect.JavaUDFH\x00R\x07javaUdf\x12\x1f\n\x0bis_distinct\x18\x07 \x01(\x08R\nisDistinctB\n\n\x08\x66unction"\x9b\x02\n\tPythonUDF\x12\x38\n\x0boutput_type\x18\x01 \x01(\x0b\x32\x17.spark.connect.DataTypeR\noutputType\x12\x1b\n\teval_type\x18\x02 \x01(\x05R\x08\x65valType\x12\x18\n\x07\x63ommand\x18\x03 \x01(\x0cR\x07\x63ommand\x12\x1d\n\npython_ver\x18\x04 \x01(\tR\tpythonVer\x12/\n\x13\x61\x64\x64itional_includes\x18\x05 \x03(\tR\x12\x61\x64\x64itionalIncludes\x12=\n\x0b\x62uffer_type\x18\x06 \x01(\x0b\x32\x17.spark.connect.DataTypeH\x00R\nbufferType\x88\x01\x01\x42\x0e\n\x0c_buffer_type"\xd6\x01\n\x0eScalarScalaUDF\x12\x18\n\x07payload\x18\x01 \x01(\x0cR\x07payload\x12\x37\n\ninputTypes\x18\x02 \x03(\x0b\x32\x17.spark.connect.DataTypeR\ninputTypes\x12\x37\n\noutputType\x18\x03 \x01(\x0b\x32\x17.spark.connect.DataTypeR\noutputType\x12\x1a\n\x08nullable\x18\x04 \x01(\x08R\x08nullable\x12\x1c\n\taggregate\x18\x05 \x01(\x08R\taggregate"\x95\x01\n\x07JavaUDF\x12\x1d\n\nclass_name\x18\x01 \x01(\tR\tclassName\x12=\n\x0boutput_type\x18\x02 \x01(\x0b\x32\x17.spark.connect.DataTypeH\x00R\noutputType\x88\x01\x01\x12\x1c\n\taggregate\x18\x03 \x01(\x08R\taggregateB\x0e\n\x0c_output_type"c\n\x18TypedAggregateExpression\x12G\n\x10scalar_scala_udf\x18\x01 \x01(\x0b\x32\x1d.spark.connect.ScalarScalaUDFR\x0escalarScalaUdf"l\n\x0c\x43\x61llFunction\x12#\n\rfunction_name\x18\x01 \x01(\tR\x0c\x66unctionName\x12\x37\n\targuments\x18\x02 \x03(\x0b\x32\x19.spark.connect.ExpressionR\targuments"\\\n\x17NamedArgumentExpression\x12\x10\n\x03key\x18\x01 \x01(\tR\x03key\x12/\n\x05value\x18\x02 \x01(\x0b\x32\x19.spark.connect.ExpressionR\x05value"\x80\x04\n\x0bMergeAction\x12\x46\n\x0b\x61\x63tion_type\x18\x01 \x01(\x0e\x32%.spark.connect.MergeAction.ActionTypeR\nactionType\x12<\n\tcondition\x18\x02 \x01(\x0b\x32\x19.spark.connect.ExpressionH\x00R\tcondition\x88\x01\x01\x12G\n\x0b\x61ssignments\x18\x03 \x03(\x0b\x32%.spark.connect.MergeAction.AssignmentR\x0b\x61ssignments\x1aj\n\nAssignment\x12+\n\x03key\x18\x01 \x01(\x0b\x32\x19.spark.connect.ExpressionR\x03key\x12/\n\x05value\x18\x02 \x01(\x0b\x32\x19.spark.connect.ExpressionR\x05value"\xa7\x01\n\nActionType\x12\x17\n\x13\x41\x43TION_TYPE_INVALID\x10\x00\x12\x16\n\x12\x41\x43TION_TYPE_DELETE\x10\x01\x12\x16\n\x12\x41\x43TION_TYPE_INSERT\x10\x02\x12\x1b\n\x17\x41\x43TION_TYPE_INSERT_STAR\x10\x03\x12\x16\n\x12\x41\x43TION_TYPE_UPDATE\x10\x04\x12\x1b\n\x17\x41\x43TION_TYPE_UPDATE_STAR\x10\x05\x42\x0c\n\n_condition"\xc5\x05\n\x12SubqueryExpression\x12\x17\n\x07plan_id\x18\x01 \x01(\x03R\x06planId\x12S\n\rsubquery_type\x18\x02 \x01(\x0e\x32..spark.connect.SubqueryExpression.SubqueryTypeR\x0csubqueryType\x12\x62\n\x11table_arg_options\x18\x03 \x01(\x0b\x32\x31.spark.connect.SubqueryExpression.TableArgOptionsH\x00R\x0ftableArgOptions\x88\x01\x01\x12G\n\x12in_subquery_values\x18\x04 \x03(\x0b\x32\x19.spark.connect.ExpressionR\x10inSubqueryValues\x1a\xea\x01\n\x0fTableArgOptions\x12@\n\x0epartition_spec\x18\x01 \x03(\x0b\x32\x19.spark.connect.ExpressionR\rpartitionSpec\x12\x42\n\norder_spec\x18\x02 \x03(\x0b\x32#.spark.connect.Expression.SortOrderR\torderSpec\x12\x37\n\x15with_single_partition\x18\x03 \x01(\x08H\x00R\x13withSinglePartition\x88\x01\x01\x42\x18\n\x16_with_single_partition"\x90\x01\n\x0cSubqueryType\x12\x19\n\x15SUBQUERY_TYPE_UNKNOWN\x10\x00\x12\x18\n\x14SUBQUERY_TYPE_SCALAR\x10\x01\x12\x18\n\x14SUBQUERY_TYPE_EXISTS\x10\x02\x12\x1b\n\x17SUBQUERY_TYPE_TABLE_ARG\x10\x03\x12\x14\n\x10SUBQUERY_TYPE_IN\x10\x04\x42\x14\n\x12_table_arg_optionsB6\n\x1eorg.apache.spark.connect.protoP\x01Z\x12internal/generatedb\x06proto3' ) _globals = globals() @@ -135,27 +135,27 @@ _globals["_COMMONINLINEUSERDEFINEDFUNCTION"]._serialized_start = 7899 _globals["_COMMONINLINEUSERDEFINEDFUNCTION"]._serialized_end = 8296 _globals["_PYTHONUDF"]._serialized_start = 8299 - _globals["_PYTHONUDF"]._serialized_end = 8503 - _globals["_SCALARSCALAUDF"]._serialized_start = 8506 - _globals["_SCALARSCALAUDF"]._serialized_end = 8720 - _globals["_JAVAUDF"]._serialized_start = 8723 - _globals["_JAVAUDF"]._serialized_end = 8872 - _globals["_TYPEDAGGREGATEEXPRESSION"]._serialized_start = 8874 - _globals["_TYPEDAGGREGATEEXPRESSION"]._serialized_end = 8973 - _globals["_CALLFUNCTION"]._serialized_start = 8975 - _globals["_CALLFUNCTION"]._serialized_end = 9083 - _globals["_NAMEDARGUMENTEXPRESSION"]._serialized_start = 9085 - _globals["_NAMEDARGUMENTEXPRESSION"]._serialized_end = 9177 - _globals["_MERGEACTION"]._serialized_start = 9180 - _globals["_MERGEACTION"]._serialized_end = 9692 - _globals["_MERGEACTION_ASSIGNMENT"]._serialized_start = 9402 - _globals["_MERGEACTION_ASSIGNMENT"]._serialized_end = 9508 - _globals["_MERGEACTION_ACTIONTYPE"]._serialized_start = 9511 - _globals["_MERGEACTION_ACTIONTYPE"]._serialized_end = 9678 - _globals["_SUBQUERYEXPRESSION"]._serialized_start = 9695 - _globals["_SUBQUERYEXPRESSION"]._serialized_end = 10404 - _globals["_SUBQUERYEXPRESSION_TABLEARGOPTIONS"]._serialized_start = 10001 - _globals["_SUBQUERYEXPRESSION_TABLEARGOPTIONS"]._serialized_end = 10235 - _globals["_SUBQUERYEXPRESSION_SUBQUERYTYPE"]._serialized_start = 10238 - _globals["_SUBQUERYEXPRESSION_SUBQUERYTYPE"]._serialized_end = 10382 + _globals["_PYTHONUDF"]._serialized_end = 8582 + _globals["_SCALARSCALAUDF"]._serialized_start = 8585 + _globals["_SCALARSCALAUDF"]._serialized_end = 8799 + _globals["_JAVAUDF"]._serialized_start = 8802 + _globals["_JAVAUDF"]._serialized_end = 8951 + _globals["_TYPEDAGGREGATEEXPRESSION"]._serialized_start = 8953 + _globals["_TYPEDAGGREGATEEXPRESSION"]._serialized_end = 9052 + _globals["_CALLFUNCTION"]._serialized_start = 9054 + _globals["_CALLFUNCTION"]._serialized_end = 9162 + _globals["_NAMEDARGUMENTEXPRESSION"]._serialized_start = 9164 + _globals["_NAMEDARGUMENTEXPRESSION"]._serialized_end = 9256 + _globals["_MERGEACTION"]._serialized_start = 9259 + _globals["_MERGEACTION"]._serialized_end = 9771 + _globals["_MERGEACTION_ASSIGNMENT"]._serialized_start = 9481 + _globals["_MERGEACTION_ASSIGNMENT"]._serialized_end = 9587 + _globals["_MERGEACTION_ACTIONTYPE"]._serialized_start = 9590 + _globals["_MERGEACTION_ACTIONTYPE"]._serialized_end = 9757 + _globals["_SUBQUERYEXPRESSION"]._serialized_start = 9774 + _globals["_SUBQUERYEXPRESSION"]._serialized_end = 10483 + _globals["_SUBQUERYEXPRESSION_TABLEARGOPTIONS"]._serialized_start = 10080 + _globals["_SUBQUERYEXPRESSION_TABLEARGOPTIONS"]._serialized_end = 10314 + _globals["_SUBQUERYEXPRESSION_SUBQUERYTYPE"]._serialized_start = 10317 + _globals["_SUBQUERYEXPRESSION_SUBQUERYTYPE"]._serialized_end = 10461 # @@protoc_insertion_point(module_scope) diff --git a/python/pyspark/sql/connect/proto/expressions_pb2.pyi b/python/pyspark/sql/connect/proto/expressions_pb2.pyi index c613ade2f43f0..22b65357c2345 100644 --- a/python/pyspark/sql/connect/proto/expressions_pb2.pyi +++ b/python/pyspark/sql/connect/proto/expressions_pb2.pyi @@ -1826,6 +1826,7 @@ class PythonUDF(google.protobuf.message.Message): COMMAND_FIELD_NUMBER: builtins.int PYTHON_VER_FIELD_NUMBER: builtins.int ADDITIONAL_INCLUDES_FIELD_NUMBER: builtins.int + BUFFER_TYPE_FIELD_NUMBER: builtins.int @property def output_type(self) -> pyspark.sql.connect.proto.types_pb2.DataType: """(Required) Output type of the Python UDF""" @@ -1840,6 +1841,11 @@ class PythonUDF(google.protobuf.message.Message): self, ) -> google.protobuf.internal.containers.RepeatedScalarFieldContainer[builtins.str]: """(Optional) Additional includes for the Python UDF.""" + @property + def buffer_type(self) -> pyspark.sql.connect.proto.types_pb2.DataType: + """(Optional) Intermediate buffer schema for an incremental Python aggregator + (see PythonAggregate). Set only for the incremental aggregator eval types. + """ def __init__( self, *, @@ -1848,15 +1854,28 @@ class PythonUDF(google.protobuf.message.Message): command: builtins.bytes = ..., python_ver: builtins.str = ..., additional_includes: collections.abc.Iterable[builtins.str] | None = ..., + buffer_type: pyspark.sql.connect.proto.types_pb2.DataType | None = ..., ) -> None: ... def HasField( - self, field_name: typing_extensions.Literal["output_type", b"output_type"] + self, + field_name: typing_extensions.Literal[ + "_buffer_type", + b"_buffer_type", + "buffer_type", + b"buffer_type", + "output_type", + b"output_type", + ], ) -> builtins.bool: ... def ClearField( self, field_name: typing_extensions.Literal[ + "_buffer_type", + b"_buffer_type", "additional_includes", b"additional_includes", + "buffer_type", + b"buffer_type", "command", b"command", "eval_type", @@ -1867,6 +1886,9 @@ class PythonUDF(google.protobuf.message.Message): b"python_ver", ], ) -> None: ... + def WhichOneof( + self, oneof_group: typing_extensions.Literal["_buffer_type", b"_buffer_type"] + ) -> typing_extensions.Literal["buffer_type"] | None: ... global___PythonUDF = PythonUDF diff --git a/python/pyspark/sql/connect/udf.py b/python/pyspark/sql/connect/udf.py index c9848634eb303..06521fb87c712 100644 --- a/python/pyspark/sql/connect/udf.py +++ b/python/pyspark/sql/connect/udf.py @@ -210,6 +210,8 @@ def to_expr(col: "ColumnOrName") -> Expression: eval_type=self.evalType, func=self.func, python_ver="%d.%d" % sys.version_info[:2], + # Set for incremental Python aggregators (see pyspark.sql.aggregator). + buffer_type=getattr(self, "bufferSchema", None), ) return CommonInlineUserDefinedFunction( function_name=self._name, @@ -303,6 +305,7 @@ def register( PythonEvalType.SQL_GROUPED_AGG_ARROW_UDF, PythonEvalType.SQL_GROUPED_AGG_PANDAS_ITER_UDF, PythonEvalType.SQL_GROUPED_AGG_ARROW_ITER_UDF, + PythonEvalType.SQL_GROUPED_AGG_ARROW_INCREMENTAL_FINAL_UDF, ]: raise PySparkTypeError( errorClass="INVALID_UDF_EVAL_TYPE", @@ -311,11 +314,18 @@ def register( "SQL_SCALAR_PANDAS_UDF, SQL_SCALAR_ARROW_UDF, " "SQL_SCALAR_PANDAS_ITER_UDF, SQL_SCALAR_ARROW_ITER_UDF, " "SQL_GROUPED_AGG_PANDAS_UDF, SQL_GROUPED_AGG_ARROW_UDF, " - "SQL_GROUPED_AGG_PANDAS_ITER_UDF or SQL_GROUPED_AGG_ARROW_ITER_UDF" + "SQL_GROUPED_AGG_PANDAS_ITER_UDF, SQL_GROUPED_AGG_ARROW_ITER_UDF " + "or SQL_GROUPED_AGG_ARROW_INCREMENTAL_FINAL_UDF" }, ) self.sparkSession._client.register_udf( - f.func, f.returnType, name, f.evalType, f.deterministic + f.func, + f.returnType, + name, + f.evalType, + f.deterministic, + # Set for the incremental aggregator (see pyspark.sql.aggregator). + buffer_type=getattr(f, "bufferSchema", None), ) return f else: diff --git a/python/pyspark/sql/pandas/_typing/__init__.pyi b/python/pyspark/sql/pandas/_typing/__init__.pyi index f989443a4dd90..442849bcbf0a5 100644 --- a/python/pyspark/sql/pandas/_typing/__init__.pyi +++ b/python/pyspark/sql/pandas/_typing/__init__.pyi @@ -68,6 +68,8 @@ ArrowScalarIterUDFType = Literal[251] ArrowGroupedAggUDFType = Literal[252] ArrowWindowAggUDFType = Literal[253] ArrowGroupedAggIterUDFType = Literal[254] +ArrowGroupedAggIncrementalPartialUDFType = Literal[255] +ArrowGroupedAggIncrementalFinalUDFType = Literal[256] # Arrow stream types # A single group of Arrow batches (e.g., one key group in groupBy). diff --git a/python/pyspark/sql/tests/arrow/test_arrow_python_aggregator.py b/python/pyspark/sql/tests/arrow/test_arrow_python_aggregator.py new file mode 100644 index 0000000000000..1bfebedbdd71c --- /dev/null +++ b/python/pyspark/sql/tests/arrow/test_arrow_python_aggregator.py @@ -0,0 +1,250 @@ +# +# Licensed to the Apache Software Foundation (ASF) under one or more +# contributor license agreements. See the NOTICE file distributed with +# this work for additional information regarding copyright ownership. +# The ASF licenses this file to You under the Apache License, Version 2.0 +# (the "License"); you may not use this file except in compliance with +# the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# +import unittest +from decimal import Decimal + +from pyspark.sql import functions as sf +from pyspark.sql.types import ( + DecimalType, + DoubleType, + LongType, + StructField, + StructType, +) +from pyspark.testing.sqlutils import ReusedSQLTestCase +from pyspark.testing.utils import ( + have_pyarrow, + pyarrow_requirement_message, +) + + +if have_pyarrow: + from pyspark.sql.aggregator import Aggregator, udaf + + class Mean(Aggregator): + @property + def bufferSchema(self): + return StructType([StructField("sum", DoubleType()), StructField("count", LongType())]) + + @property + def outputType(self): + return DoubleType() + + def zero(self): + return (0.0, 0) + + def reduce(self, buffer, value): + (v,) = value + if v is None: # ignore nulls, like SQL avg + return buffer + return (buffer[0] + v, buffer[1] + 1) + + def merge(self, b1, b2): + return (b1[0] + b2[0], b1[1] + b2[1]) + + def finish(self, buffer): + return buffer[0] / buffer[1] if buffer[1] else None + + class DecimalSum(Aggregator): + # Non-trivial output/buffer type: the result column and the intermediate buffer are both + # DecimalType, exercising explicit Arrow typing of the emitted arrays (a bare + # ``pa.array([Decimal(...)])`` would infer a decimal type whose precision/scale need not + # match the declared one). + @property + def bufferSchema(self): + return StructType([StructField("total", DecimalType(20, 4))]) + + @property + def outputType(self): + return DecimalType(20, 4) + + def zero(self): + return (Decimal(0),) + + def reduce(self, buffer, value): + (v,) = value + return buffer if v is None else (buffer[0] + Decimal(str(v)),) + + def merge(self, b1, b2): + return (b1[0] + b2[0],) + + def finish(self, buffer): + return buffer[0] + + class SumSquares(Aggregator): + @property + def bufferSchema(self): + return StructType([StructField("sumsq", DoubleType())]) + + @property + def outputType(self): + return DoubleType() + + def zero(self): + return (0.0,) + + def reduce(self, buffer, value): + (v,) = value + return (buffer[0] + float(v) * float(v),) + + def merge(self, b1, b2): + return (b1[0] + b2[0],) + + def finish(self, buffer): + return buffer[0] + + +@unittest.skipIf(not have_pyarrow, pyarrow_requirement_message) +class ArrowPythonAggregatorTestsMixin: + def _data(self): + # 100 rows across 5 keys; repartition so each key is split across partitions, + # exercising map-side PARTIAL combine + post-shuffle FINAL merge. + return ( + self.spark.range(0, 100) + .select((sf.col("id") % 5).alias("k"), sf.col("id").cast("double").alias("v")) + .repartition(4, sf.col("v") % 3) + ) + + def test_incremental_aggregator_matches_builtin_mean(self): + df = self._data() + result = df.groupBy("k").agg(udaf(Mean())(sf.col("v")).alias("m")).orderBy("k").collect() + expected = df.groupBy("k").agg(sf.avg("v").alias("m")).orderBy("k").collect() + got = {r["k"]: r["m"] for r in result} + exp = {r["k"]: r["m"] for r in expected} + self.assertEqual(got, exp) + + def test_incremental_aggregator_no_group(self): + df = self._data() + result = df.agg(udaf(Mean())(sf.col("v")).alias("m")).collect() + expected = df.agg(sf.avg("v").alias("m")).collect() + self.assertAlmostEqual(result[0]["m"], expected[0]["m"], places=6) + + def test_incremental_aggregator_empty_global_input(self): + # A global aggregation over empty input must still return one identity row: finish(zero). + empty = self._data().limit(0) + result = empty.agg(udaf(Mean())(sf.col("v")).alias("m")).collect() + expected = empty.agg(sf.avg("v").alias("m")).collect() + self.assertEqual(len(result), 1) + self.assertEqual(result[0]["m"], expected[0]["m"]) + + def test_incremental_aggregator_custom_buffer(self): + df = self._data() + result = ( + df.groupBy("k").agg(udaf(SumSquares())(sf.col("v")).alias("s")).orderBy("k").collect() + ) + expected = ( + df.groupBy("k").agg(sf.sum(sf.col("v") * sf.col("v")).alias("s")).orderBy("k").collect() + ) + got = {r["k"]: r["s"] for r in result} + exp = {r["k"]: r["s"] for r in expected} + for k in exp: + self.assertAlmostEqual(got[k], exp[k], places=6) + + def test_incremental_aggregator_decimal_output(self): + # Non-trivial output/buffer type (DecimalType), crossing the shuffle as a decimal buffer + # and emitted as a decimal result -- guards the explicit Arrow typing of both stages. + df = self._data() + result = ( + df.groupBy("k").agg(udaf(DecimalSum())(sf.col("v")).alias("s")).orderBy("k").collect() + ) + expected = df.groupBy("k").agg(sf.sum("v").alias("s")).orderBy("k").collect() + got = {r["k"]: r["s"] for r in result} + exp = {r["k"]: r["s"] for r in expected} + for k in exp: + self.assertIsInstance(got[k], Decimal) + self.assertAlmostEqual(float(got[k]), exp[k], places=4) + + def test_incremental_aggregator_null_inputs(self): + # reduce must tolerate null input values; the null-skipping Mean should match SQL avg, + # including a group whose values are all null (identity buffer -> finish returns None). + df = self.spark.createDataFrame( + [("a", 1.0), ("a", None), ("a", 3.0), ("b", None), ("b", None)], + "k string, v double", + ) + result = df.groupBy("k").agg(udaf(Mean())(sf.col("v")).alias("m")).orderBy("k").collect() + expected = df.groupBy("k").agg(sf.avg("v").alias("m")).orderBy("k").collect() + got = {r["k"]: r["m"] for r in result} + exp = {r["k"]: r["m"] for r in expected} + self.assertEqual(got, exp) + + def test_multiple_incremental_aggregators(self): + # Two aggregators with different buffer schemas over the same input in one agg call. + df = self._data() + result = ( + df.groupBy("k") + .agg( + udaf(Mean())(sf.col("v")).alias("m"), + udaf(SumSquares())(sf.col("v")).alias("s"), + ) + .orderBy("k") + .collect() + ) + expected = ( + df.groupBy("k") + .agg( + sf.avg("v").alias("m"), + sf.sum(sf.col("v") * sf.col("v")).alias("s"), + ) + .orderBy("k") + .collect() + ) + got_m = {r["k"]: r["m"] for r in result} + got_s = {r["k"]: r["s"] for r in result} + for r in expected: + self.assertAlmostEqual(got_m[r["k"]], r["m"], places=6) + self.assertAlmostEqual(got_s[r["k"]], r["s"], places=6) + + def test_result_independent_of_partition_count(self): + # Partial buffers must merge to the same result regardless of how keys are split. + base = self.spark.range(0, 60).select( + (sf.col("id") % 3).alias("k"), sf.col("id").cast("double").alias("v") + ) + results = [] + for n in (1, 2, 7): + rows = ( + base.repartition(n, sf.col("v")) + .groupBy("k") + .agg(udaf(Mean())(sf.col("v")).alias("m")) + .orderBy("k") + .collect() + ) + results.append({r["k"]: r["m"] for r in rows}) + self.assertEqual(results[0], results[1]) + self.assertEqual(results[1], results[2]) + + def test_sql_registration(self): + # Register the aggregator and invoke it from SQL text. + df = self._data() + df.createOrReplaceTempView("agg_input") + self.spark.udf.register("my_mean", udaf(Mean())) + result = self.spark.sql( + "SELECT k, my_mean(v) AS m FROM agg_input GROUP BY k ORDER BY k" + ).collect() + expected = df.groupBy("k").agg(sf.avg("v").alias("m")).orderBy("k").collect() + got = {r["k"]: r["m"] for r in result} + exp = {r["k"]: r["m"] for r in expected} + self.assertEqual(got, exp) + + +class ArrowPythonAggregatorTests(ArrowPythonAggregatorTestsMixin, ReusedSQLTestCase): + pass + + +if __name__ == "__main__": + from pyspark.testing import main + + main() diff --git a/python/pyspark/sql/tests/connect/arrow/test_parity_arrow_python_aggregator.py b/python/pyspark/sql/tests/connect/arrow/test_parity_arrow_python_aggregator.py new file mode 100644 index 0000000000000..5b49582e8391a --- /dev/null +++ b/python/pyspark/sql/tests/connect/arrow/test_parity_arrow_python_aggregator.py @@ -0,0 +1,29 @@ +# +# Licensed to the Apache Software Foundation (ASF) under one or more +# contributor license agreements. See the NOTICE file distributed with +# this work for additional information regarding copyright ownership. +# The ASF licenses this file to You under the Apache License, Version 2.0 +# (the "License"); you may not use this file except in compliance with +# the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# + +from pyspark.sql.tests.arrow.test_arrow_python_aggregator import ArrowPythonAggregatorTestsMixin +from pyspark.testing.connectutils import ReusedConnectTestCase + + +class ArrowPythonAggregatorParityTests(ArrowPythonAggregatorTestsMixin, ReusedConnectTestCase): + pass + + +if __name__ == "__main__": + from pyspark.testing import main + + main() diff --git a/python/pyspark/sql/tests/pandas/test_pandas_grouped_map.py b/python/pyspark/sql/tests/pandas/test_pandas_grouped_map.py index 92fb2d0cb06a8..d0031545d7b21 100644 --- a/python/pyspark/sql/tests/pandas/test_pandas_grouped_map.py +++ b/python/pyspark/sql/tests/pandas/test_pandas_grouped_map.py @@ -243,7 +243,8 @@ def test_register_grouped_map_udf(self): "SQL_SCALAR_PANDAS_UDF, SQL_SCALAR_ARROW_UDF, " "SQL_SCALAR_PANDAS_ITER_UDF, SQL_SCALAR_ARROW_ITER_UDF, " "SQL_GROUPED_AGG_PANDAS_UDF, SQL_GROUPED_AGG_ARROW_UDF, " - "SQL_GROUPED_AGG_PANDAS_ITER_UDF or SQL_GROUPED_AGG_ARROW_ITER_UDF" + "SQL_GROUPED_AGG_PANDAS_ITER_UDF, SQL_GROUPED_AGG_ARROW_ITER_UDF " + "or SQL_GROUPED_AGG_ARROW_INCREMENTAL_FINAL_UDF" }, ) diff --git a/python/pyspark/sql/udf.py b/python/pyspark/sql/udf.py index dbfc1b4650864..f1f6fcd70b753 100644 --- a/python/pyspark/sql/udf.py +++ b/python/pyspark/sql/udf.py @@ -517,15 +517,35 @@ def _create_judf( assert sc._jvm is not None transpiled = self.transpiled if include_transpiled else [] input_categories = self._transpiled_input_categories if include_transpiled else [] - judf = getattr(sc._jvm, "org.apache.spark.sql.execution.python.UserDefinedPythonFunction")( - self._name, - wrapped_func, - jdt, - self.evalType, - self.deterministic, - map(_to_java_column_opt, transpiled), - input_categories, - ) + # Incremental Python aggregators additionally carry the intermediate buffer schema, which + # the JVM needs at planning time to build the two-stage aggregation (see PythonAggregate). + buffer_schema = getattr(self, "bufferSchema", None) + if buffer_schema is not None: + jbuf = spark._jsparkSession.parseDataType(buffer_schema.json()) + judf = getattr( + sc._jvm, "org.apache.spark.sql.execution.python.UserDefinedPythonFunction" + )( + self._name, + wrapped_func, + jdt, + self.evalType, + self.deterministic, + map(_to_java_column_opt, transpiled), + input_categories, + jbuf, + ) + else: + judf = getattr( + sc._jvm, "org.apache.spark.sql.execution.python.UserDefinedPythonFunction" + )( + self._name, + wrapped_func, + jdt, + self.evalType, + self.deterministic, + map(_to_java_column_opt, transpiled), + input_categories, + ) return judf def __call__(self, *args: "ColumnOrName", **kwargs: "ColumnOrName") -> Column: @@ -839,6 +859,7 @@ def register( PythonEvalType.SQL_GROUPED_AGG_ARROW_UDF, PythonEvalType.SQL_GROUPED_AGG_PANDAS_ITER_UDF, PythonEvalType.SQL_GROUPED_AGG_ARROW_ITER_UDF, + PythonEvalType.SQL_GROUPED_AGG_ARROW_INCREMENTAL_FINAL_UDF, ]: raise PySparkTypeError( errorClass="INVALID_UDF_EVAL_TYPE", @@ -847,7 +868,8 @@ def register( "SQL_SCALAR_PANDAS_UDF, SQL_SCALAR_ARROW_UDF, " "SQL_SCALAR_PANDAS_ITER_UDF, SQL_SCALAR_ARROW_ITER_UDF, " "SQL_GROUPED_AGG_PANDAS_UDF, SQL_GROUPED_AGG_ARROW_UDF, " - "SQL_GROUPED_AGG_PANDAS_ITER_UDF or SQL_GROUPED_AGG_ARROW_ITER_UDF" + "SQL_GROUPED_AGG_PANDAS_ITER_UDF, SQL_GROUPED_AGG_ARROW_ITER_UDF " + "or SQL_GROUPED_AGG_ARROW_INCREMENTAL_FINAL_UDF" }, ) source_udf = _create_udf( @@ -858,6 +880,10 @@ def register( deterministic=f.deterministic, ) register_udf = source_udf._unwrapped # type: ignore[attr-defined] + # Preserve the incremental aggregator's buffer schema, which _create_udf drops. + buffer_schema = getattr(f, "bufferSchema", None) + if buffer_schema is not None: + register_udf.bufferSchema = buffer_schema return_udf = register_udf else: if returnType is None: diff --git a/python/pyspark/util.py b/python/pyspark/util.py index 76be9bcc9998e..81bea6152f98c 100644 --- a/python/pyspark/util.py +++ b/python/pyspark/util.py @@ -86,6 +86,8 @@ ArrowScalarIterUDFType, ArrowGroupedAggUDFType, ArrowGroupedAggIterUDFType, + ArrowGroupedAggIncrementalPartialUDFType, + ArrowGroupedAggIncrementalFinalUDFType, ArrowWindowAggUDFType, ) from pyspark.sql._typing import ( @@ -699,6 +701,13 @@ class PythonEvalType: SQL_WINDOW_AGG_ARROW_UDF: "ArrowWindowAggUDFType" = 253 SQL_GROUPED_AGG_ARROW_ITER_UDF: "ArrowGroupedAggIterUDFType" = 254 + # Incremental (partial + final) Arrow aggregator. See ``pyspark.sql.aggregator``. + # PARTIAL folds input rows into a per-group buffer via ``Aggregator.reduce`` on the map side; + # FINAL merges partial buffers via ``Aggregator.merge`` and produces output via + # ``Aggregator.finish`` after the shuffle. + SQL_GROUPED_AGG_ARROW_INCREMENTAL_PARTIAL_UDF: "ArrowGroupedAggIncrementalPartialUDFType" = 255 + SQL_GROUPED_AGG_ARROW_INCREMENTAL_FINAL_UDF: "ArrowGroupedAggIncrementalFinalUDFType" = 256 + SQL_TABLE_UDF: "SQLTableUDFType" = 300 SQL_ARROW_TABLE_UDF: "SQLArrowTableUDFType" = 301 SQL_ARROW_UDTF: "SQLArrowUDTFType" = 302 diff --git a/python/pyspark/worker.py b/python/pyspark/worker.py index 4e77c630777d1..7b3861dfec6ff 100644 --- a/python/pyspark/worker.py +++ b/python/pyspark/worker.py @@ -626,6 +626,15 @@ def read_single_udf(pickleSer, udf_info, eval_type, runner_conf, udf_index): # The last returnType will be the return type of UDF. Eval types are grouped below by the # shape of the value they return. + # Incremental Python aggregators: the pickled "function" is the Aggregator object itself, whose + # zero/reduce/merge/finish methods the worker calls directly. Return it unwrapped (not through + # fail_on_stopiteration, which would treat it as a plain callable). + if eval_type in ( + PythonEvalType.SQL_GROUPED_AGG_ARROW_INCREMENTAL_PARTIAL_UDF, + PythonEvalType.SQL_GROUPED_AGG_ARROW_INCREMENTAL_FINAL_UDF, + ): + return chained_func, args_offsets, kwargs_offsets, return_type + # Scalar, aggregation and window UDFs: (func, args_offsets, kwargs_offsets, return_type). if eval_type in ( PythonEvalType.SQL_ARROW_BATCHED_UDF, @@ -1946,6 +1955,8 @@ def read_udfs(pickleSer, udf_info_list, eval_type, runner_conf, eval_conf): PythonEvalType.SQL_TRANSFORM_WITH_STATE_PANDAS_INIT_STATE_UDF, PythonEvalType.SQL_TRANSFORM_WITH_STATE_PYTHON_ROW_UDF, PythonEvalType.SQL_TRANSFORM_WITH_STATE_PYTHON_ROW_INIT_STATE_UDF, + PythonEvalType.SQL_GROUPED_AGG_ARROW_INCREMENTAL_PARTIAL_UDF, + PythonEvalType.SQL_GROUPED_AGG_ARROW_INCREMENTAL_FINAL_UDF, ): # NOTE: if timezone is set here, that implies respectSessionTimeZone is True if eval_type in ( @@ -1953,6 +1964,8 @@ def read_udfs(pickleSer, udf_info_list, eval_type, runner_conf, eval_conf): PythonEvalType.SQL_GROUPED_AGG_ARROW_UDF, PythonEvalType.SQL_GROUPED_AGG_PANDAS_ITER_UDF, PythonEvalType.SQL_GROUPED_AGG_PANDAS_UDF, + PythonEvalType.SQL_GROUPED_AGG_ARROW_INCREMENTAL_PARTIAL_UDF, + PythonEvalType.SQL_GROUPED_AGG_ARROW_INCREMENTAL_FINAL_UDF, PythonEvalType.SQL_GROUPED_MAP_ARROW_ITER_UDF, PythonEvalType.SQL_GROUPED_MAP_ARROW_UDF, PythonEvalType.SQL_GROUPED_MAP_PANDAS_ITER_UDF, @@ -2199,6 +2212,95 @@ def grouped_func( # profiling is not supported for UDF return grouped_func, None, ser, ser + if eval_type == PythonEvalType.SQL_GROUPED_AGG_ARROW_INCREMENTAL_PARTIAL_UDF: + import pyarrow as pa + + # Map-side PARTIAL stage: fold each group's input rows into a per-group buffer via the + # aggregator's `reduce`, and emit one intermediate-buffer struct column per aggregator. + # Batches are streamed and folded one at a time -- only the per-aggregator buffers are + # retained, never the whole group -- so a skewed group keeps map-side memory bounded. + # `udf_func` is the Aggregator object itself (see read_single_udf). + return_schema = to_arrow_schema( + StructType( + [StructField("_%d" % i, agg.bufferSchema) for i, (agg, _, _, _) in enumerate(udfs)] + ), + timezone="UTC", + prefers_large_types=runner_conf.use_large_var_types, + ) + col_names = ["_%d" % i for i in range(len(udfs))] + # Buffer field names are invariant across groups; compute them once per aggregator. + field_names_by_udf = [[f.name for f in agg.bufferSchema.fields] for agg, _, _, _ in udfs] + + def grouped_func( + split_index: int, data: Iterator["GroupedBatch"] + ) -> Iterator[pa.RecordBatch]: + for group in data: + buffers = [agg.zero() for agg, _, _, _ in udfs] + for batch in group: + for i, (agg, args_offsets, _, _) in enumerate(udfs): + cols = [batch.column(o).to_pylist() for o in args_offsets] + buf = buffers[i] + for r in range(batch.num_rows): + buf = agg.reduce(buf, tuple(c[r] for c in cols)) + buffers[i] = buf + result_arrays = [] + for i, (agg, _, _, _) in enumerate(udfs): + field_names = field_names_by_udf[i] + struct_value = {name: buffers[i][j] for j, name in enumerate(field_names)} + result_arrays.append(pa.array([struct_value], type=return_schema.field(i).type)) + batch = pa.RecordBatch.from_arrays(result_arrays, col_names) + yield ArrowBatchTransformer.enforce_schema(batch, return_schema) + + # profiling is not supported for UDF + return grouped_func, None, ser, ser + + if eval_type == PythonEvalType.SQL_GROUPED_AGG_ARROW_INCREMENTAL_FINAL_UDF: + import pyarrow as pa + + # Post-shuffle FINAL stage: merge each group's partial buffers via the aggregator's `merge` + # and produce the output via `finish`. Buffers are streamed and merged one batch at a time. + # Every group the JVM sends yields exactly one output row; null partial-buffer rows are + # skipped, so an empty global aggregation (a single all-null buffer row injected by the + # operator) still produces `finish(zero)`. + col_names = ["_%d" % i for i in range(len(udfs))] + return_schema = to_arrow_schema( + StructType([StructField(name, rt) for name, (_, _, _, rt) in zip(col_names, udfs)]), + timezone="UTC", + prefers_large_types=runner_conf.use_large_var_types, + ) + # Buffer field names are invariant across groups and batches; compute once per aggregator. + field_names_by_udf = [[f.name for f in agg.bufferSchema.fields] for agg, _, _, _ in udfs] + + def grouped_func( + split_index: int, data: Iterator["GroupedBatch"] + ) -> Iterator[pa.RecordBatch]: + for group in data: + merged: list = [None] * len(udfs) + for batch in group: + for i, (agg, args_offsets, _, _) in enumerate(udfs): + field_names = field_names_by_udf[i] + m = merged[i] + for row in batch.column(args_offsets[0]).to_pylist(): + if row is None: + continue + partial = tuple(row[name] for name in field_names) + m = partial if m is None else agg.merge(m, partial) + merged[i] = m + results = [] + for i, (agg, _, _, _) in enumerate(udfs): + m = merged[i] if merged[i] is not None else agg.zero() + results.append(agg.finish(m)) + # Type each output array explicitly (mirroring the PARTIAL stage) so a non-trivial + # outputType or an all-None column does not depend on Arrow type inference. + result_arrays = [ + pa.array([r], type=return_schema.field(i).type) for i, r in enumerate(results) + ] + batch = pa.RecordBatch.from_arrays(result_arrays, col_names) + yield ArrowBatchTransformer.enforce_schema(batch, return_schema) + + # profiling is not supported for UDF + return grouped_func, None, ser, ser + if eval_type == PythonEvalType.SQL_GROUPED_AGG_PANDAS_UDF: import pyarrow as pa import pandas as pd diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/PythonUDF.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/PythonUDF.scala index e421d946a669c..ddd4979b07ef1 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/PythonUDF.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/PythonUDF.scala @@ -348,6 +348,55 @@ case class PythonUDAF( copy(children = newChildren) } +/** + * A serialized Python aggregator that supports true incremental (partial) aggregation, the + * analog of the Scala typed `org.apache.spark.sql.expressions.Aggregator[IN, BUF, OUT]`. Unlike + * [[PythonUDAF]] (which materializes the whole group and calls Python once), this is planned as a + * two-stage aggregation by + * [[org.apache.spark.sql.execution.python.PythonIncrementalAggregateExec]]: a map-side PARTIAL + * stage folds input rows into a per-group buffer via the aggregator's `reduce`, and a post-shuffle + * FINAL stage + * merges the partial buffers via `merge` and produces the output via `finish`. + * + * `bufferSchema` is the schema of the intermediate buffer that crosses the shuffle between the two + * stages (the analog of the Scala aggregator's `bufferEncoder`). It is exposed here rather than via + * [[aggBufferAttributes]] because, like [[PythonUDAF]], this expression is unevaluable in the JVM; + * the physical operator derives the buffer attributes from `bufferSchema` directly. + */ +case class PythonAggregate( + name: String, + func: PythonFunction, + dataType: DataType, + children: Seq[Expression], + udfDeterministic: Boolean, + bufferSchema: StructType, + evalType: Int = PythonEvalType.SQL_GROUPED_AGG_ARROW_INCREMENTAL_FINAL_UDF, + resultId: ExprId = NamedExpression.newExprId) + extends UnevaluableAggregateFunc with PythonFuncExpression { + + override def sql(isDistinct: Boolean): String = { + val distinct = if (isDistinct) "DISTINCT " else "" + s"$name($distinct${children.mkString(", ")})" + } + + override def toAggString(isDistinct: Boolean): String = { + val start = if (isDistinct) "(distinct " else "(" + name + children.mkString(start, ", ", ")") + s"#${resultId.id}$typeSuffix" + } + + override lazy val canonicalized: Expression = { + val canonicalizedChildren = children.map(_.canonicalized) + // `resultId` can be seen as cosmetic variation, as it doesn't affect the result. + this.copy(resultId = ExprId(-1)).withNewChildren(canonicalizedChildren) + } + + final override val nodePatterns: Seq[TreePattern] = Seq(PYTHON_UDF) + + override protected def withNewChildrenInternal( + newChildren: IndexedSeq[Expression]): PythonAggregate = + copy(children = newChildren) +} + abstract class UnevaluableGenerator extends Generator { final override def eval(input: InternalRow): IterableOnce[InternalRow] = throw QueryExecutionErrors.cannotEvaluateExpressionError(this) diff --git a/sql/connect/common/src/main/protobuf/spark/connect/expressions.proto b/sql/connect/common/src/main/protobuf/spark/connect/expressions.proto index 18f8f294f0c02..432bc13e918dc 100644 --- a/sql/connect/common/src/main/protobuf/spark/connect/expressions.proto +++ b/sql/connect/common/src/main/protobuf/spark/connect/expressions.proto @@ -475,6 +475,9 @@ message PythonUDF { string python_ver = 4; // (Optional) Additional includes for the Python UDF. repeated string additional_includes = 5; + // (Optional) Intermediate buffer schema for an incremental Python aggregator + // (see PythonAggregate). Set only for the incremental aggregator eval types. + optional DataType buffer_type = 6; } message ScalarScalaUDF { diff --git a/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/planner/SparkConnectPlanner.scala b/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/planner/SparkConnectPlanner.scala index 841ce26402cf0..9ccdc19ceb442 100644 --- a/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/planner/SparkConnectPlanner.scala +++ b/sql/connect/server/src/main/scala/org/apache/spark/sql/connect/planner/SparkConnectPlanner.scala @@ -2201,6 +2201,7 @@ class SparkConnectPlanner( createUserDefinedPythonFunction(fun) .builder(fun.getArgumentsList.asScala.map(transformExpression).toSeq) match { case udaf: PythonUDAF => udaf.toAggregateExpression() + case agg: PythonAggregate => agg.toAggregateExpression() case other => other } } @@ -2214,7 +2215,9 @@ class SparkConnectPlanner( func = function, dataType = transformDataType(udf.getOutputType), pythonEvalType = udf.getEvalType, - udfDeterministic = fun.getDeterministic) + udfDeterministic = fun.getDeterministic, + // Set only for incremental Python aggregators (see PythonAggregate). + bufferType = if (udf.hasBufferType) transformDataType(udf.getBufferType) else null) } private def transformPythonFunction(fun: proto.PythonUDF): SimplePythonFunction = { diff --git a/sql/core/src/main/scala/org/apache/spark/sql/execution/SparkStrategies.scala b/sql/core/src/main/scala/org/apache/spark/sql/execution/SparkStrategies.scala index 577f09d9a685b..031014ec5fa5f 100644 --- a/sql/core/src/main/scala/org/apache/spark/sql/execution/SparkStrategies.scala +++ b/sql/core/src/main/scala/org/apache/spark/sql/execution/SparkStrategies.scala @@ -720,7 +720,8 @@ abstract class SparkStrategies extends QueryPlanner[SparkPlan] { object Aggregation extends Strategy { def apply(plan: LogicalPlan): Seq[SparkPlan] = plan match { case PhysicalAggregation(groupingExpressions, aggExpressions, resultExpressions, child) - if !aggExpressions.exists(_.aggregateFunction.isInstanceOf[PythonUDAF]) => + if !aggExpressions.exists(ae => ae.aggregateFunction.isInstanceOf[PythonUDAF] || + ae.aggregateFunction.isInstanceOf[PythonAggregate]) => val (functionsWithDistinct, functionsWithoutDistinct) = aggExpressions.partition(_.isDistinct) val distinctAggChildSets = functionsWithDistinct.map { ae => @@ -799,13 +800,25 @@ abstract class SparkStrategies extends QueryPlanner[SparkPlan] { resultExpressions, planLater(child))) + case PhysicalAggregation(groupingExpressions, aggExpressions, resultExpressions, child) + if aggExpressions.forall(_.aggregateFunction.isInstanceOf[PythonAggregate]) => + Seq(execution.python.PythonIncrementalAggregateExec.plan( + groupingExpressions, + aggExpressions, + resultExpressions, + planLater(child))) + case PhysicalAggregation(_, aggExpressions, _, _) => - val groupAggPandasUDFNames = aggExpressions + // Reached when Python aggregate UDFs cannot be planned by the two cases above -- e.g. a + // grouped-agg pandas/arrow UDF or an incremental Python aggregator is mixed with other + // (SQL or differently-typed Python) aggregate functions in the same Aggregate. + val pythonUDFNames = aggExpressions .map(_.aggregateFunction) - .filter(_.isInstanceOf[PythonUDAF]) - .map(_.asInstanceOf[PythonUDAF].name) - // If cannot match the two cases above, then it's an error - throw QueryCompilationErrors.invalidPandasUDFPlacementError(groupAggPandasUDFNames.distinct) + .collect { + case p: PythonUDAF => p.name + case p: PythonAggregate => p.name + } + throw QueryCompilationErrors.invalidPandasUDFPlacementError(pythonUDFNames.distinct) case _ => Nil } diff --git a/sql/core/src/main/scala/org/apache/spark/sql/execution/python/PythonIncrementalAggregateExec.scala b/sql/core/src/main/scala/org/apache/spark/sql/execution/python/PythonIncrementalAggregateExec.scala new file mode 100644 index 0000000000000..a937296e703e3 --- /dev/null +++ b/sql/core/src/main/scala/org/apache/spark/sql/execution/python/PythonIncrementalAggregateExec.scala @@ -0,0 +1,295 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one or more + * contributor license agreements. See the NOTICE file distributed with + * this work for additional information regarding copyright ownership. + * The ASF licenses this file to You under the Apache License, Version 2.0 + * (the "License"); you may not use this file except in compliance with + * the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package org.apache.spark.sql.execution.python + +import java.io.File + +import scala.collection.mutable.ArrayBuffer + +import org.apache.spark.{JobArtifactSet, SparkEnv, TaskContext} +import org.apache.spark.api.python.{ChainedPythonFunctions, PythonEvalType} +import org.apache.spark.rdd.RDD +import org.apache.spark.sql.catalyst.InternalRow +import org.apache.spark.sql.catalyst.expressions._ +import org.apache.spark.sql.catalyst.expressions.aggregate.AggregateExpression +import org.apache.spark.sql.catalyst.plans.physical.{AllTuples, ClusteredDistribution, Distribution, Partitioning, UnspecifiedDistribution} +import org.apache.spark.sql.execution.{GroupedIterator, SparkPlan, UnaryExecNode} +import org.apache.spark.sql.execution.python.EvalPythonExec.ArgumentMetadata +import org.apache.spark.sql.types.{DataType, StructField, StructType} +import org.apache.spark.util.Utils + +/** + * Shared execution logic for the two stages of an incremental Python aggregation (see + * [[org.apache.spark.sql.catalyst.expressions.PythonAggregate]]). Both stages group the child rows + * by the grouping expressions, send each group's projected input columns to the Python worker as + * Arrow record batches, and join the single result row the worker returns per group back with the + * grouping key. + * + * The two stages differ only in: + * - which columns are sent to Python (`udfInputs`): the aggregator arguments in the PARTIAL + * stage, the intermediate buffer columns in the FINAL stage; + * - the Python eval type, which selects `reduce`-into-buffer vs. `merge`+`finish` in the worker; + * - the required child distribution (map-side/local for PARTIAL, clustered for FINAL); + * - the output attributes and the final projection. + */ +abstract class PythonIncrementalAggregateExecBase extends UnaryExecNode with PythonSQLMetrics { + + def groupingExpressions: Seq[NamedExpression] + def aggExpressions: Seq[AggregateExpression] + + protected val udfExpressions: Seq[PythonAggregate] = + aggExpressions.map(_.aggregateFunction.asInstanceOf[PythonAggregate]) + + /** The Python eval type for this stage. */ + protected def evalType: Int + + /** Per-UDF input expressions to project out of the child and send to the Python worker. */ + protected def udfInputs: Seq[Seq[Expression]] + + /** Attributes of the row the Python worker returns per group (right side of the join). */ + protected def pythonOutputAttributes: Seq[Attribute] + + /** Expressions producing this operator's output from (groupingKey ++ pythonOutput). */ + protected def outputExpressions: Seq[NamedExpression] + + /** The grouping attributes as seen in the child's output. */ + protected def groupingAttributes: Seq[Attribute] = groupingExpressions.map(_.toAttribute) + + /** + * Whether to still invoke Python on an empty partition. Only the FINAL stage of a *global* + * (no grouping) aggregation sets this: it must emit the identity row `finish(zero)` for empty + * input, matching SQL aggregate semantics. Everywhere else an empty partition yields no rows. + */ + protected def emitOnEmptyPartition: Boolean = false + + override def output: Seq[Attribute] = outputExpressions.map(_.toAttribute) + + override def producedAttributes: AttributeSet = AttributeSet(output) + + override def requiredChildOrdering: Seq[Seq[SortOrder]] = + Seq(groupingExpressions.map(SortOrder(_, Ascending))) + + override protected def doExecute(): RDD[InternalRow] = { + val inputRDD = child.execute() + + val sessionLocalTimeZone = conf.sessionLocalTimeZone + val largeVarTypes = conf.arrowUseLargeVarTypes + val pythonRunnerConf = ArrowPythonRunner.getPythonRunnerConfMap(conf) + + val pyFuncs = udfExpressions.map { u => + (ChainedPythonFunctions(Seq(u.func)), u.resultId.id) + } + + // Filter child output attributes down to only those that are UDF inputs, and eliminate + // duplicates, mirroring ArrowAggregatePythonExec. + val allInputs = new ArrayBuffer[Expression] + val dataTypes = new ArrayBuffer[DataType] + val argMetas = udfInputs.map { input => + input.map { e => + val (key, value) = e match { + case NamedArgumentExpression(key, value) => (Some(key), value) + case _ => (None, e) + } + if (allInputs.exists(_.semanticEquals(value))) { + ArgumentMetadata(allInputs.indexWhere(_.semanticEquals(value)), key) + } else { + allInputs += value + dataTypes += value.dataType + ArgumentMetadata(allInputs.length - 1, key) + } + }.toArray + }.toArray + + val aggInputSchema = StructType(dataTypes.zipWithIndex.map { case (dt, i) => + StructField(s"_$i", dt) + }.toArray) + + val jobArtifactUUID = JobArtifactSet.getCurrentJobArtifactState.map(_.uuid) + val sessionUUID = Option(session).collect { + case s if s.sessionState.conf.pythonWorkerLoggingEnabled => s.sessionUUID + } + + val groupingExprs = groupingExpressions + val childOutput = child.output + val joinedAttributes = groupingAttributes ++ pythonOutputAttributes + val resultExprs = outputExpressions + val localEvalType = evalType + + val emitIdentityOnEmpty = emitOnEmptyPartition + inputRDD.mapPartitionsInternal { iter => if (iter.isEmpty && !emitIdentityOnEmpty) iter else { + val prunedProj = UnsafeProjection.create(allInputs.toSeq, childOutput) + + val groupedItr = if (groupingExprs.isEmpty) { + Iterator((new UnsafeRow(), iter)) + } else { + GroupedIterator(iter, groupingExprs, childOutput) + } + + // For a global aggregation with empty input, feed one all-null buffer row so the Python + // worker still emits `finish(zero)`. An empty group cannot be sent through + // GroupedPythonArrowInput (it asserts a non-empty batch per group), and the worker treats a + // null partial buffer as contributing nothing to `merge`. + lazy val nullInputRow: UnsafeRow = + UnsafeProjection.create(aggInputSchema.map(_.dataType).toArray) + .apply(new GenericInternalRow(aggInputSchema.length)).copy() + val grouped = groupedItr.map { case (key, rows) => + val projected = rows.map(prunedProj) + val toSend = if (emitIdentityOnEmpty && groupingExprs.isEmpty && !projected.hasNext) { + Iterator(nullInputRow) + } else { + projected + } + (key, toSend) + } + + val context = TaskContext.get() + + val queue = HybridRowQueue(context.taskMemoryManager(), + new File(Utils.getLocalDir(SparkEnv.get.conf)), groupingExprs.length, lockFree = false) + context.addTaskCompletionListener[Unit] { _ => queue.close() } + + val projectedRowIter = grouped.map { case (groupingKey, rows) => + queue.add(groupingKey.asInstanceOf[UnsafeRow]) + rows + } + + val runner = new ArrowPythonWithNamedArgumentRunner( + pyFuncs, + localEvalType, + argMetas, + aggInputSchema, + sessionLocalTimeZone, + largeVarTypes, + pythonRunnerConf, + pythonMetrics, + jobArtifactUUID, + sessionUUID) with GroupedPythonArrowInput + + val columnarBatchIter = runner.compute(projectedRowIter, context.partitionId(), context) + + val joined = new JoinedRow + val resultProj = UnsafeProjection.create(resultExprs, joinedAttributes) + + columnarBatchIter.map(_.rowIterator.next()).map { pythonOutputRow => + val leftRow = queue.remove() + resultProj(joined(leftRow, pythonOutputRow)) + } + }} + } +} + +/** + * Map-side PARTIAL stage: folds each group's input rows into a per-group intermediate buffer via + * the aggregator's `reduce`. It requires only a local sort on the grouping expressions (no + * shuffle), so keys may be split across partitions; [[PythonIncrementalAggregateFinalExec]] merges + * resulting partial buffers after the shuffle. Its output is the grouping key columns followed by + * one intermediate-buffer struct column per aggregator. + */ +case class PythonIncrementalAggregatePartialExec( + groupingExpressions: Seq[NamedExpression], + aggExpressions: Seq[AggregateExpression], + bufferAttributes: Seq[Attribute], + child: SparkPlan) extends PythonIncrementalAggregateExecBase { + + override protected def evalType: Int = + PythonEvalType.SQL_GROUPED_AGG_ARROW_INCREMENTAL_PARTIAL_UDF + + override protected def udfInputs: Seq[Seq[Expression]] = udfExpressions.map(_.children) + + override protected def pythonOutputAttributes: Seq[Attribute] = bufferAttributes + + override protected def outputExpressions: Seq[NamedExpression] = + groupingAttributes ++ bufferAttributes + + override def requiredChildDistribution: Seq[Distribution] = + Seq(UnspecifiedDistribution) + + override def outputPartitioning: Partitioning = child.outputPartitioning + + override protected def withNewChildInternal(newChild: SparkPlan): SparkPlan = + copy(child = newChild) +} + +/** + * Post-shuffle FINAL stage: clusters the partial buffers by the grouping key, merges the buffers + * of each group via the aggregator's `merge`, and produces the output via `finish`. Its input is + * the [[PythonIncrementalAggregatePartialExec]] output (grouping key columns followed by the + * intermediate-buffer columns); it sends the buffer columns to Python and outputs + * `resultExpressions`. + */ +case class PythonIncrementalAggregateFinalExec( + groupingExpressions: Seq[NamedExpression], + aggExpressions: Seq[AggregateExpression], + bufferAttributes: Seq[Attribute], + resultExpressions: Seq[NamedExpression], + child: SparkPlan) extends PythonIncrementalAggregateExecBase { + + override protected def evalType: Int = + PythonEvalType.SQL_GROUPED_AGG_ARROW_INCREMENTAL_FINAL_UDF + + // Each aggregator reads its own intermediate-buffer column from the (shuffled) child. + override protected def udfInputs: Seq[Seq[Expression]] = bufferAttributes.map(Seq(_)) + + override protected def pythonOutputAttributes: Seq[Attribute] = + aggExpressions.map(_.resultAttribute) + + override protected def outputExpressions: Seq[NamedExpression] = resultExpressions + + // A global (no-grouping) aggregation must return the identity row even for empty input. This + // stage runs on a single partition (AllTuples), so exactly one identity row is produced. + override protected def emitOnEmptyPartition: Boolean = groupingExpressions.isEmpty + + override def requiredChildDistribution: Seq[Distribution] = { + if (groupingExpressions.isEmpty) { + AllTuples :: Nil + } else { + ClusteredDistribution(groupingExpressions) :: Nil + } + } + + override def outputPartitioning: Partitioning = child.outputPartitioning + + override protected def withNewChildInternal(newChild: SparkPlan): SparkPlan = + copy(child = newChild) +} + +object PythonIncrementalAggregateExec { + + /** + * Builds the two-stage physical plan (PARTIAL -> [Exchange, inserted by EnsureRequirements] -> + * FINAL) for a logical aggregation whose aggregate functions are all [[PythonAggregate]]. + */ + def plan( + groupingExpressions: Seq[NamedExpression], + aggExpressions: Seq[AggregateExpression], + resultExpressions: Seq[NamedExpression], + child: SparkPlan): SparkPlan = { + // One intermediate-buffer attribute per aggregator, threaded from the PARTIAL output into the + // FINAL inputs (matched by expression id). + val bufferAttributes = aggExpressions.map { ae => + val agg = ae.aggregateFunction.asInstanceOf[PythonAggregate] + AttributeReference(s"buf_${agg.resultId.id}", agg.bufferSchema, nullable = true)() + } + val partial = PythonIncrementalAggregatePartialExec( + groupingExpressions, aggExpressions, bufferAttributes, child) + // After the PARTIAL stage the grouping expressions are materialized as plain attributes. + val groupingAttributes = groupingExpressions.map(_.toAttribute) + PythonIncrementalAggregateFinalExec( + groupingAttributes, aggExpressions, bufferAttributes, resultExpressions, partial) + } +} diff --git a/sql/core/src/main/scala/org/apache/spark/sql/execution/python/UserDefinedPythonFunction.scala b/sql/core/src/main/scala/org/apache/spark/sql/execution/python/UserDefinedPythonFunction.scala index cfd1fb4ea2229..fa9a831acf8e5 100644 --- a/sql/core/src/main/scala/org/apache/spark/sql/execution/python/UserDefinedPythonFunction.scala +++ b/sql/core/src/main/scala/org/apache/spark/sql/execution/python/UserDefinedPythonFunction.scala @@ -28,7 +28,7 @@ import net.razorvine.pickle.Pickler import org.apache.spark.api.python.{PythonEvalType, PythonFunction, PythonWorkerUtils, SpecialLengths} import org.apache.spark.sql.{Column, TableArg} import org.apache.spark.sql.catalyst.analysis.UnresolvedAttribute -import org.apache.spark.sql.catalyst.expressions.{Alias, Ascending, Descending, Expression, FunctionTableSubqueryArgumentExpression, NamedArgumentExpression, NullsFirst, NullsLast, PythonUDAF, PythonUDF, PythonUDTF, PythonUDTFAnalyzeResult, PythonUDTFSelectedExpression, SortOrder, TranspiledPythonUDF, UnresolvedPolymorphicPythonUDTF, UnresolvedTableArgPlanId} +import org.apache.spark.sql.catalyst.expressions.{Alias, Ascending, Descending, Expression, FunctionTableSubqueryArgumentExpression, NamedArgumentExpression, NullsFirst, NullsLast, PythonAggregate, PythonUDAF, PythonUDF, PythonUDTF, PythonUDTFAnalyzeResult, PythonUDTFSelectedExpression, SortOrder, TranspiledPythonUDF, UnresolvedPolymorphicPythonUDTF, UnresolvedTableArgPlanId} import org.apache.spark.sql.catalyst.parser.ParserInterface import org.apache.spark.sql.catalyst.plans.logical.{Generate, LogicalPlan, NamedParametersSupport, OneRowRelation} import org.apache.spark.sql.classic.{DataFrame, Dataset, SparkSession} @@ -56,7 +56,25 @@ case class UserDefinedPythonFunction( // categories match the bound argument types; when none match, the call // falls back to the plain Python UDF. `builder` requires the two lists to // be parallel and skips transpilation otherwise. - transpiledInputTypes: JList[JList[String]] = Nil.asJava) { + transpiledInputTypes: JList[JList[String]] = Nil.asJava, + // Schema of the intermediate aggregation buffer, set only for the incremental Python + // aggregator eval types (see [[PythonAggregate]]); `null` otherwise. Nullable rather than + // `Option` so it can be passed positionally from Python over Py4J. + bufferType: DataType = null) { + + // Preserves the constructor arity used by non-aggregator Python callers (which pass no + // `bufferType`), so their positional Py4J `new UserDefinedPythonFunction(...)` still resolves. + def this( + name: String, + func: PythonFunction, + dataType: DataType, + pythonEvalType: Int, + udfDeterministic: Boolean, + transpiled: JList[Column], + transpiledInputTypes: JList[JList[String]]) = { + this(name, func, dataType, pythonEvalType, udfDeterministic, + transpiled, transpiledInputTypes, null) + } def builder(e: Seq[Expression]): Expression = { if (pythonEvalType == PythonEvalType.SQL_BATCHED_UDF @@ -66,7 +84,8 @@ case class UserDefinedPythonFunction( || pythonEvalType == PythonEvalType.SQL_GROUPED_AGG_PANDAS_ITER_UDF || pythonEvalType == PythonEvalType.SQL_SCALAR_ARROW_UDF || pythonEvalType == PythonEvalType.SQL_GROUPED_AGG_ARROW_UDF - || pythonEvalType == PythonEvalType.SQL_GROUPED_AGG_ARROW_ITER_UDF) { + || pythonEvalType == PythonEvalType.SQL_GROUPED_AGG_ARROW_ITER_UDF + || pythonEvalType == PythonEvalType.SQL_GROUPED_AGG_ARROW_INCREMENTAL_FINAL_UDF) { /* * Check if the named arguments: * - don't have duplicated names @@ -87,6 +106,14 @@ case class UserDefinedPythonFunction( || pythonEvalType == PythonEvalType.SQL_GROUPED_AGG_ARROW_UDF || pythonEvalType == PythonEvalType.SQL_GROUPED_AGG_ARROW_ITER_UDF) { PythonUDAF(name, func, dataType, e, udfDeterministic, pythonEvalType) + } else if (pythonEvalType == PythonEvalType.SQL_GROUPED_AGG_ARROW_INCREMENTAL_FINAL_UDF) { + // The incremental Python aggregator. `bufferType` (the intermediate buffer schema) must have + // been supplied when the UDF was created. The single expression carries the aggregator for + // both the PARTIAL and FINAL stages; the physical operator picks the per-stage eval type. + require(bufferType != null, + "An incremental Python aggregator requires a buffer schema.") + PythonAggregate( + name, func, dataType, e, udfDeterministic, bufferType.asInstanceOf[StructType]) } else { PythonUDF(name, func, dataType, e, pythonEvalType, udfDeterministic) } @@ -172,6 +199,7 @@ case class UserDefinedPythonFunction( case TranspiledPythonUDF(name, udaf: PythonUDAF, transpiled, inputCategories) => TranspiledPythonUDF(name, udaf.toAggregateExpression(), transpiled, inputCategories) case udaf: PythonUDAF => udaf.toAggregateExpression() + case agg: PythonAggregate => agg.toAggregateExpression() case _ => expr }) }