/
githubmirror
/
spark
Обзор
Документация
Войти
/
githubmirror
/
spark
Код
Запросы
0
Пакеты
0
Релизы
0
Аналитика
Безопасность
master
python/pyspark/worker.py
4 828 строк
208 KB
Hyukjin Kwon
[SPARK-58695][SQL][PYTHON] Support scalar pandas and Arrow UDFs inside higher-order function lambdas
13 часов назад
13 часов назад
15af805
Код
Авторство
О чём код?
# # 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. # """ Worker that receives input from Piped RDD. """ import os import sys import dataclasses import time import inspect import itertools import json import warnings from collections.abc import Iterator from typing import ( Any, Callable, Iterable, Optional, Tuple, Type, TypeVar, TYPE_CHECKING, Union, get_args, get_origin, BinaryIO, ) T = TypeVar("T") if TYPE_CHECKING: import pandas as pd import pyarrow as pa from pyspark.sql.pandas._typing import GroupedBatch from pyspark.accumulators import ( SpecialAccumulatorIds, _accumulatorRegistry, _deserialize_accumulator, ) from pyspark.sql.streaming.stateful_processor_api_client import StatefulProcessorApiClient from pyspark.sql.streaming.stateful_processor_util import TransformWithStateInPandasFuncMode from pyspark.taskcontext import BarrierTaskContext, TaskContext from pyspark.util import PythonEvalType from pyspark.serializers import ( write_int, write_long, SpecialLengths, CPickleSerializer, BatchedSerializer, ) from pyspark.sql.conversion import ( LocalDataToArrowConversion, ArrowTableToRowsConversion, ArrowBatchTransformer, PandasToArrowConversion, ) from pyspark.sql.functions import SkipRestOfInputTableException from pyspark.sql.pandas.serializers import ( ArrowStreamSerializer, ArrowStreamGroupSerializer, ArrowStreamCoGroupSerializer, ) from pyspark.sql.pandas.types import to_arrow_schema, to_arrow_type from pyspark.sql.types import ( ArrayType, BinaryType, DataType, IntegerType, LongType, MapType, Row, StringType, StructField, StructType, _create_row, _parse_datatype_json_string, ) from pyspark.util import ( fail_on_stopiteration, handle_worker_exception, with_faulthandler, start_faulthandler_periodic_traceback, ) from pyspark import _NoValue, shuffle from pyspark.errors import PySparkRuntimeError, PySparkTypeError, PySparkValueError from pyspark.worker_message import WorkerInitInfo from pyspark.worker_util import ( check_python_version, get_sock_file_to_executor, read_command, pickleSer, send_accumulator_updates, setup_broadcasts, setup_memory_limits, setup_spark_files, Conf, ) from pyspark.logger.worker_io import capture_outputs from pyspark.messages import ( SparkMessageReceiver, SparkSocketMessageReceiver, ) class RunnerConf(Conf): @property def assign_cols_by_name(self) -> bool: return ( self.get("spark.sql.legacy.execution.pandas.groupedMap.assignColumnsByName", "true") == "true" ) @property def use_large_var_types(self) -> bool: return self.get("spark.sql.execution.arrow.useLargeVarTypes", "false") == "true" @property def use_legacy_pandas_udf_conversion(self) -> bool: return ( self.get("spark.sql.legacy.execution.pythonUDF.pandas.conversion.enabled", "false") == "true" ) @property def use_legacy_pandas_udtf_conversion(self) -> bool: return ( self.get("spark.sql.legacy.execution.pythonUDTF.pandas.conversion.enabled", "false") == "true" ) @property def binary_as_bytes(self) -> bool: return self.get("spark.sql.execution.pyspark.binaryAsBytes", "true") == "true" @property def safecheck(self) -> bool: return self.get("spark.sql.execution.pandas.convertToArrowArraySafely", "false") == "true" @property def int_to_decimal_coercion_enabled(self) -> bool: return ( self.get("spark.sql.execution.pythonUDF.pandas.intToDecimalCoercionEnabled", "false") == "true" ) @property def prefer_int_ext_dtype(self) -> bool: return ( self.get("spark.sql.execution.pythonUDF.pandas.preferIntExtensionDtype", "false") == "true" ) @property def timezone(self) -> Optional[str]: return self.get("spark.sql.session.timeZone", None, lower_str=False) @property def arrow_max_records_per_batch(self) -> int: return int(self.get("spark.sql.execution.arrow.maxRecordsPerBatch", 10000)) @property def arrow_max_bytes_per_batch(self) -> int: return int(self.get("spark.sql.execution.arrow.maxBytesPerBatch", 2**31 - 1)) @property def arrow_concurrency_level(self) -> int: return int(self.get("spark.sql.execution.pythonUDF.arrow.concurrency.level", -1)) @property def udf_profiler(self) -> Optional[str]: return self.get("spark.sql.pyspark.udf.profiler", None) @property def data_source_profiler(self) -> Optional[str]: return self.get("spark.sql.pyspark.dataSource.profiler", None) class EvalConf(Conf): @property def state_value_schema(self) -> Optional[StructType]: schema = self.get("state_value_schema", None) if schema is None: return None return StructType.fromJson(json.loads(schema)) @property def grouping_key_schema(self) -> Optional[StructType]: schema = self.get("grouping_key_schema", None) if schema is None: return None return StructType.fromJson(json.loads(schema)) @property def state_server_socket_port(self) -> Optional[int | str]: port = self.get("state_server_socket_port", None) try: return int(port) except ValueError: return port @property def input_type(self) -> Optional[DataType]: input_type = self.get("input_type", None, lower_str=False) if input_type is None: return None return _parse_datatype_json_string(input_type) @property def table_arg_offsets(self) -> Optional[list[int]]: offsets = self.get("table_arg_offsets", None) if offsets is None: return None return [int(x) for x in offsets.split(",") if x] def report_times(outfile, boot, init, finish, processing_time_ms): write_int(SpecialLengths.TIMING_DATA, outfile) write_long(int(1000 * boot), outfile) write_long(int(1000 * init), outfile) write_long(int(1000 * finish), outfile) write_long(processing_time_ms, outfile) def chain(f, g): """chain two functions together""" return lambda *a: g(f(*a)) def verify_return_type(result: T, expected_type: Type[T]) -> T: """ Verify a UDF return value against an expected type. Returns ``result`` unchanged if ``isinstance(result, expected_type)``. For ``Iterator[T]``, returns a lazy iterator that checks each element against ``T`` on consumption. Raises ``PySparkTypeError`` on mismatch. """ if get_origin(expected_type) is Iterator: (element_type,) = get_args(expected_type) label = f"iterator of {_top_level_package(element_type)}.{element_type.__name__}" if not isinstance(result, Iterator): raise PySparkTypeError( errorClass="UDF_RETURN_TYPE", messageParameters={"expected": label, "actual": type(result).__name__}, ) def check_element(element: T) -> T: if not isinstance(element, element_type): raise PySparkTypeError( errorClass="UDF_RETURN_TYPE", messageParameters={ "expected": label, "actual": f"iterator of {type(element).__name__}", }, ) return element return map(check_element, result) # type: ignore[return-value] if not isinstance(result, expected_type): raise PySparkTypeError( errorClass="UDF_RETURN_TYPE", messageParameters={ "expected": f"{_top_level_package(expected_type)}.{expected_type.__name__}", "actual": type(result).__name__, }, ) return result def _top_level_package(t: type) -> str: """Return the top-level package of ``t`` (``pandas`` for ``pd.DataFrame``).""" return (t.__module__ or "").split(".", 1)[0] def verify_result_row_count(result_length: int, expected: int) -> None: """Raise if the result row count doesn't match the expected input row count.""" if result_length != expected: raise PySparkRuntimeError( errorClass="RESULT_ROWS_MISMATCH", messageParameters={ "output_length": str(result_length), "input_length": str(expected), }, ) def verify_scalar_result(result: Any, num_rows: int) -> Any: """ Verify a scalar UDF result is array-like and has the expected number of rows. Parameters ---------- result : Any The UDF result to verify. num_rows : int Expected number of rows (must match input batch size). """ try: result_length = len(result) except TypeError: raise PySparkTypeError( errorClass="UDF_RETURN_TYPE", messageParameters={ "expected": "array-like object", "actual": type(result).__name__, }, ) verify_result_row_count(result_length, num_rows) return result def verify_iterator_exhausted(iterator: Iterator) -> None: """Verify that an iterator has been fully consumed.""" try: next(iterator) except StopIteration: pass else: raise PySparkRuntimeError(errorClass="INPUT_NOT_FULLY_CONSUMED", messageParameters={}) def verify_output_row_limit( iterator: Iterator, max_rows: Union[int, Callable[[], int]], ) -> Iterator: """Yield elements while verifying total rows do not exceed a limit (fail-fast).""" total_rows = 0 for element in iterator: total_rows += len(element) if total_rows > (max_rows() if callable(max_rows) else max_rows): raise PySparkRuntimeError(errorClass="OUTPUT_EXCEEDS_INPUT_ROWS", messageParameters={}) yield element def verify_iter_result_row_count( iterator: Iterator, expected_rows: Callable[[], int], ) -> Iterator: """Yield elements and verify final row count matches expected exactly. ``expected_rows`` is a callable because the expected count is only known once the iterator is fully consumed (input rows are counted lazily as a side effect of pulling batches), so it must be read after this generator is exhausted. """ actual_rows = 0 for element in iterator: actual_rows += len(element) yield element verify_result_row_count(actual_rows, expected_rows()) def _verify_column_schema( actual_names: list, expected_names: list, *, assign_cols_by_name: bool ) -> None: """Check column names (by-name) or count (by-position) match the expected schema.""" if assign_cols_by_name: actual_set = set(actual_names) expected_set = set(expected_names) missing = sorted(expected_set.difference(actual_set)) extra = sorted(actual_set.difference(expected_set)) if missing or extra: raise PySparkRuntimeError( errorClass="RESULT_COLUMN_NAMES_MISMATCH", messageParameters={ "missing": f" Missing: {', '.join(missing)}." if missing else "", "extra": f" Unexpected: {', '.join(extra)}." if extra else "", }, ) elif len(actual_names) != len(expected_names): raise PySparkRuntimeError( errorClass="RESULT_COLUMN_SCHEMA_MISMATCH", messageParameters={ "expected": str(len(expected_names)), "actual": str(len(actual_names)), }, ) def verify_pandas_result( result: Union["pd.DataFrame", "pd.Series"], return_type: DataType, assign_cols_by_name: bool, truncate_return_schema: bool, ) -> None: import pandas as pd if not isinstance(return_type, StructType): verify_return_type(result, pd.Series) return verify_return_type(result, pd.DataFrame) # Skip schema check on a fully empty result (no rows and no columns). if result.empty and len(result.columns) == 0: return field_names = [field.name for field in return_type.fields] actual_names = ( list(result.columns[: len(field_names)]) if truncate_return_schema else list(result.columns) ) # By-name mode only applies when the result has string column names; # a numeric RangeIndex falls back to a by-position count check. by_name = assign_cols_by_name and any(isinstance(n, str) for n in result.columns) _verify_column_schema(actual_names, field_names, assign_cols_by_name=by_name) def verify_arrow_result( result: Union["pa.Table", "pa.RecordBatch"], assign_cols_by_name: bool, expected_cols_and_types: Union[dict[str, "pa.DataType"], list[tuple[str, "pa.DataType"]]], ) -> None: # Skip schema check on a fully empty result (no rows and no columns). if result.num_columns == 0 and result.num_rows == 0: return actual_names = list(result.schema.names) actual_types = list(result.schema.types) # expected_cols_and_types is a dict in by-name mode, list of (name, type) by position. if isinstance(expected_cols_and_types, dict): expected_names = list(expected_cols_and_types.keys()) else: expected_names = [name for name, _ in expected_cols_and_types] _verify_column_schema(actual_names, expected_names, assign_cols_by_name=assign_cols_by_name) if isinstance(expected_cols_and_types, dict): actual_by_name = dict(zip(actual_names, actual_types)) column_types = [ (name, expected_cols_and_types[name], actual_by_name[name]) for name in sorted(expected_cols_and_types.keys()) ] else: column_types = [ (expected_name, expected_type, actual_type) for (expected_name, expected_type), actual_type in zip( expected_cols_and_types, actual_types ) ] type_mismatch = [ (name, expected, actual) for name, expected, actual in column_types if actual != expected ] if type_mismatch: raise PySparkRuntimeError( errorClass="RESULT_COLUMN_TYPES_MISMATCH", messageParameters={ "mismatch": ", ".join( "column '{}' (expected {}, actual {})".format(name, expected, actual) for name, expected, actual in type_mismatch ) }, ) def wrap_kwargs_support(f, args_offsets, kwargs_offsets): if len(kwargs_offsets): keys = list(kwargs_offsets.keys()) len_args_offsets = len(args_offsets) if len_args_offsets > 0: def func(*args): return f(*args[:len_args_offsets], **dict(zip(keys, args[len_args_offsets:]))) else: def func(*args): return f(**dict(zip(keys, args))) return func, args_offsets + [kwargs_offsets[key] for key in keys] else: return f, args_offsets def _is_iter_based(eval_type: int) -> bool: return eval_type in ( PythonEvalType.SQL_SCALAR_PANDAS_ITER_UDF, PythonEvalType.SQL_SCALAR_ARROW_ITER_UDF, # Iterator UDFs lifted out of a higher-order function lambda keep the iterator contract: # the user function still consumes and produces an iterator of batches; the worker only # feeds it the flattened elements and re-nests the results. See ExtractPythonUDFFromLambda. PythonEvalType.SQL_SCALAR_PANDAS_ITER_ELEMENTWISE_UDF, PythonEvalType.SQL_SCALAR_ARROW_ITER_ELEMENTWISE_UDF, PythonEvalType.SQL_MAP_PANDAS_ITER_UDF, PythonEvalType.SQL_MAP_ARROW_ITER_UDF, PythonEvalType.SQL_GROUPED_MAP_PANDAS_UDF_WITH_STATE, PythonEvalType.SQL_GROUPED_MAP_ARROW_ITER_UDF, PythonEvalType.SQL_GROUPED_MAP_PANDAS_ITER_UDF, ) def wrap_perf_profiler(f, eval_type, result_id): from pyspark.sql.profiler import ProfileResultsParam, ProfileResultsParamV2, WorkerPerfProfiler accumulator = _deserialize_accumulator( SpecialAccumulatorIds.SQL_UDF_PROFIER, None, ProfileResultsParam ) accumulator_v2 = _deserialize_accumulator( SpecialAccumulatorIds.SQL_UDF_PROFIER_V2, {}, ProfileResultsParamV2 ) if _is_iter_based(eval_type): def profiling_func(*args, **kwargs): iterator = iter(f(*args, **kwargs)) while True: try: with WorkerPerfProfiler(accumulator, accumulator_v2, result_id): item = next(iterator) yield item except StopIteration: break else: def profiling_func(*args, **kwargs): with WorkerPerfProfiler(accumulator, accumulator_v2, result_id): ret = f(*args, **kwargs) return ret return profiling_func def wrap_memory_profiler(f, eval_type, result_id): from pyspark.sql.profiler import ( ProfileResultsParam, ProfileResultsParamV2, WorkerMemoryProfiler, ) import pyspark.memory_profiler_ext if not pyspark.memory_profiler_ext.has_memory_profiler: return f accumulator = _deserialize_accumulator( SpecialAccumulatorIds.SQL_UDF_PROFIER, None, ProfileResultsParam ) accumulator_v2 = _deserialize_accumulator( SpecialAccumulatorIds.SQL_UDF_PROFIER_V2, {}, ProfileResultsParamV2 ) if _is_iter_based(eval_type): def profiling_func(*args, **kwargs): g = f(*args, **kwargs) iterator = iter(g) while True: try: with WorkerMemoryProfiler(accumulator, accumulator_v2, result_id, g.gi_code): item = next(iterator) yield item except StopIteration: break else: def profiling_func(*args, **kwargs): with WorkerMemoryProfiler(accumulator, accumulator_v2, result_id, f): ret = f(*args, **kwargs) return ret return profiling_func def read_single_udf(pickleSer, udf_info, eval_type, runner_conf, udf_index): chained_func = None for udf in udf_info.udfs: f, return_type = read_command(pickleSer, udf) if chained_func is None: chained_func = f else: chained_func = chain(chained_func, f) # If chained_func is from pyspark.sql.worker, it is to read/write data source. # In this case, we check the data_source_profiler config. module = getattr(chained_func, "__module__", "") if isinstance(module, str) and module.startswith("pyspark.sql.worker."): profiler = runner_conf.data_source_profiler else: profiler = runner_conf.udf_profiler if profiler == "perf": profiling_func = wrap_perf_profiler(chained_func, eval_type, udf_info.result_id) elif profiler == "memory": profiling_func = wrap_memory_profiler(chained_func, eval_type, udf_info.result_id) else: profiling_func = chained_func if eval_type in ( PythonEvalType.SQL_SCALAR_PANDAS_ITER_UDF, PythonEvalType.SQL_SCALAR_PANDAS_ITER_ELEMENTWISE_UDF, PythonEvalType.SQL_ARROW_BATCHED_UDF, ): func = profiling_func else: # make sure StopIteration's raised in the user code are not ignored # when they are processed in a for loop, raise them as RuntimeError's instead func = fail_on_stopiteration(profiling_func) args_offsets, kwargs_offsets = udf_info.args, udf_info.kwargs # The last returnType will be the return type of UDF. Eval types are grouped below by the # shape of the value they return. # Scalar, aggregation and window UDFs: (func, args_offsets, kwargs_offsets, return_type). if eval_type in ( PythonEvalType.SQL_ARROW_BATCHED_UDF, PythonEvalType.SQL_ARROW_ELEMENTWISE_UDF, PythonEvalType.SQL_SCALAR_PANDAS_ELEMENTWISE_UDF, PythonEvalType.SQL_SCALAR_PANDAS_ITER_ELEMENTWISE_UDF, PythonEvalType.SQL_SCALAR_ARROW_ELEMENTWISE_UDF, PythonEvalType.SQL_SCALAR_ARROW_ITER_ELEMENTWISE_UDF, PythonEvalType.SQL_GROUPED_AGG_ARROW_ITER_UDF, PythonEvalType.SQL_GROUPED_AGG_ARROW_UDF, PythonEvalType.SQL_GROUPED_AGG_PANDAS_ITER_UDF, PythonEvalType.SQL_GROUPED_AGG_PANDAS_UDF, PythonEvalType.SQL_SCALAR_ARROW_ITER_UDF, PythonEvalType.SQL_SCALAR_ARROW_UDF, PythonEvalType.SQL_SCALAR_PANDAS_ITER_UDF, PythonEvalType.SQL_SCALAR_PANDAS_UDF, PythonEvalType.SQL_WINDOW_AGG_ARROW_UDF, PythonEvalType.SQL_WINDOW_AGG_PANDAS_UDF, ): return func, args_offsets, kwargs_offsets, return_type # Grouped-map and cogrouped-map UDFs: (func, args_offsets, return_type, num_udf_args). elif eval_type in ( PythonEvalType.SQL_COGROUPED_MAP_ARROW_UDF, PythonEvalType.SQL_COGROUPED_MAP_PANDAS_UDF, PythonEvalType.SQL_GROUPED_MAP_ARROW_ITER_UDF, PythonEvalType.SQL_GROUPED_MAP_ARROW_UDF, PythonEvalType.SQL_GROUPED_MAP_PANDAS_ITER_UDF, PythonEvalType.SQL_GROUPED_MAP_PANDAS_UDF, ): # signature was lost when wrapping it num_udf_args = len(inspect.getfullargspec(chained_func).args) return func, args_offsets, return_type, num_udf_args # Grouped-map-with-state and transform-with-state UDFs: (func, args_offsets, return_type). elif eval_type in ( PythonEvalType.SQL_GROUPED_MAP_PANDAS_UDF_WITH_STATE, PythonEvalType.SQL_TRANSFORM_WITH_STATE_PANDAS_INIT_STATE_UDF, PythonEvalType.SQL_TRANSFORM_WITH_STATE_PANDAS_UDF, PythonEvalType.SQL_TRANSFORM_WITH_STATE_PYTHON_ROW_INIT_STATE_UDF, PythonEvalType.SQL_TRANSFORM_WITH_STATE_PYTHON_ROW_UDF, ): return func, args_offsets, return_type # Map iterator UDFs take no offsets. elif eval_type == PythonEvalType.SQL_MAP_PANDAS_ITER_UDF: return func, None, None, return_type elif eval_type == PythonEvalType.SQL_MAP_ARROW_ITER_UDF: return func, None, None, None # Batched (plain Python) UDFs: (args_kwargs_offsets, eval func); apply kwargs binding and # convert each result to the internal representation only when the return type requires it. elif eval_type == PythonEvalType.SQL_BATCHED_UDF: func, args_kwargs_offsets = wrap_kwargs_support(func, args_offsets, kwargs_offsets) if return_type.needConversion(): toInternal = return_type.toInternal return args_kwargs_offsets, lambda *a: toInternal(func(*a)) else: return args_kwargs_offsets, lambda *a: func(*a) else: raise ValueError("Unknown eval type: {}".format(eval_type)) # Read and process a serialized user-defined table function (UDTF) from a socket. # It expects the UDTF to be in a specific format and performs various checks to # ensure the UDTF is valid. This function also prepares a mapper function for applying # the UDTF logic to input rows. def read_udtf(pickleSer, udtf_info, eval_type, runner_conf, eval_conf): if eval_type in ( # Pure Arrow stream I/O for both the legacy pandas conversion path and the # non-legacy path; the pandas (de)serialization for the legacy path and the # output struct wrapping are both handled in the func below. PythonEvalType.SQL_ARROW_TABLE_UDF, # Pure Arrow stream I/O; table-arg flattening and output coercion # are handled in the func below. PythonEvalType.SQL_ARROW_UDTF, ): ser = ArrowStreamSerializer(write_start_stream=True) else: # Each row is a group so do not batch but send one by one. ser = BatchedSerializer(CPickleSerializer(), 1) if udtf_info.pickled_analyze_result is not None: pickled_analyze_result = pickleSer.loads(udtf_info.pickled_analyze_result) else: pickled_analyze_result = None # Initially we assume that the UDTF __init__ method accepts the pickled AnalyzeResult, # although we may set this to false later if we find otherwise. handler = read_command(pickleSer, udtf_info.handler) if not isinstance(handler, type): raise PySparkRuntimeError( f"Invalid UDTF handler type. Expected a class (type 'type'), but " f"got an instance of {type(handler).__name__}." ) return_type = _parse_datatype_json_string(udtf_info.return_type) if not isinstance(return_type, StructType): raise PySparkRuntimeError( f"The return type of a UDTF must be a struct type, but got {type(return_type)}." ) # Update the handler that creates a new UDTF instance to first try calling the UDTF constructor # with one argument containing the previous AnalyzeResult. If that fails, then try a constructor # with no arguments. In this way each UDTF class instance can decide if it wants to inspect the # AnalyzeResult. udtf_init_args = inspect.getfullargspec(handler) if pickled_analyze_result is not None: if len(udtf_init_args.args) > 2: raise PySparkRuntimeError( errorClass="UDTF_CONSTRUCTOR_INVALID_IMPLEMENTS_ANALYZE_METHOD", messageParameters={"name": udtf_info.name}, ) elif len(udtf_init_args.args) == 2: prev_handler = handler def construct_udtf(): # Here we pass the AnalyzeResult to the UDTF's __init__ method. return prev_handler(dataclasses.replace(pickled_analyze_result)) handler = construct_udtf elif len(udtf_init_args.args) > 1: raise PySparkRuntimeError( errorClass="UDTF_CONSTRUCTOR_INVALID_NO_ANALYZE_METHOD", messageParameters={"name": udtf_info.name}, ) class UDTFWithPartitions: """ This implements the logic of a UDTF that accepts an input TABLE argument with one or more PARTITION BY expressions. For example, let's assume we have a table like: CREATE TABLE t (c1 INT, c2 INT) USING delta; Then for the following queries: SELECT * FROM my_udtf(TABLE (t) PARTITION BY c1, c2); The partition_child_indexes will be: 0, 1. SELECT * FROM my_udtf(TABLE (t) PARTITION BY c1, c2 + 4); The partition_child_indexes will be: 0, 2 (where we add a projection for "c2 + 4"). """ def __init__(self, create_udtf: Callable, partition_child_indexes: list): """ Creates a new instance of this class to wrap the provided UDTF with another one that checks the values of projected partitioning expressions on consecutive rows to figure out when the partition boundaries change. Parameters ---------- create_udtf: function Function to create a new instance of the UDTF to be invoked. partition_child_indexes: list List of integers identifying zero-based indexes of the columns of the input table that contain projected partitioning expressions. This class will inspect these values for each pair of consecutive input rows. When they change, this indicates the boundary between two partitions, and we will invoke the 'terminate' method on the UDTF class instance and then destroy it and create a new one to implement the desired partitioning semantics. """ self._create_udtf: Callable = create_udtf self._udtf = create_udtf() self._prev_arguments: list = list() self._partition_child_indexes: list = udtf_info.partition_child_indexes self._eval_raised_skip_rest_of_input_table: bool = False def eval(self, *args, **kwargs) -> Iterator: changed_partitions = self._check_partition_boundaries( list(args) + list(kwargs.values()) ) if changed_partitions: if hasattr(self._udtf, "terminate"): result = self._udtf.terminate() if result is not None: for row in result: yield row self._udtf = self._create_udtf() self._eval_raised_skip_rest_of_input_table = False if self._udtf.eval is not None and not self._eval_raised_skip_rest_of_input_table: # Filter the arguments to exclude projected PARTITION BY values added by Catalyst. filtered_args = [self._remove_partition_by_exprs(arg) for arg in args] filtered_kwargs = { key: self._remove_partition_by_exprs(value) for (key, value) in kwargs.items() } try: result = self._udtf.eval(*filtered_args, **filtered_kwargs) if result is not None: for row in result: yield row except SkipRestOfInputTableException: # If the 'eval' method raised this exception, then we should skip the rest of # the rows in the current partition. Set this field to True here and then for # each subsequent row in the partition, we will skip calling the 'eval' method # until we see a change in the partition boundaries. self._eval_raised_skip_rest_of_input_table = True def terminate(self) -> Iterator: if hasattr(self._udtf, "terminate"): return self._udtf.terminate() return iter(()) def cleanup(self) -> None: if hasattr(self._udtf, "cleanup"): self._udtf.cleanup() def _check_partition_boundaries(self, arguments: list) -> bool: result = False if len(self._prev_arguments) > 0: cur_table_arg = self._get_table_arg(arguments) prev_table_arg = self._get_table_arg(self._prev_arguments) cur_partitions_args = [] prev_partitions_args = [] for i in self._partition_child_indexes: cur_partitions_args.append(cur_table_arg[i]) prev_partitions_args.append(prev_table_arg[i]) result = any(k != v for k, v in zip(cur_partitions_args, prev_partitions_args)) self._prev_arguments = arguments return result def _get_table_arg(self, inputs: list) -> Row: return [x for x in inputs if type(x) is Row][0] def _remove_partition_by_exprs(self, arg: Any) -> Any: if isinstance(arg, Row): new_row_keys = [] new_row_values = [] for i, (key, value) in enumerate(zip(arg.__fields__, arg)): if i not in self._partition_child_indexes: new_row_keys.append(key) new_row_values.append(value) return _create_row(new_row_keys, new_row_values) else: return arg class ArrowUDTFWithPartition: """ Implements logic for an Arrow UDTF (SQL_ARROW_UDTF) that accepts a TABLE argument with one or more PARTITION BY expressions. Arrow UDTFs receive data as PyArrow RecordBatch objects instead of individual Row objects. This wrapper ensures the UDTF's eval() method is called separately for each unique partition key value combination. How Catalyst handles PARTITION BY and ORDER BY: ------------------------------------------------ When a UDTF is called with PARTITION BY and/or ORDER BY clauses, Catalyst adds operations to the physical plan to ensure correct data organization: Example SQL: SELECT * FROM my_udtf(TABLE(t) PARTITION BY key1, key2 ORDER BY value DESC) Physical Plan generated by Catalyst: 1. Project: Adds partition_by_0 = key1, partition_by_1 = key2 columns 2. Exchange: hashpartitioning(partition_by_0, partition_by_1, 200) - Shuffles data so rows with same partition keys go to same worker 3. Sort: [partition_by_0 ASC, partition_by_1 ASC, value DESC], local=true - First sorts by partition keys to group them together - Then sorts by ORDER BY expressions within each partition - Local sort (not global) within each worker's data 4. Project: Creates struct with all columns including partition_by_* columns 5. ArrowEvalPythonUDTF: Executes this Python UDTF wrapper Key guarantee: After the Sort operation, all rows with the same partition key values are contiguous within each RecordBatch, allowing efficient boundary detection. Example queries: SELECT * FROM my_udtf(TABLE (t) PARTITION BY c1); partition_child_indexes: [2] (refers to partition_by_0 column at index 2) SELECT * FROM my_udtf(TABLE (t) PARTITION BY c1, c2); partition_child_indexes: [2, 3] (partition_by_0 and partition_by_1 columns) SELECT * FROM my_udtf(TABLE (t) PARTITION BY c1, c2 + 4); partition_child_indexes: 0, 2 (adds a projection for "c2 + 4"). """ def __init__(self, create_udtf: Callable, partition_child_indexes: list): """ Create a new instance that wraps the provided Arrow UDTF with partitioning logic. Parameters ---------- create_udtf: function Function that creates a new instance of the Arrow UDTF to invoke. partition_child_indexes: list Zero-based indexes of input-table columns that contain projected partitioning expressions. """ self._create_udtf: Callable = create_udtf self._udtf = create_udtf() self._partition_child_indexes: list = partition_child_indexes # Track last partition key from previous batch self._last_partition_key: Optional[Tuple[Any, ...]] = None self._eval_raised_skip_rest_of_input_table: bool = False def eval(self, *args, **kwargs) -> Iterator: """Handle partitioning logic for Arrow UDTFs that receive RecordBatch objects.""" import pyarrow as pa # Get the original batch with partition columns original_batch = self._get_table_arg(list(args) + list(kwargs.values())) if not isinstance(original_batch, pa.RecordBatch): # Arrow UDTFs with PARTITION BY must have a TABLE argument that # results in a PyArrow RecordBatch raise PySparkRuntimeError( errorClass="INVALID_ARROW_UDTF_TABLE_ARGUMENT", messageParameters={ "actual_type": ( str(type(original_batch)) if original_batch is not None else "None" ) }, ) # Remove partition columns to get the filtered arguments filtered_args = [self._remove_partition_by_exprs(arg) for arg in args] filtered_kwargs = { key: self._remove_partition_by_exprs(value) for (key, value) in kwargs.items() } # Get the filtered RecordBatch (without partition columns) filtered_batch = self._get_table_arg(filtered_args + list(filtered_kwargs.values())) # Process the RecordBatch by partitions yield from self._process_arrow_batch_by_partitions( original_batch, filtered_batch, filtered_args, filtered_kwargs ) def _process_arrow_batch_by_partitions( self, original_batch, filtered_batch, filtered_args, filtered_kwargs ) -> Iterator: """Process an Arrow RecordBatch that may contain multiple partition key values. When using PARTITION BY with Arrow UDTFs, a single RecordBatch from Spark may contain rows with different partition key values. For example, with 10 distinct partition keys and 2 workers, each worker might receive a batch containing 5 different partition key values. According to UDTF PARTITION BY semantics, the UDTF's eval() method must be called separately for each unique partition key value, not for the entire batch. This method handles splitting the batch by partition boundaries and calling the UDTF appropriately. The implementation leverages two key properties: 1. Catalyst guarantees rows with the same partition key are contiguous (pre-sorted) 2. Arrow's columnar format allows efficient boundary detection Parameters: ----------- original_batch : pa.RecordBatch The original batch including partition columns, used for detecting boundaries filtered_batch : pa.RecordBatch The batch with partition columns removed, to be passed to the UDTF filtered_args : list Arguments with partition columns filtered out filtered_kwargs : dict Keyword arguments with partition columns filtered out Yields: ------- Iterator of pa.Table objects returned by the UDTF's eval() method """ import pyarrow as pa # This class should only be used when partition_child_indexes is non-empty assert self._partition_child_indexes, ( "ArrowUDTFWithPartition should only be instantiated when " "len(partition_child_indexes) > 0" ) # Detect partition boundaries. boundaries = self._detect_partition_boundaries(original_batch) # Process each contiguous partition for i in range(len(boundaries) - 1): start_idx = boundaries[i] end_idx = boundaries[i + 1] # Get the partition key for this segment partition_key = tuple( original_batch.column(idx)[start_idx].as_py() for idx in self._partition_child_indexes ) # Check if this is a continuation of the previous batch's partition # TODO: This check is only necessary for the first boundary in each batch. # The following boundaries are always for new partitions within the same batch. # This could be optimized by only checking i == 0. is_new_partition = ( self._last_partition_key is not None and partition_key != self._last_partition_key ) if is_new_partition: # Previous partition ended, call terminate if hasattr(self._udtf, "terminate"): terminate_result = self._udtf.terminate() if terminate_result is not None: yield from terminate_result # Create new UDTF instance for new partition self._udtf = self._create_udtf() self._eval_raised_skip_rest_of_input_table = False # Slice the filtered batch for this partition partition_batch = filtered_batch.slice(start_idx, end_idx - start_idx) # Update the last partition key self._last_partition_key = partition_key # Update filtered args to use the partition batch partition_filtered_args = [] for arg in filtered_args: if isinstance(arg, pa.RecordBatch): partition_filtered_args.append(partition_batch) else: partition_filtered_args.append(arg) partition_filtered_kwargs = {} for key, value in filtered_kwargs.items(): if isinstance(value, pa.RecordBatch): partition_filtered_kwargs[key] = partition_batch else: partition_filtered_kwargs[key] = value # Call the UDTF with this partition's data if not self._eval_raised_skip_rest_of_input_table: try: result = self._udtf.eval( *partition_filtered_args, **partition_filtered_kwargs ) if result is not None: yield from result except SkipRestOfInputTableException: # Skip remaining rows in this partition self._eval_raised_skip_rest_of_input_table = True # Don't terminate here - let the next batch or final terminate handle it def terminate(self) -> Iterator: if hasattr(self._udtf, "terminate"): return self._udtf.terminate() return iter(()) def cleanup(self) -> None: if hasattr(self._udtf, "cleanup"): self._udtf.cleanup() def _get_table_arg(self, inputs: list): """Get the table argument (RecordBatch) from the inputs list. For Arrow UDTFs with TABLE arguments, we can guarantee the table argument will be a pa.RecordBatch, not a Row. """ import pyarrow as pa # Find all RecordBatch arguments batches = [arg for arg in inputs if isinstance(arg, pa.RecordBatch)] if len(batches) == 0: # No RecordBatch found - this shouldn't happen for Arrow UDTFs with TABLE arguments return None elif len(batches) == 1: return batches[0] else: # Multiple RecordBatch arguments found - this is unexpected raise RuntimeError( f"Expected exactly one pa.RecordBatch argument for TABLE parameter, " f"but found {len(batches)}. Received types: " f"{[type(arg).__name__ for arg in inputs]}" ) def _detect_partition_boundaries(self, batch) -> list: """ Efficiently detect partition boundaries in a batch with contiguous partitions. Since Catalyst ensures rows with the same partition key are contiguous, we only need to find where partition values change. Returns: List of indices where each partition starts, plus the total row count. For example: [0, 3, 8, 10] means partitions are rows [0:3), [3:8), [8:10) """ boundaries = [0] # First partition starts at index 0 if batch.num_rows <= 1: boundaries.append(batch.num_rows) return boundaries # Get partition column arrays partition_arrays = [batch.column(i) for i in self._partition_child_indexes] # Find boundaries by comparing consecutive rows for row_idx in range(1, batch.num_rows): # Check if any partition column changed from previous row partition_changed = False for col_array in partition_arrays: if col_array[row_idx].as_py() != col_array[row_idx - 1].as_py(): partition_changed = True break if partition_changed: boundaries.append(row_idx) boundaries.append(batch.num_rows) # Last boundary at end return boundaries def _remove_partition_by_exprs(self, arg: Any) -> Any: """ Remove partition columns from the RecordBatch argument. Why this is needed: When a UDTF is called with TABLE(t) PARTITION BY expressions, Catalyst transforms the data: 1. Adds complex partition expressions as new columns (e.g., "c2 + 4" becomes a new column) 2. Repartitions data by partition columns using hash partitioning 3. Sends ALL columns (including partition columns) to the Python worker Partition columns serve two purposes: - Routing: decide which worker processes which partition - Boundary detection: know when one partition ends and another begins However, the user's UDTF should only receive the actual table data, not the partition columns. This method filters out partition columns before passing data to the user's UDTF eval() method. Example: - User writes: SELECT * FROM udtf(TABLE(t) PARTITION BY c1, c2) - Catalyst sends: RecordBatch with [c1, c2, c3, c4], partition_child_indexes=[0, 1] - This method removes columns at indexes 0, 1 if they are pure partition columns - UDTF.eval() receives: RecordBatch with only the non-partition columns """ import pyarrow as pa if isinstance(arg, pa.RecordBatch): # Remove partition columns from the RecordBatch keep_indices = [ i for i in range(len(arg.schema.names)) if i not in self._partition_child_indexes ] if keep_indices: # Select only the columns we want to keep keep_arrays = [arg.column(i) for i in keep_indices] keep_names = [arg.schema.names[i] for i in keep_indices] return pa.RecordBatch.from_arrays(keep_arrays, names=keep_names) else: # If no columns remain, return an empty RecordBatch with the same number of rows return pa.RecordBatch.from_arrays( [], schema=pa.schema([]), num_rows=arg.num_rows ) # For non-RecordBatch arguments (like scalar pa.Arrays), return unchanged return arg # Instantiate the UDTF class. try: if len(udtf_info.partition_child_indexes) > 0: # Determine if this is an Arrow UDTF is_arrow_udtf = eval_type == PythonEvalType.SQL_ARROW_UDTF if is_arrow_udtf: udtf = ArrowUDTFWithPartition(handler, udtf_info.partition_child_indexes) else: udtf = UDTFWithPartitions(handler, udtf_info.partition_child_indexes) else: udtf = handler() except Exception as e: raise PySparkRuntimeError( errorClass="UDTF_EXEC_ERROR", messageParameters={"method_name": "__init__", "error": str(e)}, ) # Validate the UDTF if not hasattr(udtf, "eval"): raise PySparkRuntimeError( "Failed to execute the user defined table function because it has not " "implemented the 'eval' method. Please add the 'eval' method and try " "the query again." ) # Check that the arguments provided to the UDTF call match the expected parameters defined # in the 'eval' method signature. try: inspect.signature(udtf.eval).bind(*udtf_info.args, **udtf_info.kwargs) except TypeError as e: raise PySparkRuntimeError( errorClass="UDTF_EVAL_METHOD_ARGUMENTS_DO_NOT_MATCH_SIGNATURE", messageParameters={"name": udtf_info.name, "reason": str(e)}, ) from None def build_null_checker(return_type: StructType) -> Optional[Callable[[Any], None]]: def raise_(result_column_index): raise PySparkRuntimeError( errorClass="UDTF_EXEC_ERROR", messageParameters={ "method_name": "eval' or 'terminate", "error": f"Column {result_column_index} within a returned row had a " + "value of None, either directly or within array/struct/map " + "subfields, but the corresponding column type was declared as " + "non-nullable; please update the UDTF to return a non-None value at " + "this location or otherwise declare the column type as nullable.", }, ) def checker(data_type: DataType, result_column_index: int): if isinstance(data_type, ArrayType): element_checker = checker(data_type.elementType, result_column_index) contains_null = data_type.containsNull if element_checker is None and contains_null: return None def check_array(arr): if isinstance(arr, list): for e in arr: if e is None: if not contains_null: raise_(result_column_index) elif element_checker is not None: element_checker(e) return check_array elif isinstance(data_type, MapType): key_checker = checker(data_type.keyType, result_column_index) value_checker = checker(data_type.valueType, result_column_index) value_contains_null = data_type.valueContainsNull if value_checker is None and value_contains_null: def check_map(map): if isinstance(map, dict): for k, v in map.items(): if k is None: raise_(result_column_index) elif key_checker is not None: key_checker(k) else: def check_map(map): if isinstance(map, dict): for k, v in map.items(): if k is None: raise_(result_column_index) elif key_checker is not None: key_checker(k) if v is None: if not value_contains_null: raise_(result_column_index) elif value_checker is not None: value_checker(v) return check_map elif isinstance(data_type, StructType): field_checkers = [checker(f.dataType, result_column_index) for f in data_type] nullables = [f.nullable for f in data_type] if all(c is None for c in field_checkers) and all(nullables): return None def check_struct(struct): if isinstance(struct, tuple): for value, checker, nullable in zip(struct, field_checkers, nullables): if value is None: if not nullable: raise_(result_column_index) elif checker is not None: checker(value) return check_struct else: return None field_checkers = [ checker(f.dataType, result_column_index=i) for i, f in enumerate(return_type) ] nullables = [f.nullable for f in return_type] if all(c is None for c in field_checkers) and all(nullables): return None def check(row): if isinstance(row, tuple): for i, (value, checker, nullable) in enumerate(zip(row, field_checkers, nullables)): if value is None: if not nullable: raise_(i) elif checker is not None: checker(value) return check check_output_row_against_schema = build_null_checker(return_type) if ( eval_type == PythonEvalType.SQL_ARROW_TABLE_UDF and runner_conf.use_legacy_pandas_udtf_conversion ): import pandas as pd return_type_size = len(return_type) # The output pandas DataFrame is converted as a single struct column named # "_0" against this schema. output_schema = StructType([StructField("_0", return_type)]) def verify_result(result: Any, method_name: str) -> Any: if not isinstance(result, pd.DataFrame): raise PySparkTypeError( errorClass="INVALID_ARROW_UDTF_RETURN_TYPE", messageParameters={ "return_type": type(result).__name__, "value": str(result), "func": method_name, }, ) # Validate the output schema when the result dataframe has either output # rows or columns. Note that we avoid using `df.empty` here because the # result dataframe may contain an empty row. For example, when a UDTF is # defined as follows: def eval(self): yield tuple(). if len(result) > 0 or len(result.columns) > 0: if len(result.columns) != return_type_size: raise PySparkRuntimeError( errorClass="UDTF_RETURN_SCHEMA_MISMATCH", messageParameters={ "expected": str(return_type_size), "actual": str(len(result.columns)), "func": method_name, }, ) # Verify the type and the schema of the result. verify_pandas_result( result, return_type, assign_cols_by_name=False, truncate_return_schema=False ) return result def check_return_value(res: Any, method_name: str) -> Iterator: # Check whether the result of an arrow UDTF is iterable before # using it to construct a pandas DataFrame. if res is not None: if not isinstance(res, Iterable): raise PySparkRuntimeError( errorClass="UDTF_RETURN_NOT_ITERABLE", messageParameters={ "type": type(res).__name__, "func": method_name, }, ) if check_output_row_against_schema is not None: for row in res: if row is not None: check_output_row_against_schema(row) yield row else: yield from res def convert_df_to_arrow(result: "pd.DataFrame") -> "pa.RecordBatch": # Convert the output pandas DataFrame into a single "_0" struct column, # applying the legacy pandas-to-Arrow coercions. return PandasToArrowConversion.convert( [result], output_schema, timezone=runner_conf.timezone, safecheck=runner_conf.safecheck, arrow_cast=True, assign_cols_by_name=False, int_to_decimal_coercion_enabled=runner_conf.int_to_decimal_coercion_enabled, ignore_unexpected_complex_type_values=True, is_legacy=True, ) def evaluate_rows( method: Callable, *args: list, num_rows: int = 1 ) -> Iterator["pa.RecordBatch"]: # Create tuples from the input pandas Series, each tuple represents a row # across all Series. rows = itertools.repeat((), num_rows) if len(args) == 0 else zip(*args) for row in rows: # Wrap the exception thrown from the UDTF in a PySparkRuntimeError. try: res = method(*row) except SkipRestOfInputTableException: raise except Exception as e: raise PySparkRuntimeError( errorClass="UDTF_EXEC_ERROR", messageParameters={"method_name": method.__name__, "error": str(e)}, ) result = verify_result( pd.DataFrame(list(check_return_value(res, method.__name__))), method.__name__ ) yield convert_df_to_arrow(result) eval_method, args_kwargs_offsets = wrap_kwargs_support( getattr(udtf, "eval"), udtf_info.args, udtf_info.kwargs ) terminate = getattr(udtf, "terminate", None) cleanup = getattr(udtf, "cleanup", None) def func(split_index: int, data: Iterator["pa.RecordBatch"]) -> Iterator["pa.RecordBatch"]: """Apply legacy pandas Arrow table UDF""" try: for batch in data: # Deserialize the Arrow batch into a list of pandas Series (one per # input column), then call eval once per input row. series_list = ArrowBatchTransformer.to_pandas( batch, timezone=runner_conf.timezone, schema=eval_conf.input_type, struct_in_pandas="row", ndarray_as_list=True, prefer_int_ext_dtype=runner_conf.prefer_int_ext_dtype, df_for_struct=False, ) yield from evaluate_rows( eval_method, *[series_list[o] for o in args_kwargs_offsets], num_rows=batch.num_rows, ) if terminate is not None: yield from evaluate_rows(terminate) except SkipRestOfInputTableException: if terminate is not None: yield from evaluate_rows(terminate) finally: if cleanup is not None: cleanup() return func, None, ser, ser elif ( eval_type == PythonEvalType.SQL_ARROW_TABLE_UDF and not runner_conf.use_legacy_pandas_udtf_conversion ): import pyarrow as pa arrow_return_type = to_arrow_type( return_type, timezone="UTC", prefers_large_types=runner_conf.use_large_var_types ) return_type_size = len(return_type) def verify_result(result: pa.Table, method_name: str) -> pa.Table: if not isinstance(result, pa.Table): raise PySparkTypeError( errorClass="INVALID_ARROW_UDTF_RETURN_TYPE", messageParameters={ "return_type": type(result).__name__, "value": str(result), "func": method_name, }, ) # Validate the output schema when the result dataframe has either output # rows or columns. Note that we avoid using `df.empty` here because the # result dataframe may contain an empty row. For example, when a UDTF is # defined as follows: def eval(self): yield tuple(). if result.num_rows > 0 or result.num_columns > 0: if result.num_columns != return_type_size: raise PySparkRuntimeError( errorClass="UDTF_RETURN_SCHEMA_MISMATCH", messageParameters={ "expected": str(return_type_size), "actual": str(result.num_columns), "func": method_name, }, ) # Verify the type and the schema of the result. verify_arrow_result( result, assign_cols_by_name=False, expected_cols_and_types=[(field.name, field.type) for field in arrow_return_type], ) return result def check_return_value(res: Any, method_name: str) -> Iterator: # Check whether the result of an arrow UDTF is iterable before # using it to construct a pandas DataFrame. if res is not None: if not isinstance(res, Iterable): raise PySparkRuntimeError( errorClass="UDTF_RETURN_NOT_ITERABLE", messageParameters={ "type": type(res).__name__, "func": method_name, }, ) for row in res: if not isinstance(row, tuple) and return_type_size == 1: row = (row,) if check_output_row_against_schema is not None: if row is not None: check_output_row_against_schema(row) yield row def convert_rows_to_arrow(data: Iterable, method_name: str) -> list[pa.RecordBatch]: data = list(check_return_value(data, method_name)) if len(data) == 0: # Return one empty RecordBatch to match the left side of the lateral join return [pa.RecordBatch.from_pylist(data, schema=pa.schema(list(arrow_return_type)))] def raise_conversion_error(original_exception): raise PySparkRuntimeError( errorClass="UDTF_ARROW_DATA_CONVERSION_ERROR", messageParameters={ "data": str(data), "schema": return_type.simpleString(), "arrow_schema": str(arrow_return_type), }, ) from original_exception try: table = LocalDataToArrowConversion.convert( data, return_type, runner_conf.use_large_var_types ) except PySparkValueError as e: if e.getErrorClass() == "AXIS_LENGTH_MISMATCH": raise PySparkRuntimeError( errorClass="UDTF_RETURN_SCHEMA_MISMATCH", messageParameters={ "expected": e.getMessageParameters()["expected_length"], # type: ignore[index] "actual": e.getMessageParameters()["actual_length"], # type: ignore[index] "func": method_name, }, ) from e # Fall through to general conversion error raise_conversion_error(e) except Exception as e: raise_conversion_error(e) return verify_result(table, method_name).to_batches() def evaluate_rows( method: Callable, *args: list, num_rows: int = 1 ) -> Iterator[pa.RecordBatch]: rows = itertools.repeat((), num_rows) if len(args) == 0 else zip(*args) for row in rows: # Wrap the exception thrown from the UDTF in a PySparkRuntimeError. try: res = method(*row) except SkipRestOfInputTableException: raise except Exception as e: raise PySparkRuntimeError( errorClass="UDTF_EXEC_ERROR", messageParameters={"method_name": method.__name__, "error": str(e)}, ) for batch in convert_rows_to_arrow(res, method.__name__): yield ArrowBatchTransformer.wrap_struct(batch) eval_method, args_kwargs_offsets = wrap_kwargs_support( getattr(udtf, "eval"), udtf_info.args, udtf_info.kwargs ) terminate = getattr(udtf, "terminate", None) cleanup = getattr(udtf, "cleanup", None) def func(split_index: int, data: Iterator[pa.RecordBatch]) -> Iterator[pa.RecordBatch]: """Apply Arrow table UDF""" try: converters = [ ArrowTableToRowsConversion._create_converter( f.dataType, none_on_identity=True, binary_as_bytes=runner_conf.binary_as_bytes, ) for f in eval_conf.input_type ] for batch in data: # Convert each input column to a list of Python values per row, # then call eval once per input row. pylist = [ ( [conv(v) for v in ArrowTableToRowsConversion._to_pylist(column)] if conv is not None else ArrowTableToRowsConversion._to_pylist(column) ) for column, conv in zip(batch.columns, converters) ] yield from evaluate_rows( eval_method, *[pylist[o] for o in args_kwargs_offsets], num_rows=batch.num_rows, ) if terminate is not None: yield from evaluate_rows(terminate) except SkipRestOfInputTableException: if terminate is not None: yield from evaluate_rows(terminate) finally: if cleanup is not None: cleanup() return func, None, ser, ser elif eval_type == PythonEvalType.SQL_ARROW_UDTF: import pyarrow as pa arrow_return_type = to_arrow_type( return_type, timezone="UTC", prefers_large_types=runner_conf.use_large_var_types ) return_type_size = len(return_type) target_schema = pa.schema(list(arrow_return_type)) def verify_result(result: pa.RecordBatch, method_name: str) -> pa.RecordBatch: # Validate the output schema when the result has columns if result.num_columns != return_type_size: raise PySparkRuntimeError( errorClass="UDTF_RETURN_SCHEMA_MISMATCH", messageParameters={ "expected": str(return_type_size), "actual": str(result.num_columns), "func": method_name, }, ) return result def convert_to_arrow(res: Any, method_name: str) -> Iterator[pa.RecordBatch]: # Check whether the result of a PyArrow UDTF is iterable before processing if res is None: res = iter([]) elif not isinstance(res, Iterable): raise PySparkRuntimeError( errorClass="UDTF_RETURN_NOT_ITERABLE", messageParameters={ "type": type(res).__name__, "func": method_name, }, ) # Handle PyArrow Tables/RecordBatches directly is_empty = True for item in res: is_empty = False if isinstance(item, pa.Table): yield from item.to_batches() elif isinstance(item, pa.RecordBatch): yield item else: # Arrow UDTF should only return Arrow types (RecordBatch/Table) raise PySparkRuntimeError( errorClass="UDTF_ARROW_TYPE_CONVERSION_ERROR", messageParameters={}, ) if is_empty: yield pa.RecordBatch.from_pylist([], schema=target_schema) def evaluate(method: Callable, *args: pa.RecordBatch) -> Iterator[pa.RecordBatch]: # Wrap the exception thrown from the UDTF in a PySparkRuntimeError. try: res = method(*args) except SkipRestOfInputTableException: raise except Exception as e: raise PySparkRuntimeError( errorClass="UDTF_EXEC_ERROR", messageParameters={"method_name": method.__name__, "error": str(e)}, ) for batch in convert_to_arrow(res, method.__name__): coerced = ArrowBatchTransformer.enforce_schema( verify_result(batch, method.__name__), target_schema, safecheck=True ) yield ArrowBatchTransformer.wrap_struct(coerced) eval_method, args_kwargs_offsets = wrap_kwargs_support( getattr(udtf, "eval"), udtf_info.args, udtf_info.kwargs ) terminate = getattr(udtf, "terminate", None) cleanup = getattr(udtf, "cleanup", None) table_arg_offsets = ( set(eval_conf.table_arg_offsets) if eval_conf.table_arg_offsets else set() ) def func(split_index: int, data: Iterator[pa.RecordBatch]) -> Iterator[pa.RecordBatch]: """Apply Arrow UDTF""" try: for batch in data: # Pre-processing: for each column, flatten struct columns at # table_arg_offsets into RecordBatch, keep other columns as Array. columns = [ ( ArrowBatchTransformer.flatten_struct(batch, column_index=i) if i in table_arg_offsets else batch.column(i) ) for i in range(batch.num_columns) ] # For PyArrow UDTFs, pass RecordBatches directly (no row conversion needed) yield from evaluate(eval_method, *[columns[o] for o in args_kwargs_offsets]) if terminate is not None: yield from evaluate(terminate) except SkipRestOfInputTableException: if terminate is not None: yield from evaluate(terminate) finally: if cleanup is not None: cleanup() return func, None, ser, ser else: def wrap_udtf(f, return_type): assert return_type.needConversion() toInternal = return_type.toInternal return_type_size = len(return_type) def verify_and_convert_result(result): if result is not None: if hasattr(result, "__UDT__"): # UDT object should not be returned directly. raise PySparkRuntimeError( errorClass="UDTF_INVALID_OUTPUT_ROW_TYPE", messageParameters={ "type": type(result).__name__, "func": f.__name__, }, ) if hasattr(result, "__len__") and len(result) != return_type_size: raise PySparkRuntimeError( errorClass="UDTF_RETURN_SCHEMA_MISMATCH", messageParameters={ "expected": str(return_type_size), "actual": str(len(result)), "func": f.__name__, }, ) if not (isinstance(result, (list, dict, tuple)) or hasattr(result, "__dict__")): raise PySparkRuntimeError( errorClass="UDTF_INVALID_OUTPUT_ROW_TYPE", messageParameters={ "type": type(result).__name__, "func": f.__name__, }, ) if check_output_row_against_schema is not None: check_output_row_against_schema(result) return toInternal(result) # Evaluate the function and return a tuple back to the executor. def evaluate(*a) -> tuple: try: res = f(*a) except SkipRestOfInputTableException: raise except Exception as e: raise PySparkRuntimeError( errorClass="UDTF_EXEC_ERROR", messageParameters={"method_name": f.__name__, "error": str(e)}, ) if res is None: # If the function returns None or does not have an explicit return statement, # an empty tuple is returned to the executor. # This is because directly constructing tuple(None) results in an exception. return tuple() if not isinstance(res, Iterable): raise PySparkRuntimeError( errorClass="UDTF_RETURN_NOT_ITERABLE", messageParameters={ "type": type(res).__name__, "func": f.__name__, }, ) # If the function returns a result, we map it to the internal representation and # returns the results as a tuple. return tuple(map(verify_and_convert_result, res)) return evaluate eval_func_kwargs_support, args_kwargs_offsets = wrap_kwargs_support( getattr(udtf, "eval"), udtf_info.args, udtf_info.kwargs ) eval = wrap_udtf(eval_func_kwargs_support, return_type) if hasattr(udtf, "terminate"): terminate = wrap_udtf(getattr(udtf, "terminate"), return_type) else: terminate = None cleanup = getattr(udtf, "cleanup") if hasattr(udtf, "cleanup") else None # Return an iterator of iterators. def mapper(_, it): try: for a in it: yield eval(*[a[o] for o in args_kwargs_offsets]) if terminate is not None: yield terminate() except SkipRestOfInputTableException: if terminate is not None: yield terminate() finally: if cleanup is not None: cleanup() return mapper, None, ser, ser def _elementwise_renest(flat_values, shape_lengths, is_large): """Re-nest a flat Array of per-element results into an ``array<R>`` column. ``flat_values`` holds the results for every non-null element in order; ``shape_lengths`` is the per-array element count of the iterated argument (``None`` for a null array, which stays null and consumes no elements). ``is_large`` preserves the input's list width (``ListArray`` with int32 offsets vs. ``LargeListArray`` with int64). Shared by the vectorized element-wise worker paths (scalar pandas / Arrow and their iterator variants) that back Python UDFs inside higher-order function lambdas. See ``ExtractPythonUDFFromLambda``. """ import pyarrow as pa offsets = [0] running = 0 mask = [] for n in shape_lengths: mask.append(n is None) if n is not None: running += n offsets.append(running) list_cls = pa.LargeListArray if is_large else pa.ListArray offsets_arr = pa.array(offsets, type=pa.int64() if is_large else pa.int32()) null_mask = pa.array(mask, type=pa.bool_()) return list_cls.from_arrays(offsets_arr, flat_values, mask=null_mask) def _elementwise_flatten_column(flat, element_type, is_pandas, runner_conf): """Adapt one already-flattened ``array<T>`` element column to the vectorized fn's input. ``flat`` is the flattened element ``pa.Array`` (the caller flattens once per batch and shares it across fused UDFs). Returns it unchanged for the Arrow flavor, or converted to a pandas Series / DataFrame with the element type ``T`` for the pandas flavor. Shared by the vectorized element-wise worker paths that back Python UDFs inside higher-order function lambdas. See ``ExtractPythonUDFFromLambda``. """ if not is_pandas: return flat from pyspark.sql.conversion import ArrowArrayToPandasConversion return ArrowArrayToPandasConversion.convert( flat, element_type, timezone=runner_conf.timezone, struct_in_pandas="dict", ndarray_as_list=False, prefer_int_ext_dtype=runner_conf.prefer_int_ext_dtype, df_for_struct=True, ) def _elementwise_result_to_arrow(result, return_type, arrow_element_type, is_pandas, runner_conf): """Convert one vectorized UDF result over the flat elements to a single flat Arrow Array. ``result`` is a pandas Series / DataFrame (pandas flavor) or a ``pa.Array`` (Arrow flavor); the returned array holds one element per input element. The Arrow flavor is coerced to ``arrow_element_type`` (UTC-typed); the pandas flavor is typed by ``PandasToArrowConversion`` using the session timezone, so its timestamp type may differ from ``arrow_element_type`` - callers that concatenate results must take the type from the returned array, not assume UTC. Shared by the vectorized element-wise worker paths. See ``ExtractPythonUDFFromLambda``. """ import pyarrow as pa if is_pandas: batch = PandasToArrowConversion.convert( [result], StructType([StructField("_0", return_type)]), timezone=runner_conf.timezone, safecheck=runner_conf.safecheck, arrow_cast=True, prefers_large_types=runner_conf.use_large_var_types, assign_cols_by_name=runner_conf.assign_cols_by_name, int_to_decimal_coercion_enabled=runner_conf.int_to_decimal_coercion_enabled, ) else: batch = ArrowBatchTransformer.enforce_schema( pa.RecordBatch.from_arrays([result], ["_0"]), pa.schema([pa.field("_0", arrow_element_type)]), safecheck=runner_conf.safecheck, ) # PandasToArrowConversion / enforce_schema both return a pa.RecordBatch, so column(0) is a # single pa.Array (never a ChunkedArray). return batch.column(0) def read_udfs(pickleSer, udf_info_list, eval_type, runner_conf, eval_conf): if eval_type in ( PythonEvalType.SQL_ARROW_BATCHED_UDF, PythonEvalType.SQL_ARROW_ELEMENTWISE_UDF, PythonEvalType.SQL_SCALAR_PANDAS_ELEMENTWISE_UDF, PythonEvalType.SQL_SCALAR_PANDAS_ITER_ELEMENTWISE_UDF, PythonEvalType.SQL_SCALAR_ARROW_ELEMENTWISE_UDF, PythonEvalType.SQL_SCALAR_ARROW_ITER_ELEMENTWISE_UDF, PythonEvalType.SQL_SCALAR_PANDAS_UDF, PythonEvalType.SQL_SCALAR_ARROW_UDF, PythonEvalType.SQL_COGROUPED_MAP_PANDAS_UDF, PythonEvalType.SQL_SCALAR_PANDAS_ITER_UDF, PythonEvalType.SQL_SCALAR_ARROW_ITER_UDF, PythonEvalType.SQL_MAP_PANDAS_ITER_UDF, PythonEvalType.SQL_MAP_ARROW_ITER_UDF, PythonEvalType.SQL_GROUPED_MAP_PANDAS_UDF, PythonEvalType.SQL_GROUPED_MAP_PANDAS_ITER_UDF, PythonEvalType.SQL_GROUPED_AGG_PANDAS_UDF, PythonEvalType.SQL_GROUPED_AGG_ARROW_UDF, PythonEvalType.SQL_GROUPED_AGG_ARROW_ITER_UDF, PythonEvalType.SQL_WINDOW_AGG_PANDAS_UDF, PythonEvalType.SQL_GROUPED_AGG_PANDAS_ITER_UDF, PythonEvalType.SQL_WINDOW_AGG_ARROW_UDF, PythonEvalType.SQL_GROUPED_MAP_PANDAS_UDF_WITH_STATE, PythonEvalType.SQL_GROUPED_MAP_ARROW_UDF, PythonEvalType.SQL_GROUPED_MAP_ARROW_ITER_UDF, PythonEvalType.SQL_COGROUPED_MAP_ARROW_UDF, PythonEvalType.SQL_TRANSFORM_WITH_STATE_PANDAS_UDF, 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, ): # NOTE: if timezone is set here, that implies respectSessionTimeZone is True if eval_type in ( PythonEvalType.SQL_GROUPED_AGG_ARROW_ITER_UDF, PythonEvalType.SQL_GROUPED_AGG_ARROW_UDF, PythonEvalType.SQL_GROUPED_AGG_PANDAS_ITER_UDF, PythonEvalType.SQL_GROUPED_AGG_PANDAS_UDF, PythonEvalType.SQL_GROUPED_MAP_ARROW_ITER_UDF, PythonEvalType.SQL_GROUPED_MAP_ARROW_UDF, PythonEvalType.SQL_GROUPED_MAP_PANDAS_ITER_UDF, PythonEvalType.SQL_GROUPED_MAP_PANDAS_UDF, PythonEvalType.SQL_WINDOW_AGG_ARROW_UDF, PythonEvalType.SQL_WINDOW_AGG_PANDAS_UDF, ): ser = ArrowStreamGroupSerializer(write_start_stream=True) elif eval_type in ( PythonEvalType.SQL_COGROUPED_MAP_ARROW_UDF, PythonEvalType.SQL_COGROUPED_MAP_PANDAS_UDF, ): ser = ArrowStreamCoGroupSerializer(write_start_stream=True) else: ser = ArrowStreamSerializer(write_start_stream=True) else: batch_size = int(os.environ.get("PYTHON_UDF_BATCH_SIZE", "100")) ser = BatchedSerializer(CPickleSerializer(), batch_size) udfs = [ read_single_udf(pickleSer, udf_info, eval_type, runner_conf, udf_index=udf_index) for udf_index, udf_info in enumerate(udf_info_list) ] num_udfs = len(udfs) def extract_key_value_indexes(grouped_arg_offsets): """ Helper function to extract the key and value indexes from arg_offsets for the grouped and cogrouped pandas udfs. See BasePandasGroupExec.resolveArgOffsets for equivalent scala code. Parameters ---------- grouped_arg_offsets: list List containing the key and value indexes of columns of the DataFrames to be passed to the udf. It consists of n repeating groups where n is the number of DataFrames. Each group has the following format: group[0]: length of group group[1]: length of key indexes group[2.. group[1] +2]: key attributes group[group[1] +3 group[0]]: value attributes """ parsed = [] idx = 0 while idx < len(grouped_arg_offsets): offsets_len = grouped_arg_offsets[idx] idx += 1 offsets = grouped_arg_offsets[idx : idx + offsets_len] split_index = offsets[0] + 1 offset_keys = offsets[1:split_index] offset_values = offsets[split_index:] parsed.append([offset_keys, offset_values]) idx += offsets_len return parsed if eval_type == PythonEvalType.SQL_MAP_ARROW_ITER_UDF: import pyarrow as pa assert num_udfs == 1, "One MAP_ARROW_ITER UDF expected here." udf_func: Callable[[Iterator[pa.RecordBatch]], Iterator[pa.RecordBatch]] = udfs[0][0] def func(split_index: int, data: Iterator[pa.RecordBatch]) -> Iterator[pa.RecordBatch]: """Apply mapInArrow UDF""" # Pre-processing input_batches: Iterator[pa.RecordBatch] = map( ArrowBatchTransformer.flatten_struct, data ) # invoke the UDF output_batches = udf_func(input_batches) # Post-processing verified_iter = verify_return_type( output_batches, Iterator[pa.RecordBatch], # type: ignore[type-abstract] ) yield from map(ArrowBatchTransformer.wrap_struct, verified_iter) # profiling is not supported for UDF return func, None, ser, ser if eval_type == PythonEvalType.SQL_SCALAR_ARROW_UDF: import pyarrow as pa col_names = ["_%d" % i for i in range(len(udfs))] combined_arrow_schema = to_arrow_schema( StructType([StructField(n, rt) for n, (_, _, _, rt) in zip(col_names, udfs)]), timezone="UTC", prefers_large_types=runner_conf.use_large_var_types, ) def func(split_index: int, data: Iterator[pa.RecordBatch]) -> Iterator[pa.RecordBatch]: """Apply scalar Arrow UDFs""" for batch in data: output_batch = pa.RecordBatch.from_arrays( [ udf_func( *[batch.column(o) for o in args_offsets], **{k: batch.column(v) for k, v in kwargs_offsets.items()}, ) for udf_func, args_offsets, kwargs_offsets, _ in udfs ], col_names, ) output_batch = ArrowBatchTransformer.enforce_schema( output_batch, combined_arrow_schema ) verify_scalar_result(output_batch, batch.num_rows) yield output_batch # profiling is not supported for UDF return func, None, ser, ser if eval_type == PythonEvalType.SQL_SCALAR_ARROW_ITER_UDF: import pyarrow as pa assert num_udfs == 1, "One SCALAR_ARROW_ITER UDF expected here." udf_func, args_offsets, kwargs_offsets, return_type = udfs[0] # Pre-compute target Arrow type for output coercion arrow_return_type = to_arrow_type( return_type, timezone="UTC", prefers_large_types=runner_conf.use_large_var_types ) def func(split_index: int, data: Iterator[pa.RecordBatch]) -> Iterator[pa.RecordBatch]: """Apply scalar Arrow iterator UDF""" num_input_rows = 0 def extract_args(batch: pa.RecordBatch): nonlocal num_input_rows args = tuple(batch.column(o) for o in args_offsets) num_input_rows += batch.num_rows return args[0] if len(args) == 1 else args # Extract args from input batches (streaming) args_iter = map(extract_args, data) # Call UDF and verify result type (iterator of pa.Array) verified_iter = verify_return_type( udf_func(args_iter), Iterator[pa.Array], # type: ignore[type-abstract] ) # Process results: enforce schema and assemble into RecordBatch target_schema = pa.schema([pa.field("_0", arrow_return_type)]) def process_results(): for result in verified_iter: batch = pa.RecordBatch.from_arrays([result], ["_0"]) yield ArrowBatchTransformer.enforce_schema(batch, target_schema, safecheck=True) # Apply row limit check (fail-fast) limited = verify_output_row_limit( process_results(), lambda: num_input_rows, ) # Apply row count match check (final) matched = verify_iter_result_row_count( limited, lambda: num_input_rows, ) # Yield batches yield from matched # Verify iterator consumed verify_iterator_exhausted(args_iter) # profiling is not supported for UDF return func, None, ser, ser if eval_type == PythonEvalType.SQL_GROUPED_AGG_ARROW_UDF: import pyarrow as pa # Pre-compute target schema for output coercion 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: # pyarrow.concat_batches not supported before 19.0.0 # remove this once we drop support for old versions concatenated = pa.RecordBatch.from_struct_array( pa.concat_arrays([b.to_struct_array() for b in batch_list]) ) results = [ udf_func( *[concatenated.column(o) for o in args_offsets], **{k: concatenated.column(v) for k, v in kwargs_offsets.items()}, ) for udf_func, args_offsets, kwargs_offsets, _ in udfs ] 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_ARROW_ITER_UDF: import pyarrow as pa assert num_udfs == 1, "One GROUPED_AGG_ARROW_ITER UDF expected here." udf_func, args_offsets, kwargs_offsets, return_type = udfs[0] return_schema = to_arrow_schema( StructType([StructField("_0", return_type)]), timezone="UTC", prefers_large_types=runner_conf.use_large_var_types, ) def extract_args(batch): args = tuple(batch.column(o) for o in args_offsets) return args[0] if len(args) == 1 else args def grouped_func( split_index: int, data: Iterator["GroupedBatch"] ) -> Iterator[pa.RecordBatch]: for group in data: batch_iter = map(extract_args, group) result = udf_func(batch_iter) # Drain remaining batches to maintain stream position for _ in batch_iter: pass batch = pa.RecordBatch.from_arrays([pa.array([result])], ["_0"]) 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 col_names = ["_%d" % i for i in range(len(udfs))] output_schema = StructType( [StructField(name, rt) for name, (_, _, _, rt) in zip(col_names, 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 table = pa.Table.from_batches(batch_list).combine_chunks() all_series = ArrowBatchTransformer.to_pandas( table, timezone=runner_conf.timezone, prefer_int_ext_dtype=runner_conf.prefer_int_ext_dtype, ) results = [ udf_func( *[all_series[o] for o in args_offsets], **{k: all_series[v] for k, v in kwargs_offsets.items()}, ) for udf_func, args_offsets, kwargs_offsets, _ in udfs ] result_series = [pd.Series([r]) for r in results] yield PandasToArrowConversion.convert( result_series, output_schema, timezone=runner_conf.timezone, safecheck=runner_conf.safecheck, arrow_cast=True, prefers_large_types=runner_conf.use_large_var_types, assign_cols_by_name=runner_conf.assign_cols_by_name, int_to_decimal_coercion_enabled=runner_conf.int_to_decimal_coercion_enabled, ) # profiling is not supported for UDF return grouped_func, None, ser, ser if eval_type == PythonEvalType.SQL_GROUPED_AGG_PANDAS_ITER_UDF: import pyarrow as pa import pandas as pd assert num_udfs == 1, "One GROUPED_AGG_PANDAS_ITER UDF expected here." udf_func, args_offsets, _, return_type = udfs[0] output_schema = StructType([StructField("_0", return_type)]) def extract_series( batch: "pa.RecordBatch", ) -> Union["pd.Series", tuple["pd.Series", ...]]: # Convert one RecordBatch to a pandas Series per column, then select args: # - pd.Series for a single column # - tuple[pd.Series, ...] for multiple columns all_series = ArrowBatchTransformer.to_pandas( batch, timezone=runner_conf.timezone, prefer_int_ext_dtype=runner_conf.prefer_int_ext_dtype, ) series = tuple(all_series[o] for o in args_offsets) return series[0] if len(series) == 1 else series def grouped_func( split_index: int, data: Iterator["GroupedBatch"] ) -> Iterator[pa.RecordBatch]: for group in data: series_iter = map(extract_series, group) result = udf_func(series_iter) # Drain remaining batches to maintain stream position for _ in series_iter: pass yield PandasToArrowConversion.convert( [pd.Series([result])], output_schema, timezone=runner_conf.timezone, safecheck=runner_conf.safecheck, arrow_cast=True, prefers_large_types=False, assign_cols_by_name=runner_conf.assign_cols_by_name, int_to_decimal_coercion_enabled=runner_conf.int_to_decimal_coercion_enabled, ) # profiling is not supported for UDF return grouped_func, None, ser, ser if eval_type == PythonEvalType.SQL_WINDOW_AGG_ARROW_UDF: import pyarrow as pa window_bound_types_str = runner_conf.get("window_bound_types") window_bound_types = [t.strip().lower() for t in window_bound_types_str.split(",")] 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: # pyarrow.concat_batches not supported before 19.0.0 # remove this once we drop support for old versions 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 udf_index, (udf_func, args_offsets, kwargs_offsets, _) in enumerate(udfs): bound_type = window_bound_types[udf_index] if bound_type == "unbounded": result = udf_func( *[concatenated.column(o) for o in args_offsets], **{k: concatenated.column(v) for k, v in kwargs_offsets.items()}, ) result_arrays.append(pa.repeat(result, num_rows)) elif bound_type == "bounded": begin_col = concatenated.column(args_offsets[0]) end_col = concatenated.column(args_offsets[1]) results = [] for i in range(num_rows): offset = begin_col[i].as_py() length = end_col[i].as_py() - offset slices = [ concatenated.column(o).slice(offset=offset, length=length) for o in args_offsets[2:] ] kw_slices = { k: concatenated.column(v).slice(offset=offset, length=length) for k, v in kwargs_offsets.items() } results.append(udf_func(*slices, **kw_slices)) result_arrays.append(pa.array(results)) else: raise PySparkRuntimeError( errorClass="INVALID_WINDOW_BOUND_TYPE", messageParameters={"window_bound_type": bound_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_WINDOW_AGG_PANDAS_UDF: import pyarrow as pa import pandas as pd window_bound_types_str = runner_conf.get("window_bound_types") window_bound_types = [t.strip().lower() for t in window_bound_types_str.split(",")] col_names = ["_%d" % i for i in range(len(udfs))] output_schema = StructType( [StructField(name, rt) for name, (_, _, _, rt) in zip(col_names, 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 table = pa.Table.from_batches(batch_list).combine_chunks() all_series = ArrowBatchTransformer.to_pandas( table, timezone=runner_conf.timezone, prefer_int_ext_dtype=runner_conf.prefer_int_ext_dtype, ) num_rows = table.num_rows result_series = [] for udf_index, (udf_func, args_offsets, kwargs_offsets, _) in enumerate(udfs): bound_type = window_bound_types[udf_index] if bound_type == "unbounded": result = udf_func( *[all_series[o] for o in args_offsets], **{k: all_series[v] for k, v in kwargs_offsets.items()}, ) # Repeat the scalar result to match the window (group) length. result_series.append(pd.Series([result]).repeat(num_rows)) elif bound_type == "bounded": # args_offsets[0] and args_offsets[1] are begin_index and end_index. assert len(args_offsets) >= 2, len(args_offsets) # Index operation is faster on np.ndarray, so we turn the # index series into np arrays here for performance. begin_array = all_series[args_offsets[0]].values end_array = all_series[args_offsets[1]].values series = [all_series[o] for o in args_offsets[2:]] kw_series = {k: all_series[v] for k, v in kwargs_offsets.items()} results = [] for i in range(num_rows): # Note: Creating a slice from a series for each window is # actually pretty expensive. However, there # is no easy way to reduce cost here. # Note: s.iloc[i : j] is about 30% faster than s[i: j], with # the caveat that the created slices shares the same # memory with s. Therefore, user are not allowed to # change the value of input series inside the window # function. It is rare that user needs to modify the # input series in the window function, and therefore, # it is a reasonable restriction. # Note: Calling reset_index on the slices will increase the # cost of creating slices by about 100%. Therefore, for # performance reasons we don't do it here. slices = [s.iloc[begin_array[i] : end_array[i]] for s in series] kw_slices = { k: s.iloc[begin_array[i] : end_array[i]] for k, s in kw_series.items() } results.append(udf_func(*slices, **kw_slices)) result_series.append(pd.Series(results)) else: raise PySparkRuntimeError( errorClass="INVALID_WINDOW_BOUND_TYPE", messageParameters={"window_bound_type": bound_type}, ) yield PandasToArrowConversion.convert( result_series, output_schema, timezone=runner_conf.timezone, safecheck=runner_conf.safecheck, arrow_cast=True, prefers_large_types=runner_conf.use_large_var_types, assign_cols_by_name=runner_conf.assign_cols_by_name, int_to_decimal_coercion_enabled=runner_conf.int_to_decimal_coercion_enabled, ) # profiling is not supported for UDF return grouped_func, None, ser, ser if eval_type == PythonEvalType.SQL_GROUPED_MAP_ARROW_UDF: import pyarrow as pa assert num_udfs == 1, "One GROUPED_MAP_ARROW UDF expected here." grouped_udf, arg_offsets, return_type, num_udf_args = udfs[0] parsed_offsets = extract_key_value_indexes(arg_offsets) assert len(parsed_offsets) == 1, "Expected one pair of offsets for GROUPED_MAP_ARROW UDF." arrow_return_type = to_arrow_type( return_type, timezone="UTC", prefers_large_types=runner_conf.use_large_var_types ) arrow_return_schema = pa.schema(list(arrow_return_type)) key_offsets = parsed_offsets[0][0] value_offsets = parsed_offsets[0][1] def grouped_func( split_index: int, data: Iterator["GroupedBatch"] ) -> Iterator[pa.RecordBatch]: """Apply groupBy Arrow UDF (non-iterator variant).""" for group in data: # Flatten struct column into separate columns flattened = map(ArrowBatchTransformer.flatten_struct, group) # Materialize first batch to get keys first_batch = next(flattened) keys = pa.RecordBatch.from_arrays( [first_batch.columns[o] for o in key_offsets], [first_batch.schema.names[o] for o in key_offsets], ) value_batches = ( pa.RecordBatch.from_arrays( [b.columns[o] for o in value_offsets], [b.schema.names[o] for o in value_offsets], ) for b in itertools.chain((first_batch,), flattened) ) # Call UDF value_table = pa.Table.from_batches(value_batches) if num_udf_args == 1: result = grouped_udf(value_table) else: key = tuple(c[0] for c in keys.columns) result = grouped_udf(key, value_table) verify_return_type(result, pa.Table) # Verify types (and reorder by name when configured). result = ArrowBatchTransformer.enforce_schema( result, arrow_return_schema, arrow_cast=False, reorder_by_name=runner_conf.assign_cols_by_name, ) for batch in result.to_batches(): yield ArrowBatchTransformer.wrap_struct(batch) # profiling is not supported for UDF return grouped_func, None, ser, ser if eval_type == PythonEvalType.SQL_GROUPED_MAP_ARROW_ITER_UDF: import pyarrow as pa assert num_udfs == 1, "One GROUPED_MAP_ARROW_ITER UDF expected here." grouped_udf, arg_offsets, return_type, num_udf_args = udfs[0] parsed_offsets = extract_key_value_indexes(arg_offsets) assert len(parsed_offsets) == 1, ( "Expected one pair of offsets for GROUPED_MAP_ARROW_ITER UDF." ) arrow_return_type = to_arrow_type( return_type, timezone="UTC", prefers_large_types=runner_conf.use_large_var_types ) arrow_return_schema = pa.schema(list(arrow_return_type)) key_offsets = parsed_offsets[0][0] value_offsets = parsed_offsets[0][1] def grouped_func( split_index: int, data: Iterator["GroupedBatch"] ) -> Iterator[pa.RecordBatch]: """Apply groupBy Arrow UDF (iterator variant).""" for group in data: # Flatten struct column into separate columns flattened_iter = map(ArrowBatchTransformer.flatten_struct, group) # Materialize first batch to get keys first_batch = next(flattened_iter) keys = pa.RecordBatch.from_arrays( [first_batch.columns[o] for o in key_offsets], [first_batch.schema.names[o] for o in key_offsets], ) value_batches = ( pa.RecordBatch.from_arrays( [b.columns[o] for o in value_offsets], [b.schema.names[o] for o in value_offsets], ) for b in itertools.chain((first_batch,), flattened_iter) ) # Call UDF with iterator of batches if num_udf_args == 1: result = grouped_udf(value_batches) else: key = tuple(c[0] for c in keys.columns) result = grouped_udf(key, value_batches) # Verify (and reorder by name when configured) each output batch for batch in verify_return_type(result, Iterator[pa.RecordBatch]): batch = ArrowBatchTransformer.enforce_schema( batch, arrow_return_schema, arrow_cast=False, reorder_by_name=runner_conf.assign_cols_by_name, ) yield ArrowBatchTransformer.wrap_struct(batch) # Drain remaining input batches to maintain stream position for _ in value_batches: pass # profiling is not supported for UDF return grouped_func, None, ser, ser if eval_type == PythonEvalType.SQL_GROUPED_MAP_PANDAS_UDF: import pyarrow as pa import pandas as pd assert num_udfs == 1, "One GROUPED_MAP_PANDAS UDF expected here." grouped_udf, arg_offsets, return_type, num_udf_args = udfs[0] parsed_offsets = extract_key_value_indexes(arg_offsets) assert len(parsed_offsets) == 1, "Expected one pair of offsets for GROUPED_MAP_PANDAS UDF." key_offsets = parsed_offsets[0][0] value_offsets = parsed_offsets[0][1] output_schema = StructType([StructField("_0", return_type)]) def grouped_func( split_index: int, data: Iterator[Iterator[pa.RecordBatch]], ) -> Iterator[pa.RecordBatch]: """Apply groupBy Pandas UDF (non-iterator variant). The explicit ``del`` calls below keep peakmem bounded across groups. Without them, generator locals from the previous iteration stay bound on the frame until each statement in the next iteration rebinds its slot, so the input-side DataFrames overlap with the next group's allocations and the working set grows unbounded on wide-column, large-group inputs. ``del result`` runs on resume from yield, before ``data.__next__()`` is asked for the next group. """ for group in data: all_batches = list(group) if all_batches: table = pa.Table.from_batches(all_batches).combine_chunks() else: table = pa.table({}) all_series = ArrowBatchTransformer.to_pandas( table, timezone=runner_conf.timezone, prefer_int_ext_dtype=runner_conf.prefer_int_ext_dtype, ) value_df = pd.concat([all_series[o] for o in value_offsets], axis=1) if num_udf_args == 1: result = grouped_udf(value_df) else: key = tuple(all_series[o].iloc[0] for o in key_offsets) result = grouped_udf(key, value_df) del all_batches, table, all_series, value_df verify_pandas_result( result, return_type, runner_conf.assign_cols_by_name, truncate_return_schema=False, ) yield PandasToArrowConversion.convert( [result], output_schema, timezone=runner_conf.timezone, safecheck=runner_conf.safecheck, arrow_cast=True, prefers_large_types=runner_conf.use_large_var_types, assign_cols_by_name=runner_conf.assign_cols_by_name, int_to_decimal_coercion_enabled=runner_conf.int_to_decimal_coercion_enabled, ) del result # profiling is not supported for UDF return grouped_func, None, ser, ser if eval_type == PythonEvalType.SQL_GROUPED_MAP_PANDAS_ITER_UDF: import pyarrow as pa import pandas as pd assert num_udfs == 1, "One GROUPED_MAP_PANDAS_ITER UDF expected here." grouped_udf, arg_offsets, return_type, num_udf_args = udfs[0] parsed_offsets = extract_key_value_indexes(arg_offsets) assert len(parsed_offsets) == 1, ( "Expected one pair of offsets for GROUPED_MAP_PANDAS_ITER UDF." ) key_offsets = parsed_offsets[0][0] value_offsets = parsed_offsets[0][1] output_schema = StructType([StructField("_0", return_type)]) def grouped_func( split_index: int, data: Iterator[Iterator[pa.RecordBatch]], ) -> Iterator[pa.RecordBatch]: """Apply groupBy Pandas UDF (iterator variant). The UDF receives an Iterator[pd.DataFrame] per group and returns an Iterator[pd.DataFrame]. Input batches are converted to pandas lazily so peakmem stays bounded by a single batch rather than the whole group. """ for group in data: group_iter = iter(group) # Read the first batch to extract grouping keys. first_series = ArrowBatchTransformer.to_pandas( next(group_iter), timezone=runner_conf.timezone, prefer_int_ext_dtype=runner_conf.prefer_int_ext_dtype, ) def dataframe_iter(): yield pd.concat([first_series[o] for o in value_offsets], axis=1) for batch in group_iter: series = ArrowBatchTransformer.to_pandas( batch, timezone=runner_conf.timezone, prefer_int_ext_dtype=runner_conf.prefer_int_ext_dtype, ) yield pd.concat([series[o] for o in value_offsets], axis=1) if num_udf_args == 1: result = grouped_udf(dataframe_iter()) else: key = tuple(first_series[o].iloc[0] for o in key_offsets) result = grouped_udf(key, dataframe_iter()) for df in result: verify_pandas_result( df, return_type, runner_conf.assign_cols_by_name, truncate_return_schema=False, ) yield PandasToArrowConversion.convert( [df], output_schema, timezone=runner_conf.timezone, safecheck=runner_conf.safecheck, arrow_cast=True, prefers_large_types=runner_conf.use_large_var_types, assign_cols_by_name=runner_conf.assign_cols_by_name, int_to_decimal_coercion_enabled=runner_conf.int_to_decimal_coercion_enabled, ) # Drain remaining input batches to maintain stream position. for _ in group_iter: pass # profiling is not supported for UDF return grouped_func, None, ser, ser if eval_type == PythonEvalType.SQL_COGROUPED_MAP_ARROW_UDF: import pyarrow as pa assert num_udfs == 1, "One COGROUPED_MAP_ARROW UDF expected here." cogrouped_udf, arg_offsets, return_type, num_udf_args = udfs[0] parsed_offsets = extract_key_value_indexes(arg_offsets) arrow_return_type = to_arrow_type( return_type, timezone="UTC", prefers_large_types=runner_conf.use_large_var_types ) arrow_return_schema = pa.schema(list(arrow_return_type)) select_columns = ArrowBatchTransformer.select_columns left_key_cols, left_val_cols = parsed_offsets[0] right_key_cols, right_val_cols = parsed_offsets[1] def table_from_batches(batches, cols): return pa.Table.from_batches([select_columns(b, cols) for b in batches]) def cogrouped_func( split_index: int, data: Iterator[Tuple[list[pa.RecordBatch], list[pa.RecordBatch]]], ) -> Iterator[pa.RecordBatch]: """Apply cogroupBy Arrow UDF.""" for left_batches, right_batches in data: left_keys = table_from_batches(left_batches, left_key_cols) left_values = table_from_batches(left_batches, left_val_cols) right_keys = table_from_batches(right_batches, right_key_cols) right_values = table_from_batches(right_batches, right_val_cols) if num_udf_args == 2: result = cogrouped_udf(left_values, right_values) else: key_table = left_keys if left_keys.num_rows > 0 else right_keys key = tuple(c[0] for c in key_table.columns) result = cogrouped_udf(key, left_values, right_values) verify_return_type(result, pa.Table) # Verify types (and reorder by name when configured). result = ArrowBatchTransformer.enforce_schema( result, arrow_return_schema, arrow_cast=False, reorder_by_name=runner_conf.assign_cols_by_name, ) for batch in result.to_batches(): yield ArrowBatchTransformer.wrap_struct(batch) # profiling is not supported for UDF return cogrouped_func, None, ser, ser if eval_type == PythonEvalType.SQL_MAP_PANDAS_ITER_UDF: import pyarrow as pa import pandas as pd assert num_udfs == 1, "One MAP_PANDAS_ITER UDF expected here." map_udf, _, _, return_type = udfs[0] output_schema = StructType([StructField("_0", return_type)]) iter_type_label = ( "pandas.DataFrame" if isinstance(return_type, StructType) else "pandas.Series" ) elem_type = pd.DataFrame if isinstance(return_type, StructType) else pd.Series def func( split_index: int, data: Iterator[pa.RecordBatch], ) -> Iterator[pa.RecordBatch]: """Apply mapInPandas UDF.""" def dataframe_iter(): # Input batches have a single struct column (see # MapInBatchEvaluatorFactory); convert lazily so peakmem stays # bounded by one batch. for batch in data: yield ArrowBatchTransformer.to_pandas( batch, timezone=runner_conf.timezone, prefer_int_ext_dtype=runner_conf.prefer_int_ext_dtype, df_for_struct=True, )[0] # mapInPandas accepts any iterable (e.g. a list), not just an # iterator, so the standard verify_return_type (which requires an # Iterator) is intentionally not reused here. result = map_udf(dataframe_iter()) if not isinstance(result, Iterator) and not hasattr(result, "__iter__"): raise PySparkTypeError( errorClass="UDF_RETURN_TYPE", messageParameters={ "expected": "iterator of {}".format(iter_type_label), "actual": type(result).__name__, }, ) for df in result: if not isinstance(df, elem_type): raise PySparkTypeError( errorClass="UDF_RETURN_TYPE", messageParameters={ "expected": "iterator of {}".format(iter_type_label), "actual": "iterator of {}".format(type(df).__name__), }, ) verify_pandas_result( df, return_type, assign_cols_by_name=True, truncate_return_schema=True ) yield PandasToArrowConversion.convert( [df], output_schema, timezone=runner_conf.timezone, safecheck=runner_conf.safecheck, arrow_cast=True, prefers_large_types=runner_conf.use_large_var_types, assign_cols_by_name=runner_conf.assign_cols_by_name, int_to_decimal_coercion_enabled=runner_conf.int_to_decimal_coercion_enabled, ) # profiling is not supported for UDF return func, None, ser, ser if eval_type == PythonEvalType.SQL_COGROUPED_MAP_PANDAS_UDF: import pyarrow as pa import pandas as pd assert num_udfs == 1, "One COGROUPED_MAP_PANDAS UDF expected here." cogrouped_udf, arg_offsets, return_type, num_udf_args = udfs[0] parsed_offsets = extract_key_value_indexes(arg_offsets) left_key_offsets, left_value_offsets = parsed_offsets[0] right_key_offsets, right_value_offsets = parsed_offsets[1] output_schema = StructType([StructField("_0", return_type)]) def cogrouped_func( split_index: int, data: Iterator[Tuple[list[pa.RecordBatch], list[pa.RecordBatch]]], ) -> Iterator[pa.RecordBatch]: """Apply cogroupBy Pandas UDF.""" for left_batches, right_batches in data: left_table = pa.Table.from_batches(left_batches) right_table = pa.Table.from_batches(right_batches) left_series = ArrowBatchTransformer.to_pandas( left_table, timezone=runner_conf.timezone, prefer_int_ext_dtype=runner_conf.prefer_int_ext_dtype, ) right_series = ArrowBatchTransformer.to_pandas( right_table, timezone=runner_conf.timezone, prefer_int_ext_dtype=runner_conf.prefer_int_ext_dtype, ) left_df = pd.concat([left_series[o] for o in left_value_offsets], axis=1) right_df = pd.concat([right_series[o] for o in right_value_offsets], axis=1) if num_udf_args == 2: result = cogrouped_udf(left_df, right_df) else: key_series = ( [left_series[o] for o in left_key_offsets] if not left_df.empty else [right_series[o] for o in right_key_offsets] ) key = tuple(s.iloc[0] for s in key_series) result = cogrouped_udf(key, left_df, right_df) del left_batches, right_batches, left_table, right_table del left_series, right_series, left_df, right_df verify_pandas_result( result, return_type, runner_conf.assign_cols_by_name, truncate_return_schema=False, ) yield PandasToArrowConversion.convert( [result], output_schema, timezone=runner_conf.timezone, safecheck=runner_conf.safecheck, arrow_cast=True, prefers_large_types=runner_conf.use_large_var_types, assign_cols_by_name=runner_conf.assign_cols_by_name, int_to_decimal_coercion_enabled=runner_conf.int_to_decimal_coercion_enabled, ) del result # profiling is not supported for UDF return cogrouped_func, None, ser, ser if ( eval_type == PythonEvalType.SQL_ARROW_BATCHED_UDF and not runner_conf.use_legacy_pandas_udf_conversion ): import pyarrow as pa # --- UDF preparation --- udf_infos = [] for udf_func, udf_args_offsets, udf_kwargs_offsets, udf_return_type in udfs: wrapped_func, args_kwargs_offsets = wrap_kwargs_support( udf_func, udf_args_offsets, udf_kwargs_offsets ) zero_arg = len(args_kwargs_offsets) == 0 udf_infos.append( ( wrapped_func, args_kwargs_offsets or (0,), zero_arg, to_arrow_type( udf_return_type, timezone="UTC", prefers_large_types=runner_conf.use_large_var_types, ), LocalDataToArrowConversion._create_converter( udf_return_type, none_on_identity=True, int_to_decimal_coercion_enabled=runner_conf.int_to_decimal_coercion_enabled, ), ) ) col_names = [f"_{i}" for i in range(len(udfs))] # --- Input preparation --- arrow_to_py_converters = [ ArrowTableToRowsConversion._create_converter( f.dataType, none_on_identity=True, binary_as_bytes=runner_conf.binary_as_bytes ) for f in eval_conf.input_type ] @fail_on_stopiteration def _evaluate_batch_udf(udf_func, rows): if runner_conf.arrow_concurrency_level <= 0: return [udf_func(*row) for row in rows] from concurrent.futures import ThreadPoolExecutor with ThreadPoolExecutor(max_workers=runner_conf.arrow_concurrency_level) as pool: return list(pool.map(lambda row: udf_func(*row), rows)) def func(split_index: int, data: Iterator[pa.RecordBatch]) -> Iterator[pa.RecordBatch]: for input_batch in data: num_rows = input_batch.num_rows # --- Input: Arrow -> Python columns --- columns = [ ( [conv(v) for v in ArrowTableToRowsConversion._to_pylist(col)] if conv is not None else ArrowTableToRowsConversion._to_pylist(col) ) for col, conv in zip(input_batch.itercolumns(), arrow_to_py_converters) ] if not columns: columns = [[_NoValue] * num_rows] # --- Process: evaluate each UDF row-by-row --- output_arrays = [] for udf_func, offsets, zero_arg, arrow_return_type, result_conv in udf_infos: rows = ( [() for _ in range(num_rows)] if zero_arg else list(zip(*[columns[o] for o in offsets])) ) results = _evaluate_batch_udf(udf_func, rows) verify_result_row_count(len(results), num_rows) # --- Output: Python -> Arrow --- converted = ( [result_conv(r) for r in results] if result_conv is not None else results ) try: arr = pa.array(converted, type=arrow_return_type) except pa.lib.ArrowInvalid: arr = pa.array(converted).cast( target_type=arrow_return_type, safe=runner_conf.safecheck ) output_arrays.append(arr) yield pa.RecordBatch.from_arrays(output_arrays, col_names) # profiling is not supported for UDF return func, None, ser, ser if ( eval_type == PythonEvalType.SQL_ARROW_BATCHED_UDF and runner_conf.use_legacy_pandas_udf_conversion ): import pandas as pd import pyarrow as pa # --- UDF preparation --- udf_infos = [] for udf_func, udf_args_offsets, udf_kwargs_offsets, udf_return_type in udfs: wrapped_func, args_kwargs_offsets = wrap_kwargs_support( udf_func, udf_args_offsets, udf_kwargs_offsets ) zero_arg = len(args_kwargs_offsets) == 0 # Legacy coerces String/Binary for Arrow compatibility coerce = ( str if isinstance(udf_return_type, StringType) else bytes if isinstance(udf_return_type, BinaryType) else None ) udf_infos.append( ( wrapped_func, args_kwargs_offsets or (0,), zero_arg, udf_return_type, coerce, ) ) col_names = [f"_{i}" for i in range(len(udfs))] return_schema = StructType( [StructField(name, info[3]) for name, info in zip(col_names, udf_infos)] ) @fail_on_stopiteration def _evaluate_batch_udf_legacy(udf_func, rows): if runner_conf.arrow_concurrency_level <= 0: return [udf_func(*row) for row in rows] from concurrent.futures import ThreadPoolExecutor with ThreadPoolExecutor(max_workers=runner_conf.arrow_concurrency_level) as pool: return list(pool.map(lambda row: udf_func(*row), rows)) def func(split_index: int, data: Iterator[pa.RecordBatch]) -> Iterator[pa.RecordBatch]: for input_batch in data: # --- Input: Arrow -> pandas columns --- pandas_columns = ArrowBatchTransformer.to_pandas( input_batch, timezone=runner_conf.timezone, schema=eval_conf.input_type, struct_in_pandas="row", ndarray_as_list=True, prefer_int_ext_dtype=runner_conf.prefer_int_ext_dtype, df_for_struct=False, ) num_rows = len(pandas_columns[0]) if pandas_columns else input_batch.num_rows if not pandas_columns: pandas_columns = [pd.Series([_NoValue] * num_rows)] # --- Process: evaluate each UDF row-by-row --- result_series = [] for udf_func, offsets, zero_arg, _, coerce in udf_infos: rows = ( [() for _ in range(num_rows)] if zero_arg else list(zip(*[pandas_columns[o].tolist() for o in offsets])) ) results = _evaluate_batch_udf_legacy(udf_func, rows) verify_result_row_count(len(results), num_rows) if coerce: results = [coerce(v) if v is not None else v for v in results] result_series.append(pd.Series(results)) # --- Output: pandas -> Arrow --- yield PandasToArrowConversion.convert( result_series, return_schema, timezone=runner_conf.timezone, safecheck=runner_conf.safecheck, arrow_cast=True, prefers_large_types=runner_conf.use_large_var_types, assign_cols_by_name=runner_conf.assign_cols_by_name, int_to_decimal_coercion_enabled=runner_conf.int_to_decimal_coercion_enabled, ) # profiling is not supported for UDF return func, None, ser, ser if eval_type == PythonEvalType.SQL_ARROW_ELEMENTWISE_UDF: # This path exchanges data with the JVM over Arrow, so PyArrow is required. Fail with a # clear message rather than a bare ImportError from `import pyarrow` below. from pyspark.sql.pandas.utils import require_minimum_pyarrow_version require_minimum_pyarrow_version() import pyarrow as pa import pyarrow.compute as pc # Element-wise UDFs back higher-order lambdas like transform(arr, x -> udf(x)). # ExtractPythonUDFFromLambda rewrites them so the UDF receives *all* array elements # at once (as ``array<T>``) rather than per-element. Flatten once, evaluate once over # the batch, then re-nest with input offsets. Example: array<int> -> udf -> array<int>. # UDF preparation udf_infos = [] for udf_func, udf_args_offsets, udf_kwargs_offsets, udf_return_type in udfs: wrapped_func, args_kwargs_offsets = wrap_kwargs_support( udf_func, udf_args_offsets, udf_kwargs_offsets ) # UDF returns one value per element; return type was pickled, unchanged. # This is per-element, so element type equals the declared return type. element_return_type = udf_return_type udf_infos.append( ( wrapped_func, args_kwargs_offsets, to_arrow_type( element_return_type, timezone="UTC", prefers_large_types=runner_conf.use_large_var_types, ), LocalDataToArrowConversion._create_converter( element_return_type, none_on_identity=True, int_to_decimal_coercion_enabled=runner_conf.int_to_decimal_coercion_enabled, ), ) ) col_names = [f"_{i}" for i in range(len(udfs))] # Input: every argument arrives as ``array<T>`` aligned with the iterated array. # Flatten once per column; convert elements with the array's element type. input_fields = list(eval_conf.input_type) arrow_to_py_converters = [ ArrowTableToRowsConversion._create_converter( f.dataType.elementType, none_on_identity=True, binary_as_bytes=runner_conf.binary_as_bytes, ) for f in input_fields ] @fail_on_stopiteration def _evaluate_elementwise_udf(udf_func, rows): if runner_conf.arrow_concurrency_level <= 0: return [udf_func(*row) for row in rows] from concurrent.futures import ThreadPoolExecutor with ThreadPoolExecutor(max_workers=runner_conf.arrow_concurrency_level) as pool: return list(pool.map(lambda row: udf_func(*row), rows)) def renest_spec(shape): """Offsets, list class and null mask that re-nest a flat result list by ``shape``. Null rows stay null and consume no offsets. ``ListArray`` uses int32 offsets and ``LargeListArray`` int64, so the input's list width is preserved. """ lengths = pc.list_value_length(shape).to_pylist() offsets = [] running = 0 for n in lengths: offsets.append(running) if n is not None: running += n offsets.append(running) is_large = pa.types.is_large_list(shape.type) list_cls = pa.LargeListArray if is_large else pa.ListArray offsets_arr = pa.array(offsets, type=pa.int64() if is_large else pa.int32()) null_mask = pa.array([n is None for n in lengths], type=pa.bool_()) total_elements = running return list_cls, offsets_arr, null_mask, total_elements def func(split_index: int, data: Iterator[pa.RecordBatch]) -> Iterator[pa.RecordBatch]: for input_batch in data: # Flatten each list column once to its element list, converting to Python. columns = [] for col, conv in zip(input_batch.itercolumns(), arrow_to_py_converters): values = ArrowTableToRowsConversion._to_pylist(col.flatten()) if conv is not None: values = [conv(v) for v in values] columns.append(values) # Each UDF is re-nested by *its own* first argument's shape. ExtractPythonUDFs can # fuse UDFs over differently shaped arrays into one batch, so a single shared shape # would misalign every UDF but the first. The rewrite always passes at least one # array argument, so `offsets_meta` is non-empty. output_arrays = [] for wrapped_func, offsets_meta, arrow_element_type, result_conv in udf_infos: list_cls, offsets_arr, null_mask, total_elements = renest_spec( input_batch.column(offsets_meta[0]) ) # Stream the argument tuples rather than materializing a batch-sized list. rows = zip(*[columns[o] for o in offsets_meta]) results = _evaluate_elementwise_udf(wrapped_func, rows) verify_result_row_count(len(results), total_elements) # Convert results and re-nest to array<R> using that UDF's offsets. converted = ( [result_conv(r) for r in results] if result_conv is not None else results ) try: flat_arr = pa.array(converted, type=arrow_element_type) # Broader than the SQL_ARROW_BATCHED_UDF path above (which catches only # ArrowInvalid): the element-wise wrapper commonly returns list/struct-typed # elements, whose type mismatches surface as ArrowTypeError, so both are caught # before falling back to an explicit cast. except (pa.lib.ArrowInvalid, pa.lib.ArrowTypeError): flat_arr = pa.array(converted).cast( target_type=arrow_element_type, safe=runner_conf.safecheck ) output_arrays.append( list_cls.from_arrays(offsets_arr, flat_arr, mask=null_mask) ) yield pa.RecordBatch.from_arrays(output_arrays, col_names) # profiling is not supported for UDF return func, None, ser, ser if eval_type in ( PythonEvalType.SQL_SCALAR_PANDAS_ELEMENTWISE_UDF, PythonEvalType.SQL_SCALAR_ARROW_ELEMENTWISE_UDF, ): from pyspark.sql.pandas.utils import require_minimum_pyarrow_version require_minimum_pyarrow_version() import pyarrow as pa import pyarrow.compute as pc # A scalar pandas or Arrow UDF lifted out of a higher-order function's lambda by # ExtractPythonUDFFromLambda. Each argument arrives as ``array<T>`` aligned with the # iterated array. We flatten each list column to its element column, run the *vectorized* # function once over that flat column (so it still receives a pandas Series / DataFrame or a # pa.Array, its native contract), then re-nest the flat result to ``array<R>`` using the # input's offsets - one row in, one row out, one Python round trip per batch. is_pandas = eval_type == PythonEvalType.SQL_SCALAR_PANDAS_ELEMENTWISE_UDF udf_infos = [] for udf_func, udf_args_offsets, udf_kwargs_offsets, udf_return_type in udfs: wrapped_func, args_kwargs_offsets = wrap_kwargs_support( udf_func, udf_args_offsets, udf_kwargs_offsets ) # The UDF returns one value per element, so its declared return type is the element # type of the ``array<R>`` this operator produces. arrow_element_type = to_arrow_type( udf_return_type, timezone="UTC", prefers_large_types=runner_conf.use_large_var_types ) udf_infos.append( (wrapped_func, args_kwargs_offsets, udf_return_type, arrow_element_type) ) col_names = [f"_{i}" for i in range(len(udfs))] if is_pandas: import pandas as pd # Each argument is ``array<T>``; the vectorized function must see the element type ``T``. element_types = [f.dataType.elementType for f in eval_conf.input_type] def func(split_index: int, data: Iterator[pa.RecordBatch]) -> Iterator[pa.RecordBatch]: for input_batch in data: # Flatten each list column to its element array once per batch and share it across # fused UDFs (the 102 path does the same), rather than re-flattening per UDF. flat_columns = [ _elementwise_flatten_column( col.flatten(), element_types[i], is_pandas, runner_conf ) for i, col in enumerate(input_batch.itercolumns()) ] output_arrays = [] for wrapped_func, offsets, return_type, arrow_element_type in udf_infos: # Re-nest by this UDF's first argument's shape. Different UDFs in one operator # may iterate differently shaped arrays, so each re-nests by its own argument. shape = input_batch.column(offsets[0]) shape_lengths = pc.list_value_length(shape).to_pylist() total_elements = sum(n for n in shape_lengths if n is not None) result = wrapped_func(*[flat_columns[o] for o in offsets]) if is_pandas: if not hasattr(result, "__len__"): pd_type = ( "pandas.DataFrame" if isinstance(return_type, StructType) else "pandas.Series" ) raise PySparkTypeError( errorClass="UDF_RETURN_TYPE", messageParameters={ "expected": pd_type, "actual": type(result).__name__, }, ) # struct return type must be a DataFrame (matches the base pandas path). if isinstance(return_type, StructType) and not isinstance( result, pd.DataFrame ): raise PySparkValueError( "Invalid return type. Please make sure that the UDF returns a " "pandas.DataFrame when the specified return type is StructType." ) # Verify the flat length before re-nesting so a wrong-length result raises # the friendly RESULT_ROWS_MISMATCH rather than an opaque pyarrow error. verify_result_row_count(len(result), total_elements) else: # Arrow flavor: a non-array-like result (e.g. a bare int) raises the # friendly UDF_RETURN_TYPE rather than a bare TypeError from len(), matching # the base SQL_SCALAR_ARROW_UDF path, and also checks the flat length. verify_scalar_result(result, total_elements) flat_arr = _elementwise_result_to_arrow( result, return_type, arrow_element_type, is_pandas, runner_conf ) nested = _elementwise_renest( flat_arr, shape_lengths, pa.types.is_large_list(shape.type) ) output_arrays.append(nested) yield pa.RecordBatch.from_arrays(output_arrays, col_names) # profiling is not supported for UDF return func, None, ser, ser if eval_type in ( PythonEvalType.SQL_SCALAR_PANDAS_ITER_ELEMENTWISE_UDF, PythonEvalType.SQL_SCALAR_ARROW_ITER_ELEMENTWISE_UDF, ): from pyspark.sql.pandas.utils import require_minimum_pyarrow_version require_minimum_pyarrow_version() import collections import pyarrow as pa import pyarrow.compute as pc is_pandas = eval_type == PythonEvalType.SQL_SCALAR_PANDAS_ITER_ELEMENTWISE_UDF if is_pandas: import pandas as pd assert num_udfs == 1, "One SCALAR_*_ITER_ELEMENTWISE UDF expected here." udf_func, args_offsets, kwargs_offsets, return_type = udfs[0] assert not kwargs_offsets, "Iterator UDFs do not take keyword arguments." # A scalar iterator UDF (pandas or Arrow) lifted out of a higher-order function's lambda. # The user function keeps its iterator contract: it consumes an iterator of batches and # yields an iterator of batches, one output value per input value. We preserve that by # feeding it the *flattened* elements of each input batch and, since the JVM joins UDF # output to input positionally by row (one ``array<R>`` per input ``array<T>`` row, in # order), buffering a FIFO of the per-row element counts to re-group the streamed flat # results back into arrays. Output batch boundaries need not match input ones. # Each argument is ``array<T>``; the pandas function must see each argument's own element # type ``T`` (arguments may differ, e.g. an outer column repeated into an aligned array). element_types = [f.dataType.elementType for f in eval_conf.input_type] arrow_element_type = to_arrow_type( return_type, timezone="UTC", prefers_large_types=runner_conf.use_large_var_types ) is_large = None # set from the first input batch; the list width is uniform per column. def func(split_index: int, data: Iterator[pa.RecordBatch]) -> Iterator[pa.RecordBatch]: # FIFO of per-row element counts (None for a null array) awaiting their flat results, # and the flat elements produced so far but not yet enough to complete the head shapes. pending_shapes: "collections.deque" = collections.deque() num_input_elements = 0 def extract_flat(batch: pa.RecordBatch): nonlocal is_large, num_input_elements shape = batch.column(args_offsets[0]) if is_large is None: is_large = pa.types.is_large_list(shape.type) pending_shapes.append(pc.list_value_length(shape).to_pylist()) # Flatten each argument to its element column; the user function sees the flat # elements as a pandas Series / DataFrame (pandas) or a pa.Array (Arrow). Each # argument is converted with its own element type. flat_cols = [ _elementwise_flatten_column( batch.column(o).flatten(), element_types[o], is_pandas, runner_conf ) for o in args_offsets ] num_input_elements += len(flat_cols[0]) return flat_cols[0] if len(flat_cols) == 1 else tuple(flat_cols) flat_args_iter = map(extract_flat, data) if not is_pandas: verified_iter = verify_return_type( udf_func(flat_args_iter), Iterator[pa.Array], # type: ignore[type-abstract] ) else: pandas_iter_type = ( Iterator[pd.DataFrame] if isinstance(return_type, StructType) else Iterator[pd.Series] ) verified_iter = verify_return_type(udf_func(flat_args_iter), pandas_iter_type) # Buffer the streamed flat element results and emit an ``array<R>`` row as soon as the # shape at the head of the FIFO is fully covered. A row whose length is 0 (an empty # array) or None (a null array) needs no elements, so it is emitted immediately even # before any chunk arrives - this matters when a whole partition is empty/null arrays # and the UDF yields nothing, otherwise those rows would be dropped by the positional # JVM join. Chunks are held in a list and concatenated only when a shape spans more than # one, so a UDF that yields once per input batch (the common case) never re-copies the # buffer. ``empty_type`` supplies the element type for a zero-length emit; it tracks the # most recent chunk's type (even a zero-length chunk carries the flavor's type - the # pandas flavor types timestamps with the session timezone), falling back to the # UTC-typed ``arrow_element_type`` only before any chunk arrives, so all emitted batches # share one schema. pending_chunks: "list" = [] pending_len = 0 empty_type = arrow_element_type num_output_elements = 0 def emit_ready(): nonlocal pending_chunks, pending_len while pending_shapes: lengths = pending_shapes[0] needed = sum(n for n in lengths if n is not None) if needed > pending_len: break pending_shapes.popleft() if needed == 0: flat = pa.nulls(0, type=empty_type) else: combined = ( pending_chunks[0] if len(pending_chunks) == 1 else pa.concat_arrays(pending_chunks) ) flat = combined.slice(0, needed) remainder = combined.slice(needed) pending_chunks = [remainder] if len(remainder) else [] pending_len -= needed nested = _elementwise_renest(flat, lengths, bool(is_large)) yield pa.RecordBatch.from_arrays([nested], ["_0"]) def process_results(): nonlocal pending_chunks, pending_len, empty_type, num_output_elements for result in verified_iter: if is_pandas: verify_pandas_result( result, return_type, assign_cols_by_name=True, truncate_return_schema=True, ) chunk = _elementwise_result_to_arrow( result, return_type, arrow_element_type, is_pandas, runner_conf ) num_output_elements += len(chunk) # Fail fast if the UDF over-produces, before the buffer grows unbounded (the # base iterator paths do the same via verify_output_row_limit). if num_output_elements > num_input_elements: raise PySparkRuntimeError( errorClass="OUTPUT_EXCEEDS_INPUT_ROWS", messageParameters={} ) # Even a zero-length chunk carries the flavor's element type (the pandas flavor # types timestamps with the session timezone), so always take it: otherwise # rows emitted for an all-empty batch before the first non-empty chunk would use # the UTC-typed default and disagree with later batches, breaking the output # stream's single-schema contract. empty_type = chunk.type if len(chunk): pending_chunks.append(chunk) pending_len += len(chunk) yield from emit_ready() # The iterator is exhausted: every input row's flat elements must have arrived. verify_result_row_count(num_output_elements, num_input_elements) # Flush any residual all-empty / all-null rows (they consume no elements). if pending_shapes: yield from emit_ready() verify_iterator_exhausted(flat_args_iter) yield from process_results() # profiling is not supported for UDF return func, None, ser, ser if eval_type == PythonEvalType.SQL_SCALAR_PANDAS_UDF: import pandas as pd import pyarrow as pa # --- UDF preparation --- udf_infos = [] for udf_func, udf_args_offsets, udf_kwargs_offsets, udf_return_type in udfs: wrapped_func, args_kwargs_offsets = wrap_kwargs_support( udf_func, udf_args_offsets, udf_kwargs_offsets ) udf_infos.append((wrapped_func, args_kwargs_offsets, udf_return_type)) return_schema = StructType( [StructField(f"_{i}", info[2]) for i, info in enumerate(udf_infos)] ) def func(split_index: int, data: Iterator[pa.RecordBatch]) -> Iterator[pa.RecordBatch]: for input_batch in data: num_rows = input_batch.num_rows # --- Input: Arrow -> pandas Series (struct columns become DataFrames) --- pandas_columns = ArrowBatchTransformer.to_pandas( input_batch, timezone=runner_conf.timezone, struct_in_pandas="dict", ndarray_as_list=False, prefer_int_ext_dtype=runner_conf.prefer_int_ext_dtype, df_for_struct=True, ) # --- Process: evaluate each UDF column-wise on pandas Series --- results = [] for udf_func, offsets, udf_return_type in udf_infos: result = udf_func(*[pandas_columns[o] for o in offsets]) if not hasattr(result, "__len__"): pd_type = ( "pandas.DataFrame" if isinstance(udf_return_type, StructType) else "pandas.Series" ) raise PySparkTypeError( errorClass="UDF_RETURN_TYPE", messageParameters={ "expected": pd_type, "actual": type(result).__name__, }, ) verify_result_row_count(len(result), num_rows) # struct_in_pandas="dict": UDF must return DataFrame for struct types if isinstance(udf_return_type, StructType) and not isinstance( result, pd.DataFrame ): raise PySparkValueError( "Invalid return type. Please make sure that the UDF returns a " "pandas.DataFrame when the specified return type is StructType." ) results.append(result) # --- Output: pandas -> Arrow --- yield PandasToArrowConversion.convert( results, return_schema, timezone=runner_conf.timezone, safecheck=runner_conf.safecheck, arrow_cast=True, prefers_large_types=runner_conf.use_large_var_types, assign_cols_by_name=runner_conf.assign_cols_by_name, int_to_decimal_coercion_enabled=runner_conf.int_to_decimal_coercion_enabled, ) # profiling is not supported for UDF return func, None, ser, ser if eval_type == PythonEvalType.SQL_SCALAR_PANDAS_ITER_UDF: import pandas as pd import pyarrow as pa assert num_udfs == 1, "One SCALAR_PANDAS_ITER UDF expected here." udf_func, args_offsets, kwargs_offsets, return_type = udfs[0] # Pre-compute target schema for output coercion return_schema = StructType([StructField("_0", return_type)]) expected_iter_type = ( Iterator[pd.DataFrame] if isinstance(return_type, StructType) else Iterator[pd.Series] ) def func(split_index: int, data: Iterator[pa.RecordBatch]) -> Iterator[pa.RecordBatch]: """Apply scalar pandas iterator UDF""" num_input_rows = 0 def extract_args(batch: pa.RecordBatch): nonlocal num_input_rows # Input: Arrow -> pandas Series (struct columns become DataFrames) pandas_columns = ArrowBatchTransformer.to_pandas( batch, timezone=runner_conf.timezone, struct_in_pandas="dict", ndarray_as_list=False, prefer_int_ext_dtype=runner_conf.prefer_int_ext_dtype, df_for_struct=True, ) args = tuple(pandas_columns[o] for o in args_offsets) num_input_rows += batch.num_rows return args[0] if len(args) == 1 else args # Extract args from input batches (streaming) args_iter = map(extract_args, data) # Call UDF and verify result type (iterator of pd.Series / pd.DataFrame) verified_iter = verify_return_type(udf_func(args_iter), expected_iter_type) # Process results: verify each element and convert pandas -> Arrow def process_results(): for result in verified_iter: verify_pandas_result( result, return_type, assign_cols_by_name=True, truncate_return_schema=True ) yield PandasToArrowConversion.convert( [result], return_schema, timezone=runner_conf.timezone, safecheck=runner_conf.safecheck, arrow_cast=True, prefers_large_types=runner_conf.use_large_var_types, assign_cols_by_name=runner_conf.assign_cols_by_name, int_to_decimal_coercion_enabled=runner_conf.int_to_decimal_coercion_enabled, ) # Apply row limit check (fail-fast) limited = verify_output_row_limit( process_results(), lambda: num_input_rows, ) # Apply row count match check (final) matched = verify_iter_result_row_count( limited, lambda: num_input_rows, ) # Yield batches yield from matched # Verify iterator consumed verify_iterator_exhausted(args_iter) # profiling is not supported for UDF return func, None, ser, ser if eval_type == PythonEvalType.SQL_TRANSFORM_WITH_STATE_PANDAS_UDF: import pyarrow as pa import pandas as pd assert num_udfs == 1, "One TRANSFORM_WITH_STATE_PANDAS UDF expected here." udf, arg_offsets, return_type = udfs[0] # See TransformWithStateInPandasExec for how arg_offsets are used to # distinguish between grouping attributes and data attributes parsed_offsets = extract_key_value_indexes(arg_offsets) assert len(parsed_offsets) == 1, ( "Expected one pair of offsets for TRANSFORM_WITH_STATE_PANDAS UDF." ) key_offsets = parsed_offsets[0][0] value_offsets = parsed_offsets[0][1] output_schema = StructType([StructField("_0", return_type)]) stateful_processor_api_client = StatefulProcessorApiClient( eval_conf.state_server_socket_port, eval_conf.grouping_key_schema ) arrow_max_records_per_batch = runner_conf.arrow_max_records_per_batch arrow_max_records_per_batch = ( arrow_max_records_per_batch if arrow_max_records_per_batch > 0 else 2**31 - 1 ) arrow_max_bytes_per_batch = runner_conf.arrow_max_bytes_per_batch def transform_with_state_func( split_index: int, batches: Iterator[pa.RecordBatch], ) -> Iterator[pa.RecordBatch]: """Apply transformWithStateInPandas UDF. Data chunks for the same grouping key appear sequentially in the input batches but may span batch boundaries, so rows are regrouped by key and re-chunked into pandas DataFrames bounded by arrow_max_records_per_batch and arrow_max_bytes_per_batch. The UDF is invoked once per grouping key with a lazy iterator of chunks, then once for PROCESS_TIMER and once for COMPLETE. """ total_bytes = 0 total_rows = 0 average_arrow_row_size = 0.0 def row_stream(): nonlocal total_bytes, total_rows, average_arrow_row_size for batch in batches: # Short circuit batch size stats if the batch size is # unlimited as computing batch size is computationally # expensive. if arrow_max_bytes_per_batch != 2**31 - 1 and batch.num_rows > 0: total_bytes += sum( buf.size for col in batch.columns for buf in col.buffers() if buf is not None ) total_rows += batch.num_rows average_arrow_row_size = total_bytes / total_rows data_pandas = ArrowBatchTransformer.to_pandas( batch, timezone=runner_conf.timezone, prefer_int_ext_dtype=runner_conf.prefer_int_ext_dtype, ) for row in pd.concat(data_pandas, axis=1).itertuples(index=False): batch_key = tuple(row[o] for o in key_offsets) yield (batch_key, row) def generate_data_batches(): """ Deserialize ArrowRecordBatches and return a generator of (grouping key, pandas.DataFrame) chunks. This function must avoid materializing multiple Arrow RecordBatches into memory at the same time, and data chunks from the same grouping key should appear sequentially. """ for batch_key, group_rows in itertools.groupby(row_stream(), key=lambda x: x[0]): rows = [] for _, row in group_rows: rows.append(row) if ( len(rows) >= arrow_max_records_per_batch or len(rows) * average_arrow_row_size >= arrow_max_bytes_per_batch ): yield (batch_key, pd.DataFrame(rows)) rows = [] if rows: yield (batch_key, pd.DataFrame(rows)) def convert_results(result_iter): for result in result_iter: if isinstance(return_type, StructType) and not isinstance(result, pd.DataFrame): raise PySparkValueError( "Invalid return type. Please make sure that the UDF returns a " "pandas.DataFrame when the specified return type is StructType." ) yield PandasToArrowConversion.convert( [result], output_schema, timezone=runner_conf.timezone, safecheck=runner_conf.safecheck, arrow_cast=True, assign_cols_by_name=runner_conf.assign_cols_by_name, int_to_decimal_coercion_enabled=runner_conf.int_to_decimal_coercion_enabled, ) for key, group in itertools.groupby(generate_data_batches(), key=lambda x: x[0]): # This must be a generator expression - do not materialize. values_gen = (df.iloc[:, value_offsets] for _, df in group) yield from convert_results( udf( stateful_processor_api_client, TransformWithStateInPandasFuncMode.PROCESS_DATA, key, values_gen, ) ) yield from convert_results( udf( stateful_processor_api_client, TransformWithStateInPandasFuncMode.PROCESS_TIMER, None, iter([]), ) ) yield from convert_results( udf( stateful_processor_api_client, TransformWithStateInPandasFuncMode.COMPLETE, None, iter([]), ) ) # profiling is not supported for UDF return transform_with_state_func, None, ser, ser if eval_type == PythonEvalType.SQL_TRANSFORM_WITH_STATE_PANDAS_INIT_STATE_UDF: import pyarrow as pa import pandas as pd assert num_udfs == 1, "One TRANSFORM_WITH_STATE_PANDAS_INIT_STATE UDF expected here." udf, arg_offsets, return_type = udfs[0] # See TransformWithStateInPandasExec for how arg_offsets are used to # distinguish between grouping attributes and data attributes. # parsed offsets: # [ # [groupingKeyOffsets, dedupDataOffsets], # [initStateGroupingOffsets, dedupInitDataOffsets] # ] parsed_offsets = extract_key_value_indexes(arg_offsets) key_offsets = parsed_offsets[0][0] init_key_offsets = parsed_offsets[1][0] output_schema = StructType([StructField("_0", return_type)]) stateful_processor_api_client = StatefulProcessorApiClient( eval_conf.state_server_socket_port, eval_conf.grouping_key_schema ) arrow_max_records_per_batch = runner_conf.arrow_max_records_per_batch arrow_max_records_per_batch = ( arrow_max_records_per_batch if arrow_max_records_per_batch > 0 else 2**31 - 1 ) arrow_max_bytes_per_batch = runner_conf.arrow_max_bytes_per_batch def func( split_index: int, data: Iterator[pa.RecordBatch], ) -> Iterator[pa.RecordBatch]: """Apply transformWithStateInPandas UDF with initial state. The input batches carry two struct columns, ``inputData`` and ``initState``; each batch holds one or the other but never both. Rows are flattened out of whichever struct is present, regrouped by grouping key, and re-chunked into pandas DataFrames bounded by arrow_max_records_per_batch and arrow_max_bytes_per_batch. The UDF is invoked once per grouping key with two separate lazy iterators (data DataFrames and init-state DataFrames), then once for PROCESS_TIMER and once for COMPLETE. """ total_bytes = 0 total_rows = 0 average_arrow_row_size = 0.0 def flatten_columns(cur_batch: "pa.RecordBatch", col_name: str) -> "pa.Table": struct_column = cur_batch.column(cur_batch.schema.get_field_index(col_name)) # Check if the entire column is null: an empty table (no columns) # signals the struct is absent from this batch. if struct_column.null_count == len(struct_column): return pa.Table.from_arrays([], names=[]) field_names = [ struct_column.type[i].name for i in range(struct_column.type.num_fields) ] field_arrays = [ struct_column.field(i) for i in range(struct_column.type.num_fields) ] return pa.Table.from_arrays(field_arrays, names=field_names) def to_pandas(table: "pa.Table") -> list: return ArrowBatchTransformer.to_pandas( table, timezone=runner_conf.timezone, prefer_int_ext_dtype=runner_conf.prefer_int_ext_dtype, ) def row_stream() -> Iterator[tuple]: nonlocal total_bytes, total_rows, average_arrow_row_size for batch in data: # Short circuit batch size stats if the batch size is # unlimited as computing batch size is computationally # expensive. if arrow_max_bytes_per_batch != 2**31 - 1 and batch.num_rows > 0: total_bytes += sum( buf.size for col in batch.columns for buf in col.buffers() if buf is not None ) total_rows += batch.num_rows average_arrow_row_size = total_bytes / total_rows data_table = flatten_columns(batch, "inputData") init_table = flatten_columns(batch, "initState") # Empty table has no columns. Each batch carries either # input data or init state, never both. has_data = data_table.num_columns > 0 has_init = init_table.num_columns > 0 assert not (has_data and has_init) if has_data: for row in pd.concat(to_pandas(data_table), axis=1).itertuples(index=False): batch_key = tuple(row[o] for o in key_offsets) yield (batch_key, row, None) elif has_init: for row in pd.concat(to_pandas(init_table), axis=1).itertuples(index=False): batch_key = tuple(row[o] for o in init_key_offsets) yield (batch_key, None, row) empty_dataframe = pd.DataFrame() def generate_data_batches() -> Iterator[tuple]: """ Deserialize ArrowRecordBatches and return a generator of (grouping key, data DataFrame, init-state DataFrame) chunks. This function must avoid materializing multiple Arrow RecordBatches into memory at the same time, and data chunks from the same grouping key should appear sequentially. """ for batch_key, group_rows in itertools.groupby(row_stream(), key=lambda x: x[0]): rows = [] init_state_rows = [] for _, row, init_state_row in group_rows: if row is not None: rows.append(row) if init_state_row is not None: init_state_rows.append(init_state_row) total_len = len(rows) + len(init_state_rows) if ( total_len >= arrow_max_records_per_batch or total_len * average_arrow_row_size >= arrow_max_bytes_per_batch ): yield ( batch_key, pd.DataFrame(rows) if rows else empty_dataframe.copy(), ( pd.DataFrame(init_state_rows) if init_state_rows else empty_dataframe.copy() ), ) rows = [] init_state_rows = [] if rows or init_state_rows: yield ( batch_key, pd.DataFrame(rows) if rows else empty_dataframe.copy(), ( pd.DataFrame(init_state_rows) if init_state_rows else empty_dataframe.copy() ), ) def convert_results( result_iter: Iterable["pd.DataFrame"], ) -> Iterator["pa.RecordBatch"]: # TODO(SPARK-49100): add verification that elements in result_iter are # indeed of type pd.DataFrame and conform to assigned cols for result in result_iter: if isinstance(return_type, StructType) and not isinstance(result, pd.DataFrame): raise PySparkValueError( "Invalid return type. Please make sure that the UDF returns a " "pandas.DataFrame when the specified return type is StructType." ) yield PandasToArrowConversion.convert( [result], output_schema, timezone=runner_conf.timezone, safecheck=runner_conf.safecheck, arrow_cast=True, assign_cols_by_name=runner_conf.assign_cols_by_name, int_to_decimal_coercion_enabled=runner_conf.int_to_decimal_coercion_enabled, ) for key, group in itertools.groupby(generate_data_batches(), key=lambda x: x[0]): # These must be generator expressions - do not materialize. The # UDF receives the data and init-state DataFrames as two # separate iterators, with empty chunks filtered out. group_data, group_init = itertools.tee(group, 2) state_values = (data_df for _, data_df, _ in group_data if not data_df.empty) init_states = (init_df for _, _, init_df in group_init if not init_df.empty) yield from convert_results( udf( stateful_processor_api_client, TransformWithStateInPandasFuncMode.PROCESS_DATA, key, state_values, init_states, ) ) yield from convert_results( udf( stateful_processor_api_client, TransformWithStateInPandasFuncMode.PROCESS_TIMER, None, iter([]), iter([]), ) ) yield from convert_results( udf( stateful_processor_api_client, TransformWithStateInPandasFuncMode.COMPLETE, None, iter([]), iter([]), ) ) # profiling is not supported for UDF return func, None, ser, ser if eval_type == PythonEvalType.SQL_GROUPED_MAP_PANDAS_UDF_WITH_STATE: import pyarrow as pa import pandas as pd from pyspark.sql.streaming.state import GroupState assert num_udfs == 1, "One GROUPED_MAP_PANDAS_UDF_WITH_STATE UDF expected here." # See FlatMapGroupsInPandasWithStateExec for how arg_offsets are used to # distinguish between grouping attributes and data attributes. f, arg_offsets, return_type = udfs[0] parsed_offsets = extract_key_value_indexes(arg_offsets) key_offsets = parsed_offsets[0][0] value_offsets = parsed_offsets[0][1] state_object_schema = eval_conf.state_value_schema arrow_max_records_per_batch = runner_conf.arrow_max_records_per_batch arrow_max_records_per_batch = ( arrow_max_records_per_batch if arrow_max_records_per_batch > 0 else 2**31 - 1 ) pickle_ser = CPickleSerializer() # The output RecordBatch has three struct fields, accessed by position # (_0/_1/_2, not by name): a count column indicating how many data and # state rows are present, the UDF output data, and the serialized state. result_count_df_type = StructType( [ StructField("dataCount", IntegerType()), StructField("stateCount", IntegerType()), ] ) result_state_df_type = StructType( [ StructField("properties", StringType()), StructField("keyRowAsUnsafe", BinaryType()), StructField("object", BinaryType()), StructField("oldTimeoutTimestamp", LongType()), ] ) def to_pandas(batch: "pa.RecordBatch") -> list: return ArrowBatchTransformer.to_pandas( batch, timezone=runner_conf.timezone, struct_in_pandas="dict", ndarray_as_list=False, prefer_int_ext_dtype=runner_conf.prefer_int_ext_dtype, df_for_struct=False, ) def construct_state(state_info_col: dict) -> GroupState: """Construct a state instance from the value of the state info column.""" state_properties = json.loads(state_info_col["properties"]) state_info_col_object = state_info_col["object"] if state_info_col_object: state_object = pickle_ser.loads(state_info_col_object) else: state_object = None state_properties["optionalValue"] = state_object return GroupState( keyAsUnsafe=state_info_col["keyRowAsUnsafe"], valueSchema=state_object_schema, **state_properties, ) def gen_data_and_state( batches: Iterator["pa.RecordBatch"], ) -> Iterator[tuple]: """Deserialize ArrowRecordBatches into (list of pandas.Series, state) chunks. Each batch carries the data columns plus a trailing state-info column. For every state-info row, the matching data slice is cut out via its (startOffset, numRows) and converted to pandas. A single state instance is reused across all chunks of one grouping key (the key appears sequentially and its last chunk is flagged), so grouping the output by the state object is equivalent to grouping by key. This must not materialize multiple Arrow RecordBatches at once. """ state_for_current_group = None for batch in batches: batch_schema = batch.schema data_schema = pa.schema([batch_schema[i] for i in range(0, len(batch_schema) - 1)]) state_schema = pa.schema([batch_schema[-1]]) batch_columns = batch.columns data_columns = batch_columns[0:-1] state_column = batch_columns[-1] data_batch = pa.RecordBatch.from_arrays(data_columns, schema=data_schema) state_batch = pa.RecordBatch.from_arrays([state_column], schema=state_schema) state_pandas = to_pandas(state_batch)[0] for state_idx in range(0, len(state_pandas)): state_info_col = state_pandas.iloc[state_idx] if not state_info_col: # no more data with grouping key + state break data_start_offset = state_info_col["startOffset"] num_data_rows = state_info_col["numRows"] is_last_chunk = state_info_col["isLastChunk"] if state_for_current_group: # reuse the state already built for this group state = state_for_current_group else: # first occurrence of this group, construct a new state state = construct_state(state_info_col) if is_last_chunk: # last chunk for this group, drop the cached state state_for_current_group = None elif not state_for_current_group: # more chunks expected for this group, cache the state state_for_current_group = state data_batch_for_group = data_batch.slice(data_start_offset, num_data_rows) yield (to_pandas(data_batch_for_group), state) def verify_element(result: "pd.DataFrame") -> "pd.DataFrame": if not isinstance(result, pd.DataFrame): raise PySparkTypeError( errorClass="UDF_RETURN_TYPE", messageParameters={ "expected": "iterator of pandas.DataFrame", "actual": "iterator of {}".format(type(result).__name__), }, ) # The number of columns of the result must match the return type, # but an empty result with no columns at all is acceptable. if not ( len(result.columns) == len(return_type) or (len(result.columns) == 0 and result.empty) ): raise PySparkRuntimeError( errorClass="RESULT_COLUMN_SCHEMA_MISMATCH", messageParameters={ "expected": str(len(return_type)), "actual": str(len(result.columns)), }, ) return result def apply_udf_to_group(key_series: list, value_series_gen, state: GroupState): """Adapt the deserialized chunks to the user function signature. Extract the scalar grouping key, convert each chunk of value Series into a pandas DataFrame (lazily), invoke the UDF, and validate that every returned element is a pandas.DataFrame conforming to the return type. """ key = tuple(s[0] for s in key_series) values: Iterable if state.hasTimedOut: # On timeout the UDF is called with an empty DataFrame instead # of an empty iterator. values = [ pd.DataFrame(columns=pd.concat(next(value_series_gen), axis=1).columns), ] else: values = (pd.concat(x, axis=1) for x in value_series_gen) result_iter = f(key, values, state) if isinstance(result_iter, pd.DataFrame): raise PySparkTypeError( errorClass="UDF_RETURN_TYPE", messageParameters={ "expected": "iterable of pandas.DataFrame", "actual": type(result_iter).__name__, }, ) try: iter(result_iter) except TypeError: raise PySparkTypeError( errorClass="UDF_RETURN_TYPE", messageParameters={ "expected": "iterable", "actual": type(result_iter).__name__, }, ) return (verify_element(x) for x in result_iter) def construct_state_pdf(state: GroupState) -> "pd.DataFrame": """Construct a single-row pandas DataFrame from the state instance.""" state_properties = state.json().encode("utf-8") state_key_row_as_binary = state._keyAsUnsafe if state.exists: state_object = pickle_ser.dumps(state._value_schema.toInternal(state._value)) else: state_object = None state_old_timeout_timestamp = state.oldTimeoutTimestamp state_dict = { "properties": [state_properties], "keyRowAsUnsafe": [state_key_row_as_binary], "object": [state_object], "oldTimeoutTimestamp": [state_old_timeout_timestamp], } return pd.DataFrame.from_dict(state_dict) def construct_record_batch( pdfs: list, pdf_data_cnt: int, pdf_schema: StructType, state_pdfs: list, state_data_cnt: int, ) -> "pa.RecordBatch": """Construct a count/data/state RecordBatch from output DataFrames and states. Arrow RecordBatch requires all columns to have the same number of rows, so data and state are padded with empty rows to the max of the two counts; the count column records the real (unpadded) sizes. """ max_data_cnt = max(1, max(pdf_data_cnt, state_data_cnt)) # Only the first row of the count column is meaningful; the rest # repeat the same values for friendlier compression. count_dict = { "dataCount": [pdf_data_cnt] * max_data_cnt, "stateCount": [state_data_cnt] * max_data_cnt, } count_pdf = pd.DataFrame.from_dict(count_dict) empty_row_cnt_in_data = max_data_cnt - pdf_data_cnt empty_row_cnt_in_state = max_data_cnt - state_data_cnt empty_rows_pdf = pd.DataFrame( dict.fromkeys(pdf_schema.names), index=[x for x in range(0, empty_row_cnt_in_data)], ) empty_rows_state = pd.DataFrame( columns=["properties", "keyRowAsUnsafe", "object", "oldTimeoutTimestamp"], index=[x for x in range(0, empty_row_cnt_in_state)], ) pdfs.append(empty_rows_pdf) state_pdfs.append(empty_rows_state) merged_pdf = pd.concat(pdfs, ignore_index=True) merged_state_pdf = pd.concat(state_pdfs, ignore_index=True) # Fields map to _0=count, _1=output data, _2=state data. data = [count_pdf, merged_pdf, merged_state_pdf] schema = StructType( [ StructField("_0", result_count_df_type), StructField("_1", pdf_schema), StructField("_2", result_state_df_type), ] ) return PandasToArrowConversion.convert( data, schema, timezone=runner_conf.timezone, safecheck=runner_conf.safecheck, arrow_cast=True, assign_cols_by_name=runner_conf.assign_cols_by_name, int_to_decimal_coercion_enabled=runner_conf.int_to_decimal_coercion_enabled, ) def func( split_index: int, data: Iterator["pa.RecordBatch"], ) -> Iterator["pa.RecordBatch"]: """Apply applyInPandasWithState UDF. The input batches carry the data columns plus a trailing state-info column. Chunks are regrouped by grouping key (via the reused state object), the UDF is invoked once per key with a lazy iterator of data DataFrames and its state, and the output DataFrames plus updated state are batched into count/data/state RecordBatches bounded by arrow_max_records_per_batch. """ def result_and_state_stream() -> Iterator[tuple]: # The same state object is reused across all chunks of a group, # so grouping by it is equivalent to grouping by key. for state, group in itertools.groupby(gen_data_and_state(data), key=lambda x: x[1]): # These must stay lazy - do not materialize the data chunks. data_gen = (data_pandas for data_pandas, _ in group) # Consume the first chunk to extract the grouping key series. first_elem = next(data_gen) key_series = [first_elem[o] for o in key_offsets] value_series_gen = ( [x[o] for o in value_offsets] for x in itertools.chain([first_elem], data_gen) ) yield (apply_udf_to_group(key_series, value_series_gen, state), state) pdfs: list = [] state_pdfs: list = [] pdf_data_cnt = 0 state_data_cnt = 0 for result_iter, state in result_and_state_stream(): for pdf in result_iter: # Ignore empty pandas DataFrames. if len(pdf) > 0: pdf_data_cnt += len(pdf) pdfs.append(pdf) # Flush a batch once the record threshold is exceeded. if pdf_data_cnt > arrow_max_records_per_batch: yield construct_record_batch( pdfs, pdf_data_cnt, return_type, state_pdfs, state_data_cnt ) pdfs = [] state_pdfs = [] pdf_data_cnt = 0 state_data_cnt = 0 # The state must be captured after the result iterator is fully # consumed, so the UDF has run and the state is up to date. state_pdfs.append(construct_state_pdf(state)) state_data_cnt += 1 # Flush the trailing batch if it has any data or state left. if pdf_data_cnt > 0 or state_data_cnt > 0: yield construct_record_batch( pdfs, pdf_data_cnt, return_type, state_pdfs, state_data_cnt ) # profiling is not supported for UDF return func, None, ser, ser if eval_type == PythonEvalType.SQL_TRANSFORM_WITH_STATE_PYTHON_ROW_UDF: import pyarrow as pa assert num_udfs == 1, "One TRANSFORM_WITH_STATE_PYTHON_ROW UDF expected here." udf, arg_offsets, return_type = udfs[0] # See TransformWithStateInPySparkExec for how arg_offsets are used to # distinguish between grouping attributes and data attributes. parsed_offsets = extract_key_value_indexes(arg_offsets) key_offsets = parsed_offsets[0][0] stateful_processor_api_client = StatefulProcessorApiClient( eval_conf.state_server_socket_port, eval_conf.grouping_key_schema ) def func( split_index: int, data: Iterator[pa.RecordBatch], ) -> Iterator[pa.RecordBatch]: """Apply transformWithStateInPySpark UDF over Row objects. Input batches are read row by row without materializing whole batches. Rows carrying the same grouping key appear sequentially, so they are regrouped by key and the UDF is invoked once per key with a lazy iterator of Row objects, then once for PROCESS_TIMER and once for COMPLETE. The UDF yields (iterator of Row, Spark type) pairs that are converted back into Arrow RecordBatches wrapped in a single struct column for the output stream. """ def generate_data_batches() -> Iterator[Tuple[Any, Any]]: """ Deserialize ArrowRecordBatches and return a generator of (grouping key, Row) tuples. This function must avoid materializing multiple Arrow RecordBatches into memory at the same time, and data chunks from the same grouping key should appear sequentially. """ for batch in data: DataRow = Row(*batch.schema.names) # Iterate row by row without converting the whole batch. num_cols = batch.num_columns for row_idx in range(batch.num_rows): row_key = tuple(batch[o][row_idx].as_py() for o in key_offsets) row = DataRow(*(batch.column(i)[row_idx].as_py() for i in range(num_cols))) yield row_key, row def convert_results(result_rows: Iterable[Any]) -> Iterator["pa.RecordBatch"]: # TODO(SPARK-XXXXX): add verification that elements in result_rows # are indeed of type Row and conform to assigned cols # Convert spark type to arrow type # TODO: we need to make this configurable, currently using default values. arrow_type = to_arrow_type( return_type, timezone="UTC", prefers_large_types=False, ) rows_as_dict = [row.asDict(True) for row in result_rows] pdf_schema = pa.schema(list(arrow_type)) record_batch = pa.RecordBatch.from_pylist(rows_as_dict, schema=pdf_schema) yield ArrowBatchTransformer.wrap_struct(record_batch) for key, group in itertools.groupby(generate_data_batches(), key=lambda x: x[0]): # This must be a generator expression - do not materialize. values_gen = map(lambda x: x[1], group) yield from convert_results( udf( stateful_processor_api_client, TransformWithStateInPandasFuncMode.PROCESS_DATA, key, values_gen, ) ) yield from convert_results( udf( stateful_processor_api_client, TransformWithStateInPandasFuncMode.PROCESS_TIMER, None, iter([]), ) ) yield from convert_results( udf( stateful_processor_api_client, TransformWithStateInPandasFuncMode.COMPLETE, None, iter([]), ) ) # profiling is not supported for UDF return func, None, ser, ser if eval_type == PythonEvalType.SQL_TRANSFORM_WITH_STATE_PYTHON_ROW_INIT_STATE_UDF: import pyarrow as pa assert num_udfs == 1, "One TRANSFORM_WITH_STATE_PYTHON_ROW_INIT_STATE UDF expected here." udf, arg_offsets, return_type = udfs[0] # See TransformWithStateInPandasExec for how arg_offsets are used to # distinguish between grouping attributes and data attributes. # parsed offsets: # [ # [groupingKeyOffsets, dedupDataOffsets], # [initStateGroupingOffsets, dedupInitDataOffsets] # ] parsed_offsets = extract_key_value_indexes(arg_offsets) key_offsets = parsed_offsets[0][0] init_key_offsets = parsed_offsets[1][0] stateful_processor_api_client = StatefulProcessorApiClient( eval_conf.state_server_socket_port, eval_conf.grouping_key_schema ) arrow_max_records_per_batch = runner_conf.arrow_max_records_per_batch arrow_max_records_per_batch = ( arrow_max_records_per_batch if arrow_max_records_per_batch > 0 else 2**31 - 1 ) def func( split_index: int, data: Iterator[pa.RecordBatch], ) -> Iterator[pa.RecordBatch]: """Apply transformWithStateInPySpark UDF with initial state over Rows. The input batches carry two struct columns, ``inputData`` and ``initState``; each batch holds one or the other but never both. Rows are flattened out of whichever struct is present, regrouped by grouping key, and re-chunked into Row lists bounded by arrow_max_records_per_batch. The UDF is invoked once per chunk with two separate iterators (data Rows and init-state Rows), then once for PROCESS_TIMER and once for COMPLETE. The UDF yields (iterator of Row, Spark type) pairs that are converted back into Arrow RecordBatches wrapped in a single struct column for the output stream. """ def extract_rows( cur_batch: "pa.RecordBatch", col_name: str, offsets: list ) -> Optional[Iterator[Tuple[Any, Any]]]: data_column = cur_batch.column(cur_batch.schema.get_field_index(col_name)) # Check if the entire column is null. if data_column.null_count == len(data_column): return None data_field_names = [ data_column.type[i].name for i in range(data_column.type.num_fields) ] data_field_arrays = [ data_column.field(i) for i in range(data_column.type.num_fields) ] DataRow = Row(*data_field_names) table = pa.Table.from_arrays(data_field_arrays, names=data_field_names) if table.num_rows == 0: return None def row_iterator() -> Iterator[Tuple[Any, Any]]: for row_idx in range(table.num_rows): key = tuple(table.column(o)[row_idx].as_py() for o in offsets) row = DataRow( *(table.column(i)[row_idx].as_py() for i in range(table.num_columns)) ) yield (key, row) return row_iterator() def row_stream() -> Iterator[Tuple[Any, Optional[Any], Optional[Any]]]: # The arrow batch is written in the schema: # schema: StructType = new StructType() # .add("inputData", dataSchema) # .add("initState", initStateSchema) # We parse each batch into tuples of (key, inputData, initState). # Each batch will have either init_data or input_data, not both. for batch in data: input_result = extract_rows(batch, "inputData", key_offsets) init_result = extract_rows(batch, "initState", init_key_offsets) assert not (input_result is not None and init_result is not None) if input_result is not None: for key, input_data_row in input_result: yield (key, input_data_row, None) elif init_result is not None: for key, init_state_row in init_result: yield (key, None, init_state_row) def generate_data_batches() -> Iterator[Tuple[Any, Tuple[Any, Any]]]: """ Deserialize ArrowRecordBatches and return a generator of (grouping key, (data Rows iterator, init-state Rows iterator)) chunks bounded by arrow_max_records_per_batch. This function must avoid materializing multiple Arrow RecordBatches into memory at the same time, and data chunks from the same grouping key should appear sequentially. """ for k, group_rows in itertools.groupby(row_stream(), key=lambda x: x[0]): input_rows: list = [] init_rows: list = [] for _, input_row, init_row in group_rows: if input_row is not None: input_rows.append(input_row) if init_row is not None: init_rows.append(init_row) total_len = len(input_rows) + len(init_rows) if total_len >= arrow_max_records_per_batch: yield (k, (iter(input_rows), iter(init_rows))) input_rows = [] init_rows = [] if input_rows or init_rows: yield (k, (iter(input_rows), iter(init_rows))) def convert_results(result_rows: Iterable[Any]) -> Iterator["pa.RecordBatch"]: # TODO(SPARK-XXXXX): add verification that elements in result_rows # are indeed of type Row and conform to assigned cols # Convert spark type to arrow type # TODO: we need to make this configurable, currently using default values. arrow_type = to_arrow_type( return_type, timezone="UTC", prefers_large_types=False, ) rows_as_dict = [row.asDict(True) for row in result_rows] pdf_schema = pa.schema(list(arrow_type)) record_batch = pa.RecordBatch.from_pylist(rows_as_dict, schema=pdf_schema) yield ArrowBatchTransformer.wrap_struct(record_batch) for key, group in itertools.groupby(generate_data_batches(), key=lambda x: x[0]): # These must be generator expressions - do not materialize. for _, (values_gen, init_states_gen) in group: yield from convert_results( udf( stateful_processor_api_client, TransformWithStateInPandasFuncMode.PROCESS_DATA, key, values_gen, init_states_gen, ) ) yield from convert_results( udf( stateful_processor_api_client, TransformWithStateInPandasFuncMode.PROCESS_TIMER, None, iter([]), iter([]), ) ) yield from convert_results( udf( stateful_processor_api_client, TransformWithStateInPandasFuncMode.COMPLETE, None, iter([]), iter([]), ) ) # profiling is not supported for UDF return func, None, ser, ser else: def mapper(a): result = tuple(f(*[a[o] for o in arg_offsets]) for arg_offsets, f in udfs) # In the special case of a single UDF this will return a single result rather # than a tuple of results; this is the format that the JVM side expects. if len(result) == 1: return result[0] else: return result def func(_, it): return map(mapper, it) # profiling is not supported for UDF return func, None, ser, ser def invoke_udf(message_receiver: SparkMessageReceiver, outfile: BinaryIO): """ This function is the main processing function for worker.py. It receives messages from the JVM, processes the data, and sends back results. This method goes through three phases: Initialization -> Processing -> Finish/Cleanup """ try: boot_time = time.time() # Initialization init_message = message_receiver.get_init_message() init_info = WorkerInitInfo.from_stream(init_message) start_faulthandler_periodic_traceback() check_python_version(init_info.python_version) memory_limit_mb = int(os.environ.get("PYSPARK_EXECUTOR_MEMORY_MB", "-1")) setup_memory_limits(memory_limit_mb) TaskContext._setTaskContext(init_info.task_context.to_task_context()) shuffle.MemoryBytesSpilled = 0 shuffle.DiskBytesSpilled = 0 setup_spark_files(init_info.spark_files_dir, init_info.python_includes) setup_broadcasts( init_info.broadcast.variables, init_info.broadcast.conn_info, init_info.broadcast.auth_secret, ) _accumulatorRegistry.clear() eval_type = init_info.eval_type runner_conf = RunnerConf(init_info.runner_conf) eval_conf = EvalConf(init_info.eval_conf) if eval_type == PythonEvalType.NON_UDF: assert isinstance(init_info.udf_info, (bytes, memoryview)) func, profiler, deserializer, serializer = read_command(pickleSer, init_info.udf_info) elif eval_type in ( PythonEvalType.SQL_TABLE_UDF, PythonEvalType.SQL_ARROW_TABLE_UDF, PythonEvalType.SQL_ARROW_UDTF, ): func, profiler, deserializer, serializer = read_udtf( pickleSer, init_info.udf_info, eval_type, runner_conf, eval_conf ) else: func, profiler, deserializer, serializer = read_udfs( pickleSer, init_info.udf_info, eval_type, runner_conf, eval_conf ) init_time = time.time() # Processing # Fetch the input data stream input_data_stream = message_receiver.get_data_stream() def process(): iterator = deserializer.load_stream(input_data_stream) out_iter = func(init_info.split_index, iterator) try: serializer.dump_stream(out_iter, outfile) finally: if hasattr(out_iter, "close"): out_iter.close() def pipelined_process(): """ Pipelined variant of process() that pre-fetches input batches in a background reader thread while the main thread computes the UDF and writes output. This allows input deserialization to overlap with UDF computation. """ import queue import threading queue_depth = int(os.environ.get("SPARK_PIPELINED_UDF_QUEUE_DEPTH", "2")) _SENTINEL = object() input_queue = queue.Queue(maxsize=queue_depth) reader_error = [None] # Event to signal the reader thread to stop (set by main thread on # exception or completion). The reader checks this after each failed # put attempt instead of polling with a timeout. stop_event = threading.Event() def _reader_thread(): try: for batch in deserializer.load_stream(input_data_stream): # Some serializers (e.g., ArrowStreamGroupSerializer) yield lazy # iterators that still read from the input stream. Materialize them here so # the main thread can consume them without touching the stream. if hasattr(batch, "__next__"): batch = list(batch) # Block on put, but wake up when stop_event is set. # stop_event.wait() returns immediately if already set. while not stop_event.is_set(): try: input_queue.put(batch, timeout=0.1) break except queue.Full: continue if stop_event.is_set(): return except Exception as e: reader_error[0] = e finally: # Enqueue sentinel so the consumer knows we're done. while not stop_event.is_set(): try: input_queue.put(_SENTINEL, timeout=0.1) break except queue.Full: continue t = threading.Thread( target=_reader_thread, name="pyspark-pipelined-reader", daemon=True ) t.start() def _queued_iter(): while True: item = input_queue.get() if item is _SENTINEL: if reader_error[0] is not None: raise reader_error[0] return yield item out_iter = func(init_info.split_index, _queued_iter()) try: serializer.dump_stream(out_iter, outfile) finally: if hasattr(out_iter, "close"): out_iter.close() # Signal reader thread to stop, drain the queue so it can unblock, # then wait for it to finish. stop_event.set() try: while not input_queue.empty(): input_queue.get_nowait() except Exception: pass # If the reader is still blocked in input_data_stream.read(), the stop_event # check only fires between put attempts -- it cannot interrupt a syscall. # Force-closing the stream here would break worker reuse (the next task uses # the same socket fd), so we settle for a bounded join and a loud warning # so an undetected leak shows up in the worker log. t.join(timeout=5) if t.is_alive(): warnings.warn( "pipelined reader thread did not exit within 5s; " "it may still be blocked in input_data_stream.read() and could " "read data intended for a subsequent reused-worker task. " "Consider disabling spark.python.worker.reuse if this recurs.", RuntimeWarning, ) is_pipelined = os.environ.get("SPARK_PIPELINED_UDF") == "1" if is_pipelined and hasattr(serializer, "_flush_per_batch"): serializer._flush_per_batch = True run_process = pipelined_process if is_pipelined else process processing_start_time = time.time() with capture_outputs(): if profiler: profiler.profile(run_process) else: run_process() processing_time_ms = int(1000 * (time.time() - processing_start_time)) # Cleanup # Reset task context to None. This is a guard code to avoid residual context when worker # reuse. TaskContext._setTaskContext(None) BarrierTaskContext._setTaskContext(None) except BaseException as e: handle_worker_exception(e, outfile) sys.exit(-1) finish_time = time.time() report_times(outfile, boot_time, init_time, finish_time, processing_time_ms) write_long(shuffle.MemoryBytesSpilled, outfile) write_long(shuffle.DiskBytesSpilled, outfile) # Mark the beginning of the accumulators section of the output write_int(SpecialLengths.END_OF_DATA_SECTION, outfile) send_accumulator_updates(outfile) # Check end of stream — raises if the finish signal is not received correctly. # Note: this call might fail due to other reasons (e.g. channel broke) # which will terminate the worker process. try: message_receiver.get_finish_signal_from_stream() write_int(SpecialLengths.END_OF_STREAM, outfile) except Exception: # Write a different value to tell JVM to not reuse this worker write_int(SpecialLengths.END_OF_DATA_SECTION, outfile) sys.exit(-1) @with_faulthandler def main(infile, outfile): # Instantiate socket message readers for executing the UDF socket_reader = SparkSocketMessageReceiver(infile) invoke_udf(socket_reader, outfile) if __name__ == "__main__": with get_sock_file_to_executor() as sock_file: main(sock_file, sock_file)