blob: e7af151cbc7b02f6399948814317afb3879a5012 [file]
#
# 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 dataclasses
import inspect
import itertools
import json
import os
import sys
import time
import warnings
from collections.abc import Iterator
from typing import (
TYPE_CHECKING,
Any,
BinaryIO,
Callable,
Iterable,
Optional,
Tuple,
Type,
TypeVar,
Union,
get_args,
get_origin,
)
T = TypeVar("T")
if TYPE_CHECKING:
import pandas as pd
import pyarrow as pa
from pyspark.sql.pandas._typing import GroupedBatch
from pyspark import _NoValue, shuffle
from pyspark.accumulators import (
SpecialAccumulatorIds,
_accumulatorRegistry,
_deserialize_accumulator,
)
from pyspark.errors import PySparkRuntimeError, PySparkTypeError, PySparkValueError
from pyspark.logger.worker_io import capture_outputs
from pyspark.messages import (
SparkMessageReceiver,
SparkSocketMessageReceiver,
)
from pyspark.serializers import (
BatchedSerializer,
CPickleSerializer,
SpecialLengths,
write_int,
write_long,
)
from pyspark.sql.conversion import (
ArrowBatchTransformer,
ArrowTableToRowsConversion,
LocalDataToArrowConversion,
PandasToArrowConversion,
)
from pyspark.sql.functions import SkipRestOfInputTableException
from pyspark.sql.pandas.serializers import (
ArrowStreamCoGroupSerializer,
ArrowStreamGroupSerializer,
ArrowStreamSerializer,
)
from pyspark.sql.pandas.types import to_arrow_schema, to_arrow_type
from pyspark.sql.streaming.stateful_processor_api_client import StatefulProcessorApiClient
from pyspark.sql.streaming.stateful_processor_util import TransformWithStateInPandasFuncMode
from pyspark.sql.types import (
ArrayType,
BinaryType,
DataType,
IntegerType,
LongType,
MapType,
Row,
StringType,
StructField,
StructType,
_create_row,
_parse_datatype_json_string,
)
from pyspark.taskcontext import BarrierTaskContext, TaskContext
from pyspark.util import (
PythonEvalType,
fail_on_stopiteration,
handle_worker_exception,
start_faulthandler_periodic_traceback,
with_faulthandler,
)
from pyspark.worker_message import WorkerInitInfo
from pyspark.worker_util import (
Conf,
check_python_version,
get_sock_file_to_executor,
pickleSer,
read_command,
send_accumulator_updates,
setup_broadcasts,
setup_memory_limits,
setup_spark_files,
)
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 elementwise_nesting(self) -> Optional[list]:
# Per-UDF nesting depth (parallel to the UDF list) for the element-wise lift: how many
# ``array`` levels the worker flattens off each argument and re-nests onto the result. A UDF
# in a single lambda is depth 1; one lifted out of nested lambdas is deeper. Absent/empty
# means depth 1 for every UDF. See ExtractPythonUDFFromLambda.
raw = self.get("elementwise_nesting", None)
if raw is None or raw == "":
return None
return [int(x) for x in raw.split(",")]
@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))
# Sentinel standing in for NaN grouping-key values in the map-side incremental-aggregate combine.
# ``float('nan') != float('nan')``, so distinct NaN objects would never collide in a dict; mapping
# them to one sentinel lets the map-side combine collapse them (correctness does not depend on it,
# as the FINAL stage re-groups authoritatively).
_NAN_GROUPING_KEY = object()
def _hashable_grouping_key(key_values: Tuple[Any, ...]) -> Any:
"""
Canonicalize a grouping-key value tuple (extracted from Arrow via ``to_pylist``) into a
hashable form usable as a ``dict`` key for the map-side PARTIAL combine of incremental Python
aggregators.
Lists (from ``array`` columns) become tuples and dicts (from ``struct`` columns) become tuples
of ``(name, value)`` pairs, so nested complex keys hash by value; NaN floats map to a sentinel.
This is a best-effort combine only: an exotic unhashable value falls back to a unique object so
the row forms its own group, and the FINAL stage merges any keys left uncollapsed here.
"""
def canon(v: Any) -> Any:
if isinstance(v, float) and v != v:
return _NAN_GROUPING_KEY
if isinstance(v, list):
return tuple(canon(x) for x in v)
if isinstance(v, dict):
return tuple((k, canon(x)) for k, x in v.items())
return v
try:
canonical = tuple(canon(v) for v in key_values)
hash(canonical)
return canonical
except TypeError:
return object()
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):
import pyspark.memory_profiler_ext
from pyspark.sql.profiler import (
ProfileResultsParam,
ProfileResultsParamV2,
WorkerMemoryProfiler,
)
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.
# Incremental Python aggregators: the pickled "function" is the Aggregator object itself, whose
# zero/reduce/merge/finish methods the worker calls directly. Return it unwrapped (not through
# fail_on_stopiteration, which would treat it as a plain callable).
if eval_type in (
PythonEvalType.SQL_GROUPED_AGG_ARROW_INCREMENTAL_PARTIAL_UDF,
PythonEvalType.SQL_GROUPED_AGG_ARROW_INCREMENTAL_FINAL_UDF,
PythonEvalType.SQL_WINDOW_AGG_ARROW_INCREMENTAL_UDF,
):
return chained_func, args_offsets, kwargs_offsets, return_type
# Scalar, aggregation and window UDFs: (func, args_offsets, kwargs_offsets, return_type).
if eval_type in (
PythonEvalType.SQL_ARROW_BATCHED_UDF,
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_leaf_type(data_type, depth):
"""The element type ``depth`` ``ArrayType`` levels below ``data_type``.
A lifted UDF's argument arrives as ``array^depth<T>`` (one ``array`` level per enclosing higher-
order function lambda); this peels them off to the scalar leaf ``T`` the user function sees. See
``ExtractPythonUDFFromLambda``.
"""
for _ in range(depth):
data_type = data_type.elementType
return data_type
def _elementwise_flatten_deep(col, depth):
"""Flatten ``depth`` list levels off ``col``, keeping each level's shape for re-nesting.
Returns ``(leaf, shape_levels, is_large_levels)``: ``leaf`` is the fully flattened element
``pa.Array`` (the leaves of the ``depth``-deep nesting), ``shape_levels[k]`` is the per-slot
length (``None`` for a null slot) at level ``k`` (0 = outermost), and ``is_large_levels[k]``
whether that level is a ``LargeListArray``. ``depth`` is 1 for a UDF in a single lambda and more
for one lifted out of nested lambdas. Shared by the element-wise worker paths. See
``ExtractPythonUDFFromLambda``.
"""
import pyarrow as pa
import pyarrow.compute as pc
shape_levels = []
is_large_levels = []
cur = col
for _ in range(depth):
shape_levels.append(pc.list_value_length(cur).to_pylist())
is_large_levels.append(pa.types.is_large_list(cur.type))
cur = cur.flatten()
return cur, shape_levels, is_large_levels
def _elementwise_flatten_leaf(col, depth):
"""Flatten ``depth`` list levels off ``col`` to its leaf ``pa.Array``, without capturing shape.
A lifted UDF re-nests its result by the *first* argument's per-level shapes only, so the other
arguments need just their leaves. This skips the ``pc.list_value_length(...).to_pylist()`` and
``is_large`` bookkeeping ``_elementwise_flatten_deep`` does for the first argument. See
``ExtractPythonUDFFromLambda``.
"""
cur = col
for _ in range(depth):
cur = cur.flatten()
return cur
def _elementwise_renest_deep(flat_values, shape_levels, is_large_levels):
"""Re-nest a flat leaf Array back through ``len(shape_levels)`` list levels, innermost first.
Inverse of ``_elementwise_flatten_deep``: rebuilds the ``array^depth<R>`` result from the flat
per-leaf results and the per-level shapes captured while flattening the input. For ``depth`` 1
this is a single ``_elementwise_renest``.
"""
result = flat_values
for lengths, is_large in zip(reversed(shape_levels), reversed(is_large_levels)):
result = _elementwise_renest(result, lengths, is_large)
return result
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,
PythonEvalType.SQL_GROUPED_AGG_ARROW_INCREMENTAL_PARTIAL_UDF,
PythonEvalType.SQL_GROUPED_AGG_ARROW_INCREMENTAL_FINAL_UDF,
PythonEvalType.SQL_WINDOW_AGG_ARROW_INCREMENTAL_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,
# The map-side PARTIAL stage streams ordinary (multi-group) batches and hash-combines
# inside the worker, so it uses the plain (non-grouped) stream serializer below. Only
# the post-shuffle FINAL stage receives one Arrow stream per group.
PythonEvalType.SQL_GROUPED_AGG_ARROW_INCREMENTAL_FINAL_UDF,
PythonEvalType.SQL_GROUPED_MAP_ARROW_ITER_UDF,
PythonEvalType.SQL_GROUPED_MAP_ARROW_UDF,
PythonEvalType.SQL_GROUPED_MAP_PANDAS_ITER_UDF,
PythonEvalType.SQL_GROUPED_MAP_PANDAS_UDF,
PythonEvalType.SQL_WINDOW_AGG_ARROW_UDF,
PythonEvalType.SQL_WINDOW_AGG_PANDAS_UDF,
PythonEvalType.SQL_WINDOW_AGG_ARROW_INCREMENTAL_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_ARROW_INCREMENTAL_PARTIAL_UDF:
import pyarrow as pa
# Map-side PARTIAL stage: hash-combine input rows into a per-group buffer via the
# aggregator's `reduce`. Ordinary (multi-group) batches are streamed in; the worker keeps
# one running buffer per distinct grouping key -- never whole groups of rows -- and, at end
# of partition, emits one row per key: the grouping key columns followed by one buffer
# struct column per aggregator. Because the FINAL stage re-groups these authoritatively
# after the shuffle, the worker's grouping only needs to be a best-effort combine: any keys
# it fails to collapse (e.g. NaN, which compares unequal to itself) are merged downstream.
#
# The leading `num_grouping_keys` input columns are the grouping keys (see the operator);
# `grouping_key_schema` carries their names/types so the emitted key columns round-trip.
grouping_key_schema = eval_conf.grouping_key_schema
num_grouping_keys = (
len(grouping_key_schema.fields) if grouping_key_schema is not None else 0
)
buffer_col_names = ["_%d" % i for i in range(len(udfs))]
buffer_arrow_types = [
to_arrow_type(
agg.bufferSchema,
timezone="UTC",
prefers_large_types=runner_conf.use_large_var_types,
)
for agg, _, _, _ in udfs
]
# Buffer field names are invariant across groups; compute them once per aggregator.
field_names_by_udf = [[f.name for f in agg.bufferSchema.fields] for agg, _, _, _ in udfs]
# The aggregator's `reduce` receives a single positional tuple, so any named arguments at
# the call site are appended after the positional ones, in call order (kwargs_offsets
# preserves that order). This mirrors how a Python call `f(*args, **kwargs)` would order
# them into one value tuple.
input_offsets_by_udf = [
list(args_offsets) + list(kwargs_offsets.values())
for _, args_offsets, kwargs_offsets, _ in udfs
]
# Input columns actually consumed per batch: the leading grouping keys plus every
# aggregator input, deduplicated. A UDF input may reuse a grouping-key column (the operator
# dedups its projection), so converting by distinct offset avoids repeated `to_pylist()`.
needed_offsets = sorted(
set(range(num_grouping_keys)) | {o for offsets in input_offsets_by_udf for o in offsets}
)
# Cap the map-side buffer so a high-cardinality partition -- exactly where partial
# aggregation degenerates -- cannot grow the per-key dict without bound and OOM the worker.
# When the distinct-key count reaches the cap we flush the whole map as one batch and start
# fresh; end-of-partition buffers are emitted in equally bounded chunks. This is safe
# because the FINAL stage re-groups the emitted partial buffers authoritatively after the
# shuffle and merges any duplicate keys that early flushes produce. A non-positive
# maxRecordsPerBatch means "unbounded" (mirroring the reader side).
max_records = runner_conf.arrow_max_records_per_batch
cap = max_records if max_records > 0 else None
def func(split_index: int, data: Iterator[pa.RecordBatch]) -> Iterator[pa.RecordBatch]:
# hashable key -> (representative key value tuple, list of per-aggregator buffers)
groups: "dict[Any, tuple]" = {}
key_field_types: Optional[list] = None
def make_batch(entries: list) -> pa.RecordBatch:
arrays = []
names = []
for j in range(num_grouping_keys):
arrays.append(
pa.array(
[e[0][j] for e in entries],
type=key_field_types[j], # type: ignore
)
)
names.append("k_%d" % j)
for i, (agg, _, _, _) in enumerate(udfs):
field_names = field_names_by_udf[i]
structs = [
{name: e[1][i][t] for t, name in enumerate(field_names)} for e in entries
]
arrays.append(pa.array(structs, type=buffer_arrow_types[i]))
names.append(buffer_col_names[i])
return pa.RecordBatch.from_arrays(arrays, names)
for batch in data:
if key_field_types is None:
key_field_types = [batch.schema.field(j).type for j in range(num_grouping_keys)]
pylist_by_offset = {o: batch.column(o).to_pylist() for o in needed_offsets}
key_cols = [pylist_by_offset[j] for j in range(num_grouping_keys)]
udf_cols = [
[pylist_by_offset[o] for o in input_offsets_by_udf[i]] for i in range(len(udfs))
]
for r in range(batch.num_rows):
key_values = tuple(key_cols[j][r] for j in range(num_grouping_keys))
hashable_key = _hashable_grouping_key(key_values)
entry = groups.get(hashable_key)
if entry is None:
buffers = [agg.zero() for agg, _, _, _ in udfs]
groups[hashable_key] = (key_values, buffers)
else:
buffers = entry[1]
for i, (agg, _, _, _) in enumerate(udfs):
cols_i = udf_cols[i]
buffers[i] = agg.reduce(buffers[i], tuple(c[r] for c in cols_i))
if cap is not None and len(groups) >= cap:
yield make_batch(list(groups.values()))
groups = {}
if not groups:
# Empty partition (or fully flushed above): emit nothing more. The FINAL stage
# supplies the global-aggregation identity row when there is no input at all.
return
entries = list(groups.values())
if cap is None:
yield make_batch(entries)
else:
for start in range(0, len(entries), cap):
yield make_batch(entries[start : start + cap])
# profiling is not supported for UDF
return func, None, ser, ser
if eval_type == PythonEvalType.SQL_GROUPED_AGG_ARROW_INCREMENTAL_FINAL_UDF:
import pyarrow as pa
# Post-shuffle FINAL stage: merge each group's partial buffers via the aggregator's `merge`
# and produce the output via `finish`. Buffers are streamed and merged one batch at a time.
# Every group the JVM sends yields exactly one output row; null partial-buffer rows are
# skipped, so an empty global aggregation (a single all-null buffer row injected by the
# operator) still produces `finish(zero)`.
col_names = ["_%d" % i for i in range(len(udfs))]
return_schema = to_arrow_schema(
StructType([StructField(name, rt) for name, (_, _, _, rt) in zip(col_names, udfs)]),
timezone="UTC",
prefers_large_types=runner_conf.use_large_var_types,
)
# Buffer field names are invariant across groups and batches; compute once per aggregator.
field_names_by_udf = [[f.name for f in agg.bufferSchema.fields] for agg, _, _, _ in udfs]
def grouped_func(
split_index: int, data: Iterator["GroupedBatch"]
) -> Iterator[pa.RecordBatch]:
for group in data:
merged: list = [None] * len(udfs)
for batch in group:
for i, (agg, args_offsets, _, _) in enumerate(udfs):
field_names = field_names_by_udf[i]
m = merged[i]
for row in batch.column(args_offsets[0]).to_pylist():
if row is None:
continue
partial = tuple(row[name] for name in field_names)
m = partial if m is None else agg.merge(m, partial)
merged[i] = m
results = []
for i, (agg, _, _, _) in enumerate(udfs):
m = merged[i] if merged[i] is not None else agg.zero()
results.append(agg.finish(m))
# Type each output array explicitly (mirroring the PARTIAL stage) so a non-trivial
# outputType or an all-None column does not depend on Arrow type inference.
result_arrays = [
pa.array([r], type=return_schema.field(i).type) for i, r in enumerate(results)
]
batch = pa.RecordBatch.from_arrays(result_arrays, col_names)
yield ArrowBatchTransformer.enforce_schema(batch, return_schema)
# profiling is not supported for UDF
return grouped_func, None, ser, ser
if eval_type == PythonEvalType.SQL_GROUPED_AGG_PANDAS_UDF:
import pandas as pd
import pyarrow as pa
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 pandas as pd
import pyarrow as pa
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_ARROW_INCREMENTAL_UDF:
import pyarrow as pa
# Window aggregation with an incremental ``Aggregator``. The operator sends each frame --
# the whole partition for an unbounded frame, or per-row ``[begin, end)`` slices for a
# bounded one -- and the worker folds the frame's rows with ``reduce`` (from a fresh
# ``zero``) and produces the value with ``finish``, one output value per input row. A window
# has no shuffle, so the intermediate buffer never leaves the worker (unlike the two-stage
# groupBy path); ``merge`` is not used here.
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 fold(agg: Any, buffer: Any, value_cols: list, start: int, end: int) -> Any:
# Fold rows ``[start, end)`` (each a tuple across ``value_cols``, matching the call-site
# argument order) into ``buffer`` via the aggregator's ``reduce``.
for r in range(start, end):
buffer = agg.reduce(buffer, tuple(c[r] for c in value_cols))
return buffer
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, (agg, args_offsets, kwargs_offsets, _) in enumerate(udfs):
bound_type = window_bound_types[udf_index]
result_type = return_schema.field(udf_index).type
if bound_type == "unbounded":
# One frame spanning the whole partition: compute once, repeat per row.
value_cols = [concatenated.column(o).to_pylist() for o in args_offsets] + [
concatenated.column(v).to_pylist() for v in kwargs_offsets.values()
]
result = agg.finish(fold(agg, agg.zero(), value_cols, 0, num_rows))
result_arrays.append(pa.array([result] * num_rows, type=result_type))
elif bound_type == "bounded":
# Per-row frame ``[begin, end)``. Materialize the aggregator's input columns
# once; frames index into them by row.
begin_col = concatenated.column(args_offsets[0])
end_col = concatenated.column(args_offsets[1])
data_offsets = list(args_offsets[2:]) + list(kwargs_offsets.values())
value_cols = [concatenated.column(o).to_pylist() for o in data_offsets]
# When consecutive frames share the same lower bound and only grow on the
# right (e.g. rowsBetween(unboundedPreceding, currentRow)), extend the
# running buffer by the newly-included rows instead of refolding from
# ``zero`` -- O(n) overall rather than O(n^2). Otherwise -- the lower bound
# advanced (a row left the window, which ``reduce`` cannot subtract) or the
# frame shrank -- refold the frame from ``zero``.
results = []
have_running = False
running: Any = None
prev_begin = -1
prev_end = 0
for i in range(num_rows):
begin = begin_col[i].as_py()
end = end_col[i].as_py()
if have_running and begin == prev_begin and end >= prev_end:
running = fold(agg, running, value_cols, prev_end, end)
else:
running = fold(agg, agg.zero(), value_cols, begin, end)
have_running = True
prev_begin, prev_end = begin, end
results.append(agg.finish(running))
result_arrays.append(pa.array(results, type=result_type))
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 pandas as pd
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))]
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 pandas as pd
import pyarrow as pa
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 pandas as pd
import pyarrow as pa
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 pandas as pd
import pyarrow as pa
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 pandas as pd
import pyarrow as pa
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
# 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 each argument down to its
# leaves, evaluate once over the batch, then re-nest with the input offsets. A UDF lifted
# out of nested lambdas (e.g. transform(arr, i -> transform(i, x -> udf(x)))) flattens more
# than one ``array`` level - its per-UDF depth comes from ``elementwise_nesting``.
# Example: array<array<int>> -> udf(depth 2) -> array<array<int>>.
# UDF preparation
input_fields = list(eval_conf.input_type)
nesting = eval_conf.elementwise_nesting
udf_infos = []
for udf_index, udf in enumerate(udfs):
udf_func, udf_args_offsets, udf_kwargs_offsets, udf_return_type = udf
wrapped_func, args_kwargs_offsets = wrap_kwargs_support(
udf_func, udf_args_offsets, udf_kwargs_offsets
)
depth = nesting[udf_index] if nesting is not None else 1
# Each argument arrives as ``array^depth<T>``; convert its leaves with the element type
# ``T`` reached by peeling ``depth`` array levels.
arg_converters = [
ArrowTableToRowsConversion._create_converter(
_elementwise_leaf_type(input_fields[o].dataType, depth),
none_on_identity=True,
binary_as_bytes=runner_conf.binary_as_bytes,
)
for o in args_kwargs_offsets
]
udf_infos.append(
(
wrapped_func,
args_kwargs_offsets,
depth,
arg_converters,
# UDF returns one value per element; return type was pickled, unchanged. This is
# per-element, so element type equals the declared return type.
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))]
@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 func(split_index: int, data: Iterator[pa.RecordBatch]) -> Iterator[pa.RecordBatch]:
for input_batch in data:
# Each UDF is re-nested by *its own* first argument's shape. ExtractPythonUDFs can
# fuse UDFs over differently shaped/nested 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` is non-empty.
output_arrays = []
for info in udf_infos:
(
wrapped_func,
offsets,
depth,
arg_converters,
arrow_element_type,
result_conv,
) = info
# Flatten each argument `depth` list levels to its leaves; the first argument's
# per-level shapes drive the re-nest.
leaf0, shape_levels, is_large_levels = _elementwise_flatten_deep(
input_batch.column(offsets[0]), depth
)
columns = []
for i, (o, conv) in enumerate(zip(offsets, arg_converters)):
leaf = (
leaf0
if i == 0
else _elementwise_flatten_leaf(input_batch.column(o), depth)
)
values = ArrowTableToRowsConversion._to_pylist(leaf)
if conv is not None:
values = [conv(v) for v in values]
columns.append(values)
total_elements = len(columns[0])
# Stream the argument tuples rather than materializing a batch-sized list.
rows = zip(*columns)
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(
_elementwise_renest_deep(flat_arr, shape_levels, is_large_levels)
)
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
# 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
input_fields = list(eval_conf.input_type)
nesting = eval_conf.elementwise_nesting
udf_infos = []
for udf_index, udf in enumerate(udfs):
udf_func, udf_args_offsets, udf_kwargs_offsets, udf_return_type = udf
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
)
depth = nesting[udf_index] if nesting is not None else 1
# Each argument arrives as ``array^depth<T>``; the vectorized function must see the leaf
# element type ``T`` reached by peeling ``depth`` array levels.
arg_leaf_types = [
_elementwise_leaf_type(input_fields[o].dataType, depth) for o in args_kwargs_offsets
]
udf_infos.append(
(
wrapped_func,
args_kwargs_offsets,
udf_return_type,
arrow_element_type,
depth,
arg_leaf_types,
)
)
col_names = [f"_{i}" for i in range(len(udfs))]
if is_pandas:
import pandas as pd
def func(split_index: int, data: Iterator[pa.RecordBatch]) -> Iterator[pa.RecordBatch]:
for input_batch in data:
output_arrays = []
for (
wrapped_func,
offsets,
return_type,
arrow_element_type,
depth,
arg_leaf_types,
) in udf_infos:
# Flatten each argument `depth` list levels to its leaves and adapt to the
# vectorized fn's input. Different UDFs in one operator may iterate differently
# shaped or differently nested arrays, so each flattens and re-nests by its own
# argument (the first argument's per-level shapes drive the re-nest).
leaf0, shape_levels, is_large_levels = _elementwise_flatten_deep(
input_batch.column(offsets[0]), depth
)
total_elements = len(leaf0)
flat_columns = [
_elementwise_flatten_column(
leaf0
if i == 0
else _elementwise_flatten_leaf(input_batch.column(o), depth),
t,
is_pandas,
runner_conf,
)
for i, (o, t) in enumerate(zip(offsets, arg_leaf_types))
]
result = wrapped_func(*flat_columns)
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_deep(flat_arr, shape_levels, is_large_levels)
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
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 arrives as ``array^depth<T>``; the vectorized function must see the leaf
# element type ``T`` (arguments may differ, e.g. an outer column repeated into an aligned
# array). ``depth`` > 1 for a UDF lifted out of nested lambdas.
input_fields = list(eval_conf.input_type)
nesting = eval_conf.elementwise_nesting
depth = nesting[0] if nesting is not None else 1
arg_leaf_types = [
_elementwise_leaf_type(input_fields[o].dataType, depth) for o in args_offsets
]
arrow_element_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]:
# FIFO of per-input-batch nesting shapes awaiting their flat leaf results. Each entry
# (per-level shapes, per-level list width, leaf count) re-nests one input batch's worth
# of rows back to ``array^depth<R>`` once that many leaf results have streamed in.
pending_shapes: "collections.deque" = collections.deque()
num_input_elements = 0
def extract_flat(batch: pa.RecordBatch):
nonlocal num_input_elements
# Flatten each argument `depth` levels to its leaves; the first argument's per-level
# shapes re-nest this batch's rows. The user function sees the flat leaves as a
# pandas Series / DataFrame (pandas) or a pa.Array (Arrow), each with its leaf type.
leaf0, shape_levels, is_large_levels = _elementwise_flatten_deep(
batch.column(args_offsets[0]), depth
)
pending_shapes.append((shape_levels, is_large_levels, len(leaf0)))
num_input_elements += len(leaf0)
flat_cols = [
_elementwise_flatten_column(
leaf0 if i == 0 else _elementwise_flatten_leaf(batch.column(o), depth),
arg_leaf_types[i],
is_pandas,
runner_conf,
)
for i, o in enumerate(args_offsets)
]
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:
shape_levels, is_large_levels, needed = pending_shapes[0]
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_deep(flat, shape_levels, is_large_levels)
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 pandas as pd
import pyarrow as pa
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 pandas as pd
import pyarrow as pa
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 pandas as pd
import pyarrow as pa
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
elif eval_type == PythonEvalType.SQL_BATCHED_UDF:
# Plain Python (pickle) UDFs, the only eval type reaching this branch. read_single_udf
# prepared each UDF as an (arg_offsets, eval_func) pair. Apply every one to each input
# row: a single UDF yields its bare result, multiple UDFs yield a tuple of results,
# which is the shape the JVM side expects. num_udfs is fixed, so the single-result
# case is handled once here rather than by unwrapping a one-element tuple per row.
def func(split_index: int, data: Iterator[Any]) -> Iterator[Any]:
if num_udfs == 1:
arg_offsets, f = udfs[0]
return (f(*[row[offset] for offset in arg_offsets]) for row in data)
return (
tuple(f(*[row[offset] for offset in arg_offsets]) for arg_offsets, f in udfs)
for row in data
)
# profiling is not supported for UDF
return func, None, ser, ser
else:
raise ValueError("Unknown eval type: {}".format(eval_type))
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)