From abadf1036f3a6be5895b972f8b06e67d37291d5a Mon Sep 17 00:00:00 2001 From: Hyukjin Kwon Date: Wed, 12 Aug 2026 14:22:26 +0900 Subject: [PATCH 01/11] [WIP][SQL][PYTHON] Support incremental Python aggregators via Arrow Add a Python analog of the Scala typed `Aggregator[IN, BUF, OUT]` with true incremental (partial) aggregation. Users subclass `Aggregator` (`zero`/`reduce`/`merge`/`finish` + `bufferSchema`) and wrap it with `arrow_udaf(...)` for use in `groupBy().agg(...)`. Unlike grouped-agg pandas/arrow UDFs (whole-group materialization), this is planned as a two-stage aggregation with map-side combine: a PARTIAL stage folds each group's input rows into a per-group Arrow buffer via `reduce`, the buffers are shuffled by the grouping key, and a FINAL stage merges the partial buffers via `merge` and produces the output via `finish`. - New eval types SQL_GROUPED_AGG_ARROW_INCREMENTAL_PARTIAL/FINAL_UDF (Python `PythonEvalType` and JVM `PythonEvalType`). - New Catalyst expression `PythonAggregate` carrying the intermediate buffer schema (unevaluable in the JVM, like `PythonUDAF`). - New physical operators `PythonIncrementalAggregate{Partial,Final}Exec`, routed in `SparkStrategies` as Partial -> Exchange -> Final; the buffer crosses the shuffle as an Arrow struct column. - Worker handlers: reduce-into-buffer (partial) and merge+finish (final). - `arrow_udaf` / `Aggregator` API under `pyspark.sql.pandas.aggregator`. Buffer schema is threaded to the JVM via a new nullable `bufferType` on `UserDefinedPythonFunction`. Out of scope for now (follow-ups): distinct, mixing with SQL aggregates, window/streaming, Spark Connect, SQL registration, and a typed-vs-pickled buffer perf variant. Co-authored-by: Isaac --- .../spark/api/python/PythonRunner.scala | 12 + dev/sparktestsupport/modules.py | 1 + python/pyspark/sql/pandas/aggregator.py | 175 ++++++++++++ .../arrow/test_arrow_python_aggregator.py | 158 +++++++++++ python/pyspark/sql/udf.py | 38 ++- python/pyspark/util.py | 7 + python/pyspark/worker.py | 103 +++++++ .../sql/catalyst/expressions/PythonUDF.scala | 47 +++ .../spark/sql/execution/SparkStrategies.scala | 11 +- .../PythonIncrementalAggregateExec.scala | 267 ++++++++++++++++++ .../python/UserDefinedPythonFunction.scala | 34 ++- 11 files changed, 840 insertions(+), 13 deletions(-) create mode 100644 python/pyspark/sql/pandas/aggregator.py create mode 100644 python/pyspark/sql/tests/arrow/test_arrow_python_aggregator.py create mode 100644 sql/core/src/main/scala/org/apache/spark/sql/execution/python/PythonIncrementalAggregateExec.scala 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..1f541576a222c 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", diff --git a/python/pyspark/sql/pandas/aggregator.py b/python/pyspark/sql/pandas/aggregator.py new file mode 100644 index 0000000000000..922f17b5249c3 --- /dev/null +++ b/python/pyspark/sql/pandas/aggregator.py @@ -0,0 +1,175 @@ +# +# 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 PySparkTypeError +from pyspark.sql.types import DataType, StructType +from pyspark.util import PythonEvalType + +__all__ = ["Aggregator", "arrow_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.2.0 + + Examples + -------- + A mean aggregator:: + + from pyspark.sql.pandas.aggregator import Aggregator, arrow_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 + 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 = arrow_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 NotImplementedError( + "An Aggregator is not directly callable; wrap it with arrow_udaf(...)." + ) + + +def arrow_udaf(agg: "Aggregator") -> Any: + """ + Turn an :class:`Aggregator` instance into a callable usable in ``groupBy().agg(...)``. + + .. versionadded:: 4.2.0 + + Parameters + ---------- + agg : :class:`Aggregator` + The aggregator instance. + + Returns + ------- + function + A callable that, applied to input columns, produces an aggregate :class:`Column`. + """ + from pyspark.sql.udf import UserDefinedFunction + + 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. + udf_obj.bufferSchema = agg.bufferSchema # type: ignore[attr-defined] + return udf_obj._wrapped() 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..80ee22f5c40f4 --- /dev/null +++ b/python/pyspark/sql/tests/arrow/test_arrow_python_aggregator.py @@ -0,0 +1,158 @@ +# +# 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 pyspark.sql import functions as sf +from pyspark.sql.types import ( + DoubleType, + LongType, + StructField, + StructType, +) +from pyspark.testing.sqlutils import ( + ReusedSQLTestCase, + have_pyarrow, + pyarrow_requirement_message, +) + + +if have_pyarrow: + from pyspark.sql.pandas.aggregator import Aggregator, arrow_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 + return (buffer[0] + (v or 0.0), 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 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(arrow_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(arrow_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_custom_buffer(self): + df = self._data() + result = ( + df.groupBy("k") + .agg(arrow_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_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(arrow_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]) + + +class ArrowPythonAggregatorTests(ArrowPythonAggregatorTestsMixin, ReusedSQLTestCase): + pass + + +if __name__ == "__main__": + from pyspark.sql.tests.arrow.test_arrow_python_aggregator import * # noqa: F401 + + try: + import xmlrunner # type: ignore[import] + + testRunner = xmlrunner.XMLTestRunner(output="target/test-reports", verbosity=2) + except ImportError: + testRunner = None + unittest.main(testRunner=testRunner, verbosity=2) diff --git a/python/pyspark/sql/udf.py b/python/pyspark/sql/udf.py index dbfc1b4650864..23b8cd664e037 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: diff --git a/python/pyspark/util.py b/python/pyspark/util.py index 76be9bcc9998e..183d6ac6e9768 100644 --- a/python/pyspark/util.py +++ b/python/pyspark/util.py @@ -699,6 +699,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.pandas.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..71c1472007815 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,96 @@ 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. + # `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))] + + def grouped_func( + split_index: int, data: Iterator["GroupedBatch"] + ) -> Iterator[pa.RecordBatch]: + for group in data: + batch_list = list(group) + if not batch_list: + continue + if hasattr(pa, "concat_batches"): + concatenated = pa.concat_batches(batch_list) + else: + concatenated = pa.RecordBatch.from_struct_array( + pa.concat_arrays([b.to_struct_array() for b in batch_list]) + ) + num_rows = concatenated.num_rows + result_arrays = [] + for i, (agg, args_offsets, _, _) in enumerate(udfs): + cols = [concatenated.column(o).to_pylist() for o in args_offsets] + buffer = agg.zero() + for r in range(num_rows): + buffer = agg.reduce(buffer, tuple(c[r] for c in cols)) + field_names = [f.name for f in agg.bufferSchema.fields] + struct_value = {name: buffer[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`. Each aggregator's single input column is its + # intermediate-buffer struct column. + 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, + ) + + def grouped_func( + split_index: int, data: Iterator["GroupedBatch"] + ) -> Iterator[pa.RecordBatch]: + for group in data: + batch_list = list(group) + if not batch_list: + continue + if hasattr(pa, "concat_batches"): + concatenated = pa.concat_batches(batch_list) + else: + concatenated = pa.RecordBatch.from_struct_array( + pa.concat_arrays([b.to_struct_array() for b in batch_list]) + ) + results = [] + for agg, args_offsets, _, _ in udfs: + field_names = [f.name for f in agg.bufferSchema.fields] + buffer_rows = concatenated.column(args_offsets[0]).to_pylist() + merged = None + for row in buffer_rows: + partial = tuple(row[name] for name in field_names) + merged = partial if merged is None else agg.merge(merged, partial) + if merged is None: + merged = agg.zero() + results.append(agg.finish(merged)) + result_arrays = [pa.array([r]) for r in 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..ee26d211b8333 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,53 @@ 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 [[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/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..e2a6ed1154ae3 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,6 +800,14 @@ 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 .map(_.aggregateFunction) 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..5c7af2f69cd7e --- /dev/null +++ b/sql/core/src/main/scala/org/apache/spark/sql/execution/python/PythonIncrementalAggregateExec.scala @@ -0,0 +1,267 @@ +/* + * 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) + + 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 + + inputRDD.mapPartitionsInternal { iter => if (iter.isEmpty) iter else { + val prunedProj = UnsafeProjection.create(allInputs.toSeq, childOutput) + + val groupedItr = if (groupingExprs.isEmpty) { + Iterator((new UnsafeRow(), iter)) + } else { + GroupedIterator(iter, groupingExprs, childOutput) + } + val grouped = groupedItr.map { case (key, rows) => (key, rows.map(prunedProj)) } + + 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 + + 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 }) } From b8310d57620e6777ed7b987a319db22d1ddf14f2 Mon Sep 17 00:00:00 2001 From: Hyukjin Kwon Date: Wed, 12 Aug 2026 14:45:09 +0900 Subject: [PATCH 02/11] [SQL][PYTHON][CONNECT] Support incremental Python aggregators over Spark Connect Wire the incremental Python aggregator (arrow_udaf / Aggregator) through Spark Connect so it works in remote sessions as well as classic. - Proto: add optional `buffer_type` (DataType) to the `PythonUDF` message and regenerate the Python stubs. - Connect client: `PythonUDF` expression wrapper carries `buffer_type` and serializes it into the proto; `UserDefinedFunction` forwards a `bufferSchema` attribute. `arrow_udaf` now dispatches on `is_remote()` to build the Connect UDF in a remote session. - Connect server: `SparkConnectPlanner.createUserDefinedPythonFunction` threads `buffer_type` into `UserDefinedPythonFunction`, and `transformPythonFuncExpression` builds `PythonAggregate` for the incremental eval type. Execution then reuses the same operators/worker code as classic. - Test: `ArrowPythonAggregatorParityTests` runs the same mixin under `ReusedConnectTestCase`. Co-authored-by: Isaac --- dev/sparktestsupport/modules.py | 1 + python/pyspark/sql/connect/expressions.py | 5 ++ .../sql/connect/proto/expressions_pb2.py | 48 +++++++++---------- .../sql/connect/proto/expressions_pb2.pyi | 24 +++++++++- python/pyspark/sql/connect/udf.py | 2 + python/pyspark/sql/pandas/aggregator.py | 7 ++- .../test_parity_arrow_python_aggregator.py | 29 +++++++++++ .../protobuf/spark/connect/expressions.proto | 3 ++ .../connect/planner/SparkConnectPlanner.scala | 5 +- 9 files changed, 97 insertions(+), 27 deletions(-) create mode 100644 python/pyspark/sql/tests/connect/arrow/test_parity_arrow_python_aggregator.py diff --git a/dev/sparktestsupport/modules.py b/dev/sparktestsupport/modules.py index 1f541576a222c..e1daa401c1dd0 100644 --- a/dev/sparktestsupport/modules.py +++ b/dev/sparktestsupport/modules.py @@ -1254,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/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..c259938320bdb 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.pandas.aggregator). + buffer_type=getattr(self, "bufferSchema", None), ) return CommonInlineUserDefinedFunction( function_name=self._name, diff --git a/python/pyspark/sql/pandas/aggregator.py b/python/pyspark/sql/pandas/aggregator.py index 922f17b5249c3..93fb03cd0d7ba 100644 --- a/python/pyspark/sql/pandas/aggregator.py +++ b/python/pyspark/sql/pandas/aggregator.py @@ -141,7 +141,12 @@ def arrow_udaf(agg: "Aggregator") -> Any: function A callable that, applied to input columns, produces an aggregate :class:`Column`. """ - from pyspark.sql.udf import UserDefinedFunction + from pyspark.sql.utils import is_remote + + if is_remote(): + from pyspark.sql.connect.udf import UserDefinedFunction + else: + from pyspark.sql.udf import UserDefinedFunction if not isinstance(agg, Aggregator): raise PySparkTypeError( 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/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 = { From 03083b9f6138a8f82f3630f862347d7a299539ac Mon Sep 17 00:00:00 2001 From: Hyukjin Kwon Date: Wed, 12 Aug 2026 14:55:11 +0900 Subject: [PATCH 03/11] [SQL][PYTHON] Rename arrow_udaf to udaf and require PyArrow Name the factory `udaf` to mirror Scala's `functions.udaf(agg)`, and require a supported PyArrow version up front (via require_minimum_pyarrow_version) with a clear error, since the aggregator transfers its intermediate buffer as Arrow. Co-authored-by: Isaac --- python/pyspark/sql/pandas/aggregator.py | 24 ++++++++++++++----- .../arrow/test_arrow_python_aggregator.py | 10 ++++---- 2 files changed, 23 insertions(+), 11 deletions(-) diff --git a/python/pyspark/sql/pandas/aggregator.py b/python/pyspark/sql/pandas/aggregator.py index 93fb03cd0d7ba..d641d5b479d49 100644 --- a/python/pyspark/sql/pandas/aggregator.py +++ b/python/pyspark/sql/pandas/aggregator.py @@ -25,7 +25,7 @@ from pyspark.sql.types import DataType, StructType from pyspark.util import PythonEvalType -__all__ = ["Aggregator", "arrow_udaf"] +__all__ = ["Aggregator", "udaf"] class Aggregator(ABC): @@ -50,7 +50,7 @@ class Aggregator(ABC): -------- A mean aggregator:: - from pyspark.sql.pandas.aggregator import Aggregator, arrow_udaf + from pyspark.sql.pandas.aggregator import Aggregator, udaf from pyspark.sql.types import StructType, StructField, DoubleType, LongType class Mean(Aggregator): @@ -78,7 +78,7 @@ def merge(self, b1, b2): def finish(self, buffer): return buffer[0] / buffer[1] if buffer[1] else None - mean = arrow_udaf(Mean()) + mean = udaf(Mean()) df.groupBy("k").agg(mean(df.v)).show() """ @@ -121,13 +121,17 @@ def finish(self, buffer: Tuple[Any, ...]) -> Any: # a function -- the worker calls :meth:`zero`/:meth:`reduce`/:meth:`merge`/:meth:`finish`. def __call__(self, *args: Any, **kwargs: Any) -> Any: raise NotImplementedError( - "An Aggregator is not directly callable; wrap it with arrow_udaf(...)." + "An Aggregator is not directly callable; wrap it with udaf(...)." ) -def arrow_udaf(agg: "Aggregator") -> Any: +def udaf(agg: "Aggregator") -> Any: """ - Turn an :class:`Aggregator` instance into a callable usable in ``groupBy().agg(...)``. + 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.2.0 @@ -140,9 +144,17 @@ def arrow_udaf(agg: "Aggregator") -> Any: ------- 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. """ + 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: diff --git a/python/pyspark/sql/tests/arrow/test_arrow_python_aggregator.py b/python/pyspark/sql/tests/arrow/test_arrow_python_aggregator.py index 80ee22f5c40f4..57c1bc6b85cde 100644 --- a/python/pyspark/sql/tests/arrow/test_arrow_python_aggregator.py +++ b/python/pyspark/sql/tests/arrow/test_arrow_python_aggregator.py @@ -31,7 +31,7 @@ if have_pyarrow: - from pyspark.sql.pandas.aggregator import Aggregator, arrow_udaf + from pyspark.sql.pandas.aggregator import Aggregator, udaf class Mean(Aggregator): @property @@ -94,7 +94,7 @@ def _data(self): def test_incremental_aggregator_matches_builtin_mean(self): df = self._data() result = ( - df.groupBy("k").agg(arrow_udaf(Mean())(sf.col("v")).alias("m")).orderBy("k").collect() + 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} @@ -103,7 +103,7 @@ def test_incremental_aggregator_matches_builtin_mean(self): def test_incremental_aggregator_no_group(self): df = self._data() - result = df.agg(arrow_udaf(Mean())(sf.col("v")).alias("m")).collect() + 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) @@ -111,7 +111,7 @@ def test_incremental_aggregator_custom_buffer(self): df = self._data() result = ( df.groupBy("k") - .agg(arrow_udaf(SumSquares())(sf.col("v")).alias("s")) + .agg(udaf(SumSquares())(sf.col("v")).alias("s")) .orderBy("k") .collect() ) @@ -133,7 +133,7 @@ def test_result_independent_of_partition_count(self): rows = ( base.repartition(n, sf.col("v")) .groupBy("k") - .agg(arrow_udaf(Mean())(sf.col("v")).alias("m")) + .agg(udaf(Mean())(sf.col("v")).alias("m")) .orderBy("k") .collect() ) From 8b0448b08ebf3e3b3dce726d71b94fdb6ae7a8ea Mon Sep 17 00:00:00 2001 From: Hyukjin Kwon Date: Wed, 12 Aug 2026 14:57:24 +0900 Subject: [PATCH 04/11] [SQL][PYTHON] Move aggregator module to pyspark.sql and set version to 4.4.0 Relocate `aggregator.py` from `pyspark.sql.pandas` to `pyspark.sql` (import as `pyspark.sql.aggregator`), and set the `versionadded` for `Aggregator`/`udaf` to 4.4.0. Update the references in util.py, connect/udf.py, and the test. Co-authored-by: Isaac --- python/pyspark/sql/{pandas => }/aggregator.py | 6 +++--- python/pyspark/sql/connect/udf.py | 2 +- .../pyspark/sql/tests/arrow/test_arrow_python_aggregator.py | 2 +- python/pyspark/util.py | 2 +- 4 files changed, 6 insertions(+), 6 deletions(-) rename python/pyspark/sql/{pandas => }/aggregator.py (98%) diff --git a/python/pyspark/sql/pandas/aggregator.py b/python/pyspark/sql/aggregator.py similarity index 98% rename from python/pyspark/sql/pandas/aggregator.py rename to python/pyspark/sql/aggregator.py index d641d5b479d49..5816679e2e303 100644 --- a/python/pyspark/sql/pandas/aggregator.py +++ b/python/pyspark/sql/aggregator.py @@ -44,13 +44,13 @@ class Aggregator(ABC): to the aggregator call. :meth:`merge` must be associative and commutative, since the framework may combine partial buffers in any order. - .. versionadded:: 4.2.0 + .. versionadded:: 4.4.0 Examples -------- A mean aggregator:: - from pyspark.sql.pandas.aggregator import Aggregator, udaf + from pyspark.sql.aggregator import Aggregator, udaf from pyspark.sql.types import StructType, StructField, DoubleType, LongType class Mean(Aggregator): @@ -133,7 +133,7 @@ def udaf(agg: "Aggregator") -> Any: The aggregator is executed with true incremental (partial) aggregation and transfers its intermediate buffer as Arrow; PyArrow is therefore required. - .. versionadded:: 4.2.0 + .. versionadded:: 4.4.0 Parameters ---------- diff --git a/python/pyspark/sql/connect/udf.py b/python/pyspark/sql/connect/udf.py index c259938320bdb..ca508bfc282c9 100644 --- a/python/pyspark/sql/connect/udf.py +++ b/python/pyspark/sql/connect/udf.py @@ -210,7 +210,7 @@ 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.pandas.aggregator). + # Set for incremental Python aggregators (see pyspark.sql.aggregator). buffer_type=getattr(self, "bufferSchema", None), ) return CommonInlineUserDefinedFunction( diff --git a/python/pyspark/sql/tests/arrow/test_arrow_python_aggregator.py b/python/pyspark/sql/tests/arrow/test_arrow_python_aggregator.py index 57c1bc6b85cde..5a55d78e7ba2f 100644 --- a/python/pyspark/sql/tests/arrow/test_arrow_python_aggregator.py +++ b/python/pyspark/sql/tests/arrow/test_arrow_python_aggregator.py @@ -31,7 +31,7 @@ if have_pyarrow: - from pyspark.sql.pandas.aggregator import Aggregator, udaf + from pyspark.sql.aggregator import Aggregator, udaf class Mean(Aggregator): @property diff --git a/python/pyspark/util.py b/python/pyspark/util.py index 183d6ac6e9768..84b8a41ed957e 100644 --- a/python/pyspark/util.py +++ b/python/pyspark/util.py @@ -699,7 +699,7 @@ 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.pandas.aggregator``. + # 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. From 3e2bb59c46368f0dc47f15681f6cb05f936c8411 Mon Sep 17 00:00:00 2001 From: Hyukjin Kwon Date: Wed, 12 Aug 2026 15:30:01 +0900 Subject: [PATCH 05/11] [SQL][PYTHON][CONNECT] Support spark.udf.register for the incremental aggregator Allow `spark.udf.register(name, udaf(agg))` so the incremental Python aggregator can be invoked from SQL text (`SELECT my_agg(v) FROM t GROUP BY k`), matching Scala's `spark.udf.register(name, functions.udaf(agg))`. - Classic and Connect `UDFRegistration.register` accept SQL_GROUPED_AGG_ARROW_INCREMENTAL_FINAL_UDF and thread the buffer schema through (classic `register` reconstructs the UDF and would otherwise drop it; Connect passes it via `SparkConnectClient.register_udf` -> the PythonUDF proto). - `udaf` sets `bufferSchema` on the returned wrapper too, so it survives registration. The Connect server already builds `PythonAggregate` in `handleRegisterUserDefinedFunction` via the shared `createUserDefinedPythonFunction`. - Test: `test_sql_registration` in the shared mixin (runs classic + Connect). Co-authored-by: Isaac --- python/pyspark/sql/aggregator.py | 7 +++++-- python/pyspark/sql/connect/client/core.py | 3 +++ python/pyspark/sql/connect/udf.py | 12 ++++++++++-- .../sql/tests/arrow/test_arrow_python_aggregator.py | 13 +++++++++++++ python/pyspark/sql/udf.py | 8 +++++++- 5 files changed, 38 insertions(+), 5 deletions(-) diff --git a/python/pyspark/sql/aggregator.py b/python/pyspark/sql/aggregator.py index 5816679e2e303..470e94f67f314 100644 --- a/python/pyspark/sql/aggregator.py +++ b/python/pyspark/sql/aggregator.py @@ -187,6 +187,9 @@ def udaf(agg: "Aggregator") -> Any: deterministic=True, ) # Threaded to the JVM in UserDefinedFunction._create_judf so PythonAggregate can plan the - # two-stage aggregation. + # 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] - return udf_obj._wrapped() + 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/udf.py b/python/pyspark/sql/connect/udf.py index ca508bfc282c9..06521fb87c712 100644 --- a/python/pyspark/sql/connect/udf.py +++ b/python/pyspark/sql/connect/udf.py @@ -305,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", @@ -313,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/tests/arrow/test_arrow_python_aggregator.py b/python/pyspark/sql/tests/arrow/test_arrow_python_aggregator.py index 5a55d78e7ba2f..98271183ac9e9 100644 --- a/python/pyspark/sql/tests/arrow/test_arrow_python_aggregator.py +++ b/python/pyspark/sql/tests/arrow/test_arrow_python_aggregator.py @@ -141,6 +141,19 @@ def test_result_independent_of_partition_count(self): 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 diff --git a/python/pyspark/sql/udf.py b/python/pyspark/sql/udf.py index 23b8cd664e037..f1f6fcd70b753 100644 --- a/python/pyspark/sql/udf.py +++ b/python/pyspark/sql/udf.py @@ -859,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", @@ -867,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( @@ -878,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: From 19a53da66390c2aabb5fa7072e89d88901b51f0d Mon Sep 17 00:00:00 2001 From: Hyukjin Kwon Date: Wed, 12 Aug 2026 19:38:54 +0900 Subject: [PATCH 06/11] [SQL][PYTHON] Fix custom-error lint and test import - Use PySparkNotImplementedError instead of a raw NotImplementedError in Aggregator.__call__ (PySpark custom-errors linter). - Import have_pyarrow / pyarrow_requirement_message from pyspark.testing.utils (not sqlutils), which was causing the aggregator test modules to fail at import. Co-authored-by: Isaac --- python/pyspark/sql/aggregator.py | 7 ++++--- .../sql/tests/arrow/test_arrow_python_aggregator.py | 4 ++-- 2 files changed, 6 insertions(+), 5 deletions(-) diff --git a/python/pyspark/sql/aggregator.py b/python/pyspark/sql/aggregator.py index 470e94f67f314..35c750f6c3035 100644 --- a/python/pyspark/sql/aggregator.py +++ b/python/pyspark/sql/aggregator.py @@ -21,7 +21,7 @@ from abc import ABC, abstractmethod from typing import Any, Tuple -from pyspark.errors import PySparkTypeError +from pyspark.errors import PySparkNotImplementedError, PySparkTypeError from pyspark.sql.types import DataType, StructType from pyspark.util import PythonEvalType @@ -120,8 +120,9 @@ def finish(self, buffer: Tuple[Any, ...]) -> Any: # 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 NotImplementedError( - "An Aggregator is not directly callable; wrap it with udaf(...)." + raise PySparkNotImplementedError( + errorClass="NOT_IMPLEMENTED", + messageParameters={"feature": "calling an Aggregator directly; wrap it with udaf(...)"}, ) diff --git a/python/pyspark/sql/tests/arrow/test_arrow_python_aggregator.py b/python/pyspark/sql/tests/arrow/test_arrow_python_aggregator.py index 98271183ac9e9..f2348390f8e2c 100644 --- a/python/pyspark/sql/tests/arrow/test_arrow_python_aggregator.py +++ b/python/pyspark/sql/tests/arrow/test_arrow_python_aggregator.py @@ -23,8 +23,8 @@ StructField, StructType, ) -from pyspark.testing.sqlutils import ( - ReusedSQLTestCase, +from pyspark.testing.sqlutils import ReusedSQLTestCase +from pyspark.testing.utils import ( have_pyarrow, pyarrow_requirement_message, ) From a12f8256b2eb1a89cf48d08cf070211422d518f8 Mon Sep 17 00:00:00 2001 From: Hyukjin Kwon Date: Thu, 13 Aug 2026 08:14:59 +0900 Subject: [PATCH 07/11] [SQL][PYTHON] Address review: stream partial batches and fix empty global aggregation Two blocking review items: - PARTIAL/FINAL worker handlers now stream Arrow batches and fold them one at a time into the per-aggregator buffers, instead of `list(group)` + concatenating the whole group first. Map-side peak memory is bounded by a single batch (plus the buffers), not the whole group -- the point of the incremental API. - A global (no-grouping) aggregation over empty input now returns the identity row `finish(zero)` instead of no row. GroupedPythonArrowInput cannot transmit an empty group, so the FINAL stage (which runs on a single AllTuples partition) injects one all-null buffer row; the worker skips null partial buffers and so merges nothing, yielding `finish(zero)`. Added a focused test (`df.limit(0).agg(udaf(...))`). Co-authored-by: Isaac --- .../arrow/test_arrow_python_aggregator.py | 8 +++ python/pyspark/worker.py | 67 +++++++++---------- .../PythonIncrementalAggregateExec.scala | 32 ++++++++- 3 files changed, 68 insertions(+), 39 deletions(-) diff --git a/python/pyspark/sql/tests/arrow/test_arrow_python_aggregator.py b/python/pyspark/sql/tests/arrow/test_arrow_python_aggregator.py index f2348390f8e2c..20c0b583a2598 100644 --- a/python/pyspark/sql/tests/arrow/test_arrow_python_aggregator.py +++ b/python/pyspark/sql/tests/arrow/test_arrow_python_aggregator.py @@ -107,6 +107,14 @@ def test_incremental_aggregator_no_group(self): 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 = ( diff --git a/python/pyspark/worker.py b/python/pyspark/worker.py index 71c1472007815..81b5ea392cd22 100644 --- a/python/pyspark/worker.py +++ b/python/pyspark/worker.py @@ -2217,6 +2217,8 @@ def grouped_func( # 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( @@ -2231,24 +2233,18 @@ def grouped_func( split_index: int, data: Iterator["GroupedBatch"] ) -> Iterator[pa.RecordBatch]: for group in data: - batch_list = list(group) - if not batch_list: - continue - if hasattr(pa, "concat_batches"): - concatenated = pa.concat_batches(batch_list) - else: - concatenated = pa.RecordBatch.from_struct_array( - pa.concat_arrays([b.to_struct_array() for b in batch_list]) - ) - num_rows = concatenated.num_rows + 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, args_offsets, _, _) in enumerate(udfs): - cols = [concatenated.column(o).to_pylist() for o in args_offsets] - buffer = agg.zero() - for r in range(num_rows): - buffer = agg.reduce(buffer, tuple(c[r] for c in cols)) + for i, (agg, _, _, _) in enumerate(udfs): field_names = [f.name for f in agg.bufferSchema.fields] - struct_value = {name: buffer[j] for j, name in enumerate(field_names)} + 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) ) @@ -2262,8 +2258,10 @@ def grouped_func( 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`. Each aggregator's single input column is its - # intermediate-buffer struct column. + # 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)]), @@ -2275,26 +2273,21 @@ def grouped_func( split_index: int, data: Iterator["GroupedBatch"] ) -> Iterator[pa.RecordBatch]: for group in data: - batch_list = list(group) - if not batch_list: - continue - if hasattr(pa, "concat_batches"): - concatenated = pa.concat_batches(batch_list) - else: - concatenated = pa.RecordBatch.from_struct_array( - pa.concat_arrays([b.to_struct_array() for b in batch_list]) - ) + merged: list = [None] * len(udfs) + for batch in group: + for i, (agg, args_offsets, _, _) in enumerate(udfs): + field_names = [f.name for f in agg.bufferSchema.fields] + 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 agg, args_offsets, _, _ in udfs: - field_names = [f.name for f in agg.bufferSchema.fields] - buffer_rows = concatenated.column(args_offsets[0]).to_pylist() - merged = None - for row in buffer_rows: - partial = tuple(row[name] for name in field_names) - merged = partial if merged is None else agg.merge(merged, partial) - if merged is None: - merged = agg.zero() - results.append(agg.finish(merged)) + for i, (agg, _, _, _) in enumerate(udfs): + m = merged[i] if merged[i] is not None else agg.zero() + results.append(agg.finish(m)) result_arrays = [pa.array([r]) for r in results] batch = pa.RecordBatch.from_arrays(result_arrays, col_names) yield ArrowBatchTransformer.enforce_schema(batch, return_schema) 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 index 5c7af2f69cd7e..a937296e703e3 100644 --- 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 @@ -70,6 +70,13 @@ abstract class PythonIncrementalAggregateExecBase extends UnaryExecNode with Pyt /** 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) @@ -123,7 +130,8 @@ abstract class PythonIncrementalAggregateExecBase extends UnaryExecNode with Pyt val resultExprs = outputExpressions val localEvalType = evalType - inputRDD.mapPartitionsInternal { iter => if (iter.isEmpty) iter else { + val emitIdentityOnEmpty = emitOnEmptyPartition + inputRDD.mapPartitionsInternal { iter => if (iter.isEmpty && !emitIdentityOnEmpty) iter else { val prunedProj = UnsafeProjection.create(allInputs.toSeq, childOutput) val groupedItr = if (groupingExprs.isEmpty) { @@ -131,7 +139,23 @@ abstract class PythonIncrementalAggregateExecBase extends UnaryExecNode with Pyt } else { GroupedIterator(iter, groupingExprs, childOutput) } - val grouped = groupedItr.map { case (key, rows) => (key, rows.map(prunedProj)) } + + // 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() @@ -226,6 +250,10 @@ case class PythonIncrementalAggregateFinalExec( 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 From 39170897afbbe35e55d02f662b97bb505c2b9e0c Mon Sep 17 00:00:00 2001 From: Hyukjin Kwon Date: Thu, 13 Aug 2026 08:35:57 +0900 Subject: [PATCH 08/11] [SQL][PYTHON] Clarify mixed-UDF error and test multiple aggregators - invalidPandasUDFPlacementError now also names incremental PythonAggregate functions (not just grouped-agg PythonUDAF) when Python aggregate UDFs are mixed with other aggregate functions in one Aggregate. - Add a test with two incremental aggregators (different buffer schemas) over the same input, covering multi-UDF partial/final planning and execution. Co-authored-by: Isaac --- .../arrow/test_arrow_python_aggregator.py | 27 +++++++++++++++++++ .../spark/sql/execution/SparkStrategies.scala | 14 ++++++---- 2 files changed, 36 insertions(+), 5 deletions(-) diff --git a/python/pyspark/sql/tests/arrow/test_arrow_python_aggregator.py b/python/pyspark/sql/tests/arrow/test_arrow_python_aggregator.py index 20c0b583a2598..cd8614d207a0d 100644 --- a/python/pyspark/sql/tests/arrow/test_arrow_python_aggregator.py +++ b/python/pyspark/sql/tests/arrow/test_arrow_python_aggregator.py @@ -131,6 +131,33 @@ def test_incremental_aggregator_custom_buffer(self): for k in exp: self.assertAlmostEqual(got[k], exp[k], places=6) + 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( 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 e2a6ed1154ae3..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 @@ -809,12 +809,16 @@ abstract class SparkStrategies extends QueryPlanner[SparkPlan] { 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 } From 429d7703f1d743bd12a28757420d486be86889d9 Mon Sep 17 00:00:00 2001 From: Hyukjin Kwon Date: Thu, 13 Aug 2026 12:10:10 +0900 Subject: [PATCH 09/11] [DO-NOT-MERGE][PYTHON] Fix CI: ruff errors and INVALID_UDF_EVAL_TYPE message - Define ArrowGroupedAggIncremental{Partial,Final}UDFType Literal aliases and import them under TYPE_CHECKING so the eval-type annotations resolve (F821). - Use the standard `from pyspark.testing import main` test footer instead of `import *` (F403 / RUF100); reformat with ruff. - Update the INVALID_UDF_EVAL_TYPE expected message in test_pandas_grouped_map to include SQL_GROUPED_AGG_ARROW_INCREMENTAL_FINAL_UDF now that the incremental aggregator is registerable via spark.udf.register. Co-authored-by: Isaac --- .../pyspark/sql/pandas/_typing/__init__.pyi | 2 ++ .../arrow/test_arrow_python_aggregator.py | 23 ++++--------------- .../tests/pandas/test_pandas_grouped_map.py | 3 ++- python/pyspark/util.py | 2 ++ 4 files changed, 11 insertions(+), 19 deletions(-) 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 index cd8614d207a0d..f3801d6fd044a 100644 --- a/python/pyspark/sql/tests/arrow/test_arrow_python_aggregator.py +++ b/python/pyspark/sql/tests/arrow/test_arrow_python_aggregator.py @@ -36,9 +36,7 @@ class Mean(Aggregator): @property def bufferSchema(self): - return StructType( - [StructField("sum", DoubleType()), StructField("count", LongType())] - ) + return StructType([StructField("sum", DoubleType()), StructField("count", LongType())]) @property def outputType(self): @@ -93,9 +91,7 @@ def _data(self): 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() - ) + 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} @@ -118,10 +114,7 @@ def test_incremental_aggregator_empty_global_input(self): 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() + 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() @@ -195,12 +188,6 @@ class ArrowPythonAggregatorTests(ArrowPythonAggregatorTestsMixin, ReusedSQLTestC if __name__ == "__main__": - from pyspark.sql.tests.arrow.test_arrow_python_aggregator import * # noqa: F401 - - try: - import xmlrunner # type: ignore[import] + from pyspark.testing import main - testRunner = xmlrunner.XMLTestRunner(output="target/test-reports", verbosity=2) - except ImportError: - testRunner = None - unittest.main(testRunner=testRunner, verbosity=2) + 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/util.py b/python/pyspark/util.py index 84b8a41ed957e..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 ( From bb092e03de7480031e91184c18d604743caef493 Mon Sep 17 00:00:00 2001 From: Hyukjin Kwon Date: Thu, 13 Aug 2026 13:54:27 +0900 Subject: [PATCH 10/11] [DO-NOT-MERGE][PYTHON] Fix ruff format in aggregator.py and worker.py Reformat with ruff 0.14.0 to match the CI-pinned version (files were previously formatted with an older local ruff). Co-authored-by: Isaac --- python/pyspark/sql/aggregator.py | 5 ++--- python/pyspark/worker.py | 4 +--- 2 files changed, 3 insertions(+), 6 deletions(-) diff --git a/python/pyspark/sql/aggregator.py b/python/pyspark/sql/aggregator.py index 35c750f6c3035..a9942b4529685 100644 --- a/python/pyspark/sql/aggregator.py +++ b/python/pyspark/sql/aggregator.py @@ -18,6 +18,7 @@ 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 @@ -105,9 +106,7 @@ def reduce(self, buffer: Tuple[Any, ...], value: Tuple[Any, ...]) -> Tuple[Any, ... @abstractmethod - def merge( - self, buffer1: Tuple[Any, ...], buffer2: Tuple[Any, ...] - ) -> Tuple[Any, ...]: + def merge(self, buffer1: Tuple[Any, ...], buffer2: Tuple[Any, ...]) -> Tuple[Any, ...]: """Merge two partial buffers into one. Must be associative and commutative.""" ... diff --git a/python/pyspark/worker.py b/python/pyspark/worker.py index 81b5ea392cd22..65bcc2f88ff10 100644 --- a/python/pyspark/worker.py +++ b/python/pyspark/worker.py @@ -2245,9 +2245,7 @@ def grouped_func( for i, (agg, _, _, _) in enumerate(udfs): field_names = [f.name for f in agg.bufferSchema.fields] 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) - ) + 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) From 1d6c9391afaa5dff7429767040b8d02315490cb2 Mon Sep 17 00:00:00 2001 From: Hyukjin Kwon Date: Thu, 13 Aug 2026 14:30:13 +0900 Subject: [PATCH 11/11] [DO-NOT-MERGE][PYTHON] Fix mypy incompatible import in aggregator.py Silence mypy's [assignment] error on the classic UserDefinedFunction import in the is_remote() dispatch, matching the Connect/classic dispatch pattern. Co-authored-by: Isaac --- python/pyspark/sql/aggregator.py | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) diff --git a/python/pyspark/sql/aggregator.py b/python/pyspark/sql/aggregator.py index a9942b4529685..1de1628b4b9cb 100644 --- a/python/pyspark/sql/aggregator.py +++ b/python/pyspark/sql/aggregator.py @@ -158,7 +158,9 @@ def udaf(agg: "Aggregator") -> Any: if is_remote(): from pyspark.sql.connect.udf import UserDefinedFunction else: - from pyspark.sql.udf import UserDefinedFunction + # 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(