From b9da167e29cd6b6cd390b7fe2e5f613d532c42e7 Mon Sep 17 00:00:00 2001 From: Spenser Sun Date: Thu, 17 Sep 2026 22:23:07 +0000 Subject: [PATCH 1/5] [PYTHON] Route element-wise UDF conversion through batch interfaces --- python/pyspark/sql/conversion.py | 109 ++++++ python/pyspark/sql/tests/test_conversion.py | 52 +++ python/pyspark/worker.py | 374 +++++++++----------- 3 files changed, 326 insertions(+), 209 deletions(-) diff --git a/python/pyspark/sql/conversion.py b/python/pyspark/sql/conversion.py index 55e009ce5c64e..aed8617f5b345 100644 --- a/python/pyspark/sql/conversion.py +++ b/python/pyspark/sql/conversion.py @@ -180,6 +180,115 @@ def select_columns(cls, batch: "pa.RecordBatch", column_indices: list[int]) -> " [batch.schema.names[i] for i in column_indices], ) + @staticmethod + def concat_batches(batches: Sequence["pa.RecordBatch"]) -> "pa.RecordBatch": + """Concatenate same-schema RecordBatches by row. + + A single batch is returned unchanged. PyArrow before 19.0.0 has no ``concat_batches``; + the fallback concatenates the equivalent StructArrays and converts the result back to a + RecordBatch. Element-wise iterator UDFs use this when one input batch's flattened result + spans multiple output chunks. + """ + import pyarrow as pa + + assert batches + if len(batches) == 1: + return batches[0] + if hasattr(pa, "concat_batches"): + return pa.concat_batches(batches) + return pa.RecordBatch.from_struct_array( + pa.concat_arrays([batch.to_struct_array() for batch in batches]) + ) + + @staticmethod + def flatten_elementwise_inputs( + batch: "pa.RecordBatch", input_column_indices: Sequence[int], depth: int + ) -> tuple["pa.RecordBatch", list[list[Optional[int]]], list[bool]]: + """Flatten ``depth`` list levels from an element-wise UDF's input columns. + + Returns ``(flat_input_batch, shape_levels, is_large_levels)``. ``flat_input_batch`` + contains each selected input's fully flattened leaf Array under a positional ``_N`` name. + ``shape_levels[k]`` contains the per-slot list length at level ``k`` (0 is outermost), + using ``None`` for a null list. ``is_large_levels[k]`` records whether that level uses + ``LargeListArray`` and therefore requires int64 rather than int32 offsets when rebuilt. + + Only the first selected column supplies shape and list-width metadata. The other inputs are + aligned to it by ``ExtractPythonUDFFromLambda``, so recording their shapes would repeat + the ``list_value_length(...).to_pylist()`` work without changing re-nesting. ``depth`` is 1 + for a UDF in one higher-order-function lambda and greater for nested lambdas. + + Shared by the row, scalar pandas / Arrow, and iterator element-wise worker paths. See + ``ExtractPythonUDFFromLambda``. + """ + import pyarrow as pa + import pyarrow.compute as pc + + assert input_column_indices + assert depth > 0 + + flat_inputs = [] + shape_levels = [] + is_large_levels = [] + for input_index, column_index in enumerate(input_column_indices): + current = batch.column(column_index) + for _ in range(depth): + if input_index == 0: + shape_levels.append(pc.list_value_length(current).to_pylist()) + is_large_levels.append(pa.types.is_large_list(current.type)) + current = current.flatten() + flat_inputs.append(current) + + return ( + pa.RecordBatch.from_arrays( + flat_inputs, names=[f"_{index}" for index in range(len(flat_inputs))] + ), + shape_levels, + is_large_levels, + ) + + @staticmethod + def renest_elementwise_outputs( + flat_outputs: Sequence[tuple["pa.RecordBatch", list[list[Optional[int]]], list[bool]]], + column_names: Sequence[str], + ) -> "pa.RecordBatch": + """Rebuild nested list columns from flattened element-wise UDF result batches. + + Each input tuple contains a one-column flat result batch plus the ``shape_levels`` and + ``is_large_levels`` returned by ``flatten_elementwise_inputs`` for that UDF. Levels are + rebuilt from innermost to outermost. A ``None`` length creates a null list and consumes no + flat values; a zero length creates an empty, non-null list. ``is_large_levels`` preserves + each input level's int32 ``ListArray`` versus int64 ``LargeListArray`` offset width. + + Different fused UDFs may carry different shapes, so every result batch is rebuilt with its + own metadata before the columns are assembled into one output RecordBatch. This is the + batch-level inverse of ``flatten_elementwise_inputs``. + """ + import pyarrow as pa + + assert len(flat_outputs) == len(column_names) + nested_columns = [] + for flat_batch, shape_levels, is_large_levels in flat_outputs: + assert flat_batch.num_columns == 1 + result = flat_batch.column(0) + for shape_lengths, is_large in zip(reversed(shape_levels), reversed(is_large_levels)): + offsets = [0] + running = 0 + nulls = [] + for length in shape_lengths: + nulls.append(length is None) + if length is not None: + running += length + offsets.append(running) + list_type = pa.LargeListArray if is_large else pa.ListArray + result = list_type.from_arrays( + pa.array(offsets, type=pa.int64() if is_large else pa.int32()), + result, + mask=pa.array(nulls, type=pa.bool_()), + ) + nested_columns.append(result) + + return pa.RecordBatch.from_arrays(nested_columns, names=column_names) + @staticmethod def wrap_struct(batch: "pa.RecordBatch") -> "pa.RecordBatch": """ diff --git a/python/pyspark/sql/tests/test_conversion.py b/python/pyspark/sql/tests/test_conversion.py index b03584d468c32..670109816f4f6 100644 --- a/python/pyspark/sql/tests/test_conversion.py +++ b/python/pyspark/sql/tests/test_conversion.py @@ -115,6 +115,58 @@ def test_flatten_struct_empty_batch(self): self.assertEqual(flattened.num_rows, 0) self.assertEqual(flattened.num_columns, 2) + def test_concat_batches(self): + import pyarrow as pa + + batches = [ + pa.RecordBatch.from_arrays([pa.array([1, 2])], ["x"]), + pa.RecordBatch.from_arrays([pa.array([3])], ["x"]), + ] + result = ArrowBatchTransformer.concat_batches(batches) + self.assertEqual(result.column(0).to_pylist(), [1, 2, 3]) + self.assertIs(ArrowBatchTransformer.concat_batches(batches[:1]), batches[0]) + + def test_flatten_elementwise_inputs_and_renest_outputs(self): + import pyarrow as pa + + int_values = pa.array( + [[[1, 2], []], None, [[3], None]], + type=pa.large_list(pa.list_(pa.int64())), + ) + string_values = pa.array( + [[["a", "b"], []], None, [["c"], None]], + type=pa.large_list(pa.list_(pa.string())), + ) + batch = pa.RecordBatch.from_arrays([int_values, string_values], ["ints", "strings"]) + + flat, shape_levels, is_large_levels = ArrowBatchTransformer.flatten_elementwise_inputs( + batch, [0, 1], depth=2 + ) + self.assertEqual(flat.schema.names, ["_0", "_1"]) + self.assertEqual(flat.column(0).to_pylist(), [1, 2, 3]) + self.assertEqual(flat.column(1).to_pylist(), ["a", "b", "c"]) + self.assertEqual(shape_levels, [[2, None, 2], [2, 0, 1, None]]) + self.assertEqual(is_large_levels, [True, False]) + + restored = ArrowBatchTransformer.renest_elementwise_outputs( + [ + ( + pa.RecordBatch.from_arrays([flat.column(0)], ["_0"]), + shape_levels, + is_large_levels, + ), + ( + pa.RecordBatch.from_arrays([flat.column(1)], ["_0"]), + shape_levels, + is_large_levels, + ), + ], + ["ints", "strings"], + ) + self.assertEqual(restored.schema.names, ["ints", "strings"]) + self.assertTrue(restored.column(0).equals(int_values)) + self.assertTrue(restored.column(1).equals(string_values)) + def test_wrap_struct_basic(self): """Test wrapping columns into a struct.""" import pyarrow as pa diff --git a/python/pyspark/worker.py b/python/pyspark/worker.py index bf8117a13174c..2aceb328c5904 100644 --- a/python/pyspark/worker.py +++ b/python/pyspark/worker.py @@ -1644,38 +1644,10 @@ def mapper(_, it): return mapper, ser -def _elementwise_renest(flat_values, shape_lengths, is_large): - """Re-nest a flat Array of per-element results into an ``array`` 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): +def _elementwise_udf_input_type(data_type, depth): """The element type ``depth`` ``ArrayType`` levels below ``data_type``. - A lifted UDF's argument arrives as ``array^depth`` (one ``array`` level per enclosing higher- + A lifted UDF's input arrives as ``array^depth`` (one ``array`` level per enclosing higher- order function lambda); this peels them off to the scalar leaf ``T`` the user function sees. See ``ExtractPythonUDFFromLambda``. """ @@ -1684,94 +1656,73 @@ def _elementwise_leaf_type(data_type, depth): return data_type -def _elementwise_flatten_deep(col, depth): - """Flatten ``depth`` list levels off ``col``, keeping each level's shape for re-nesting. +def _elementwise_flat_batch_to_pandas_or_arrow_udf_inputs( + flat_batch, input_schema, is_pandas, runner_conf +): + """Adapt one flattened input batch to a pandas or Arrow element-wise UDF's inputs. - 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 + ``flat_batch`` contains the aligned leaf Arrays produced once per input batch by + ``ArrowBatchTransformer.flatten_elementwise_inputs``. The Arrow flavor receives those Arrays + unchanged. The pandas flavor converts the entire batch to Series / DataFrames using the leaf + ``input_schema`` and the same options as the ordinary scalar pandas UDF path. + The scalar pandas / Arrow UDFs and their iterator variants receive whole pandas Series / + DataFrames or ``pa.Array`` objects. The row-at-a-time ``SQL_ARROW_ELEMENTWISE_UDF`` follows the + same element-wise flatten/invoke/re-nest path, but converts the flat Arrays to Python values and + invokes the function once per element tuple, so it does not use this adapter. -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`` 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`` 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``. + ``flatten_elementwise_inputs`` assigns positional ``_N`` field names only to form a + RecordBatch. Clear them from Series results so this batch wrapper does not expose new names to + user code; struct inputs remain DataFrames with their real child-field names. Shared by the + scalar and iterator element-wise worker paths that back Python UDFs in higher-order-function + lambdas. """ if not is_pandas: - return flat + return flat_batch.columns + + import pandas as pd - return ArrowToPandasConversion._convert_array( - flat, - element_type, + results = ArrowToPandasConversion.to_pandas( + flat_batch, timezone=runner_conf.timezone, + schema=input_schema, 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``. + for result in results: + if isinstance(result, pd.Series): + # Do not expose synthetic flattened-batch field names to the UDF. + result.name = None + return results + + +def _elementwise_pandas_or_arrow_udf_output_to_flat_batch( + output, return_type, output_schema, is_pandas, runner_conf +): + """Convert one pandas or Arrow UDF result over flat elements to a one-column Arrow batch. + + ``output`` is a pandas Series / DataFrame (pandas flavor) or a ``pa.Array`` (Arrow flavor); the + returned batch holds one row per input element. The Arrow flavor is coerced to the UTC-typed + ``output_schema``. The pandas flavor is typed by ``PandasToArrowConversion`` using the session + timezone, so its timestamp type may differ from ``output_schema``; iterator callers must take + the schema from the returned batch rather than assume UTC when buffering chunks. + + This adapter is only for the scalar and iterator pandas / Arrow UDF paths. The row-at-a-time + ``SQL_ARROW_ELEMENTWISE_UDF`` shares their element-wise flattening and re-nesting, but collects + ordinary Python result values and converts them with its Python-value converter instead. + + Keep each UDF result as a one-column RecordBatch so the worker can pass it directly to + ``ArrowBatchTransformer.renest_elementwise_outputs``. The transformer extracts the flat Array, + rebuilds its list levels, and assembles the final output batch. Shared by the scalar and + iterator pandas / Arrow element-wise worker paths. See ``ExtractPythonUDFFromLambda``. """ import pyarrow as pa if is_pandas: - batch = PandasToArrowConversion.from_pandas( - [result], + return PandasToArrowConversion.from_pandas( + [output], StructType([StructField("_0", return_type)]), timezone=runner_conf.timezone, safecheck=runner_conf.safecheck, @@ -1781,14 +1732,11 @@ def _elementwise_result_to_arrow(result, return_type, arrow_element_type, is_pan 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)]), + return ArrowBatchTransformer.enforce_schema( + pa.RecordBatch.from_arrays([output], names=output_schema.names), + output_schema, 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): @@ -2803,7 +2751,7 @@ def func(split_index: int, data: Iterator[pa.RecordBatch]) -> Iterator[pa.Record # ``T`` reached by peeling ``depth`` array levels. arg_converters = [ ArrowTableToRowsConversion._create_converter( - _elementwise_leaf_type(input_fields[o].dataType, depth), + _elementwise_udf_input_type(input_fields[o].dataType, depth), none_on_identity=True, binary_as_bytes=runner_conf.binary_as_bytes, ) @@ -2815,8 +2763,8 @@ def func(split_index: int, data: Iterator[pa.RecordBatch]) -> Iterator[pa.Record 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. + # The UDF returns one value per element, so its Arrow element type is the + # declared return type rather than the surrounding array operator type. to_arrow_type( udf_return_type, timezone="UTC", @@ -2846,7 +2794,7 @@ def func(split_index: int, data: Iterator[pa.RecordBatch]) -> Iterator[pa.Record # 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 = [] + output_batches = [] for info in udf_infos: ( wrapped_func, @@ -2856,48 +2804,46 @@ def func(split_index: int, data: Iterator[pa.RecordBatch]) -> Iterator[pa.Record 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 + flat_batch, shape_levels, is_large_levels = ( + ArrowBatchTransformer.flatten_elementwise_inputs( + input_batch, offsets, 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) + for column, conv in zip(flat_batch.columns, arg_converters): + values = ArrowTableToRowsConversion._to_pylist(column) if conv is not None: values = [conv(v) for v in values] columns.append(values) - total_elements = len(columns[0]) + total_elements = flat_batch.num_rows # 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 using that UDF's offsets. + # Convert the flat results before re-nesting them with this UDF's own shapes. 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. + # Broader than SQL_ARROW_BATCHED_UDF, which catches only ArrowInvalid: an + # element-wise UDF commonly returns list/struct values whose mismatches surface + # as ArrowTypeError, so both errors use the explicit cast fallback here. 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) + output_batches.append( + ( + pa.RecordBatch.from_arrays([flat_arr], ["_0"]), + shape_levels, + is_large_levels, + ) ) - yield pa.RecordBatch.from_arrays(output_arrays, col_names) + yield ArrowBatchTransformer.renest_elementwise_outputs(output_batches, col_names) return func, ser @@ -2913,9 +2859,9 @@ def func(split_index: int, data: Iterator[pa.RecordBatch]) -> Iterator[pa.Record # A scalar pandas or Arrow UDF lifted out of a higher-order function's lambda by # ExtractPythonUDFFromLambda. Each argument arrives as ``array`` 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`` using the + # iterated array. We flatten each list column to its element column, run the pandas or Arrow + # 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`` 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 @@ -2929,23 +2875,37 @@ def func(split_index: int, data: Iterator[pa.RecordBatch]) -> Iterator[pa.Record ) # The UDF returns one value per element, so its declared return type is the element # type of the ``array`` this operator produces. - arrow_element_type = to_arrow_type( - udf_return_type, timezone="UTC", prefers_large_types=runner_conf.use_large_var_types + udf_output_schema = pa.schema( + [ + pa.field( + "_0", + 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``; 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_input_schema = StructType( + [ + StructField( + f"_{i}", _elementwise_udf_input_type(input_fields[o].dataType, depth) + ) + for i, o in enumerate(args_kwargs_offsets) + ] + ) udf_infos.append( ( wrapped_func, args_kwargs_offsets, udf_return_type, - arrow_element_type, + udf_output_schema, depth, - arg_leaf_types, + udf_input_schema, ) ) col_names = [f"_{i}" for i in range(len(udfs))] @@ -2955,36 +2915,30 @@ def func(split_index: int, data: Iterator[pa.RecordBatch]) -> Iterator[pa.Record def func(split_index: int, data: Iterator[pa.RecordBatch]) -> Iterator[pa.RecordBatch]: for input_batch in data: - output_arrays = [] + output_batches = [] for ( wrapped_func, offsets, return_type, - arrow_element_type, + udf_output_schema, depth, - arg_leaf_types, + udf_input_schema, ) 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, + # pandas or Arrow function'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 shapes drive the re-nest). + flat_batch, shape_levels, is_large_levels = ( + ArrowBatchTransformer.flatten_elementwise_inputs( + input_batch, offsets, depth ) - for i, (o, t) in enumerate(zip(offsets, arg_leaf_types)) - ] + ) + total_elements = flat_batch.num_rows + udf_inputs = _elementwise_flat_batch_to_pandas_or_arrow_udf_inputs( + flat_batch, udf_input_schema, is_pandas, runner_conf + ) - result = wrapped_func(*flat_columns) + result = wrapped_func(*udf_inputs) if is_pandas: if not hasattr(result, "__len__"): pd_type = ( @@ -3016,13 +2970,12 @@ def func(split_index: int, data: Iterator[pa.RecordBatch]) -> Iterator[pa.Record # 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 + flat_output_batch = _elementwise_pandas_or_arrow_udf_output_to_flat_batch( + result, return_type, udf_output_schema, is_pandas, runner_conf ) - nested = _elementwise_renest_deep(flat_arr, shape_levels, is_large_levels) - output_arrays.append(nested) + output_batches.append((flat_output_batch, shape_levels, is_large_levels)) - yield pa.RecordBatch.from_arrays(output_arrays, col_names) + yield ArrowBatchTransformer.renest_elementwise_outputs(output_batches, col_names) return func, ser @@ -3059,11 +3012,23 @@ def func(split_index: int, data: Iterator[pa.RecordBatch]) -> Iterator[pa.Record 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 + udf_input_schema = StructType( + [ + StructField(f"_{i}", _elementwise_udf_input_type(input_fields[o].dataType, depth)) + for i, o in enumerate(args_offsets) + ] + ) + udf_output_schema = pa.schema( + [ + pa.field( + "_0", + 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]: @@ -3078,21 +3043,15 @@ def extract_flat(batch: pa.RecordBatch): # 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 + flat_batch, shape_levels, is_large_levels = ( + ArrowBatchTransformer.flatten_elementwise_inputs(batch, args_offsets, 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) + pending_shapes.append((shape_levels, is_large_levels, flat_batch.num_rows)) + num_input_elements += flat_batch.num_rows + udf_inputs = _elementwise_flat_batch_to_pandas_or_arrow_udf_inputs( + flat_batch, udf_input_schema, is_pandas, runner_conf + ) + return udf_inputs[0] if len(udf_inputs) == 1 else tuple(udf_inputs) flat_args_iter = map(extract_flat, data) @@ -3116,14 +3075,14 @@ def extract_flat(batch: pa.RecordBatch): # 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 + # buffer. ``empty_schema`` supplies the type for a zero-length emit; it tracks the most + # recent chunk's schema (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. + # UTC-typed ``udf_output_schema`` only before any chunk arrives, so all emitted batches + # use one schema. pending_chunks: "list" = [] pending_len = 0 - empty_type = arrow_element_type + empty_schema = udf_output_schema num_output_elements = 0 def emit_ready(): @@ -3134,22 +3093,19 @@ def emit_ready(): break pending_shapes.popleft() if needed == 0: - flat = pa.nulls(0, type=empty_type) + flat_batch = pa.RecordBatch.from_pylist([], schema=empty_schema) else: - combined = ( - pending_chunks[0] - if len(pending_chunks) == 1 - else pa.concat_arrays(pending_chunks) - ) - flat = combined.slice(0, needed) + combined = ArrowBatchTransformer.concat_batches(pending_chunks) + flat_batch = combined.slice(0, needed) remainder = combined.slice(needed) - pending_chunks = [remainder] if len(remainder) else [] + pending_chunks = [remainder] if remainder.num_rows else [] pending_len -= needed - nested = _elementwise_renest_deep(flat, shape_levels, is_large_levels) - yield pa.RecordBatch.from_arrays([nested], ["_0"]) + yield ArrowBatchTransformer.renest_elementwise_outputs( + [(flat_batch, shape_levels, is_large_levels)], ["_0"] + ) def process_results(): - nonlocal pending_chunks, pending_len, empty_type, num_output_elements + nonlocal pending_chunks, pending_len, empty_schema, num_output_elements for result in verified_iter: if is_pandas: verify_pandas_result( @@ -3158,10 +3114,10 @@ def process_results(): assign_cols_by_name=True, truncate_return_schema=True, ) - chunk = _elementwise_result_to_arrow( - result, return_type, arrow_element_type, is_pandas, runner_conf + chunk = _elementwise_pandas_or_arrow_udf_output_to_flat_batch( + result, return_type, udf_output_schema, is_pandas, runner_conf ) - num_output_elements += len(chunk) + num_output_elements += chunk.num_rows # 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: @@ -3173,10 +3129,10 @@ def process_results(): # 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): + empty_schema = chunk.schema + if chunk.num_rows: pending_chunks.append(chunk) - pending_len += len(chunk) + pending_len += chunk.num_rows yield from emit_ready() # The iterator is exhausted: every input row's flat elements must have arrived. From 8932b05ac17d8d8016576fc88287e8d46d334691 Mon Sep 17 00:00:00 2001 From: Spenser Sun Date: Mon, 21 Sep 2026 21:26:16 +0000 Subject: [PATCH 2/5] [SPARK-59624][PYTHON] Refine element-wise batch interface helpers --- python/pyspark/sql/conversion.py | 96 +----------- python/pyspark/sql/tests/test_conversion.py | 41 ----- python/pyspark/worker.py | 157 ++++++++++++++------ 3 files changed, 117 insertions(+), 177 deletions(-) diff --git a/python/pyspark/sql/conversion.py b/python/pyspark/sql/conversion.py index aed8617f5b345..3633c1d1a69ba 100644 --- a/python/pyspark/sql/conversion.py +++ b/python/pyspark/sql/conversion.py @@ -180,14 +180,13 @@ def select_columns(cls, batch: "pa.RecordBatch", column_indices: list[int]) -> " [batch.schema.names[i] for i in column_indices], ) - @staticmethod - def concat_batches(batches: Sequence["pa.RecordBatch"]) -> "pa.RecordBatch": + @classmethod + def concat_batches(cls, batches: Sequence["pa.RecordBatch"]) -> "pa.RecordBatch": """Concatenate same-schema RecordBatches by row. A single batch is returned unchanged. PyArrow before 19.0.0 has no ``concat_batches``; the fallback concatenates the equivalent StructArrays and converts the result back to a - RecordBatch. Element-wise iterator UDFs use this when one input batch's flattened result - spans multiple output chunks. + RecordBatch. """ import pyarrow as pa @@ -200,95 +199,6 @@ def concat_batches(batches: Sequence["pa.RecordBatch"]) -> "pa.RecordBatch": pa.concat_arrays([batch.to_struct_array() for batch in batches]) ) - @staticmethod - def flatten_elementwise_inputs( - batch: "pa.RecordBatch", input_column_indices: Sequence[int], depth: int - ) -> tuple["pa.RecordBatch", list[list[Optional[int]]], list[bool]]: - """Flatten ``depth`` list levels from an element-wise UDF's input columns. - - Returns ``(flat_input_batch, shape_levels, is_large_levels)``. ``flat_input_batch`` - contains each selected input's fully flattened leaf Array under a positional ``_N`` name. - ``shape_levels[k]`` contains the per-slot list length at level ``k`` (0 is outermost), - using ``None`` for a null list. ``is_large_levels[k]`` records whether that level uses - ``LargeListArray`` and therefore requires int64 rather than int32 offsets when rebuilt. - - Only the first selected column supplies shape and list-width metadata. The other inputs are - aligned to it by ``ExtractPythonUDFFromLambda``, so recording their shapes would repeat - the ``list_value_length(...).to_pylist()`` work without changing re-nesting. ``depth`` is 1 - for a UDF in one higher-order-function lambda and greater for nested lambdas. - - Shared by the row, scalar pandas / Arrow, and iterator element-wise worker paths. See - ``ExtractPythonUDFFromLambda``. - """ - import pyarrow as pa - import pyarrow.compute as pc - - assert input_column_indices - assert depth > 0 - - flat_inputs = [] - shape_levels = [] - is_large_levels = [] - for input_index, column_index in enumerate(input_column_indices): - current = batch.column(column_index) - for _ in range(depth): - if input_index == 0: - shape_levels.append(pc.list_value_length(current).to_pylist()) - is_large_levels.append(pa.types.is_large_list(current.type)) - current = current.flatten() - flat_inputs.append(current) - - return ( - pa.RecordBatch.from_arrays( - flat_inputs, names=[f"_{index}" for index in range(len(flat_inputs))] - ), - shape_levels, - is_large_levels, - ) - - @staticmethod - def renest_elementwise_outputs( - flat_outputs: Sequence[tuple["pa.RecordBatch", list[list[Optional[int]]], list[bool]]], - column_names: Sequence[str], - ) -> "pa.RecordBatch": - """Rebuild nested list columns from flattened element-wise UDF result batches. - - Each input tuple contains a one-column flat result batch plus the ``shape_levels`` and - ``is_large_levels`` returned by ``flatten_elementwise_inputs`` for that UDF. Levels are - rebuilt from innermost to outermost. A ``None`` length creates a null list and consumes no - flat values; a zero length creates an empty, non-null list. ``is_large_levels`` preserves - each input level's int32 ``ListArray`` versus int64 ``LargeListArray`` offset width. - - Different fused UDFs may carry different shapes, so every result batch is rebuilt with its - own metadata before the columns are assembled into one output RecordBatch. This is the - batch-level inverse of ``flatten_elementwise_inputs``. - """ - import pyarrow as pa - - assert len(flat_outputs) == len(column_names) - nested_columns = [] - for flat_batch, shape_levels, is_large_levels in flat_outputs: - assert flat_batch.num_columns == 1 - result = flat_batch.column(0) - for shape_lengths, is_large in zip(reversed(shape_levels), reversed(is_large_levels)): - offsets = [0] - running = 0 - nulls = [] - for length in shape_lengths: - nulls.append(length is None) - if length is not None: - running += length - offsets.append(running) - list_type = pa.LargeListArray if is_large else pa.ListArray - result = list_type.from_arrays( - pa.array(offsets, type=pa.int64() if is_large else pa.int32()), - result, - mask=pa.array(nulls, type=pa.bool_()), - ) - nested_columns.append(result) - - return pa.RecordBatch.from_arrays(nested_columns, names=column_names) - @staticmethod def wrap_struct(batch: "pa.RecordBatch") -> "pa.RecordBatch": """ diff --git a/python/pyspark/sql/tests/test_conversion.py b/python/pyspark/sql/tests/test_conversion.py index 670109816f4f6..549fa39c58dc5 100644 --- a/python/pyspark/sql/tests/test_conversion.py +++ b/python/pyspark/sql/tests/test_conversion.py @@ -126,47 +126,6 @@ def test_concat_batches(self): self.assertEqual(result.column(0).to_pylist(), [1, 2, 3]) self.assertIs(ArrowBatchTransformer.concat_batches(batches[:1]), batches[0]) - def test_flatten_elementwise_inputs_and_renest_outputs(self): - import pyarrow as pa - - int_values = pa.array( - [[[1, 2], []], None, [[3], None]], - type=pa.large_list(pa.list_(pa.int64())), - ) - string_values = pa.array( - [[["a", "b"], []], None, [["c"], None]], - type=pa.large_list(pa.list_(pa.string())), - ) - batch = pa.RecordBatch.from_arrays([int_values, string_values], ["ints", "strings"]) - - flat, shape_levels, is_large_levels = ArrowBatchTransformer.flatten_elementwise_inputs( - batch, [0, 1], depth=2 - ) - self.assertEqual(flat.schema.names, ["_0", "_1"]) - self.assertEqual(flat.column(0).to_pylist(), [1, 2, 3]) - self.assertEqual(flat.column(1).to_pylist(), ["a", "b", "c"]) - self.assertEqual(shape_levels, [[2, None, 2], [2, 0, 1, None]]) - self.assertEqual(is_large_levels, [True, False]) - - restored = ArrowBatchTransformer.renest_elementwise_outputs( - [ - ( - pa.RecordBatch.from_arrays([flat.column(0)], ["_0"]), - shape_levels, - is_large_levels, - ), - ( - pa.RecordBatch.from_arrays([flat.column(1)], ["_0"]), - shape_levels, - is_large_levels, - ), - ], - ["ints", "strings"], - ) - self.assertEqual(restored.schema.names, ["ints", "strings"]) - self.assertTrue(restored.column(0).equals(int_values)) - self.assertTrue(restored.column(1).equals(string_values)) - def test_wrap_struct_basic(self): """Test wrapping columns into a struct.""" import pyarrow as pa diff --git a/python/pyspark/worker.py b/python/pyspark/worker.py index 2aceb328c5904..fc3b27d033229 100644 --- a/python/pyspark/worker.py +++ b/python/pyspark/worker.py @@ -1656,26 +1656,94 @@ def _elementwise_udf_input_type(data_type, depth): return data_type +def _elementwise_flatten_inputs(batch, input_column_indices, depth): + """Flatten element-wise UDF inputs and capture the first input's list shape. + + Returns ``(flat_batch, list_lengths_by_level, large_list_by_level)``. ``flat_batch`` contains + one fully flattened leaf Array per selected input under a positional ``_N`` name. + ``list_lengths_by_level[k]`` contains the per-slot length at level ``k`` (0 is outermost), using + ``None`` for a null list. ``large_list_by_level[k]`` records whether that level uses + ``LargeListArray``. + + Only the first selected input supplies shape metadata because + ``ExtractPythonUDFFromLambda`` aligns the other inputs to it. ``depth`` is 1 for a UDF in a + single lambda and greater for one lifted out of nested lambdas. + """ + import pyarrow as pa + import pyarrow.compute as pc + + assert input_column_indices + assert depth > 0 + + flat_inputs = [] + list_lengths_by_level = [] + large_list_by_level = [] + for input_index, column_index in enumerate(input_column_indices): + current = batch.column(column_index) + for _ in range(depth): + if input_index == 0: + list_lengths_by_level.append(pc.list_value_length(current).to_pylist()) + large_list_by_level.append(pa.types.is_large_list(current.type)) + current = current.flatten() + flat_inputs.append(current) + + return ( + pa.RecordBatch.from_arrays( + flat_inputs, names=[f"_{index}" for index in range(len(flat_inputs))] + ), + list_lengths_by_level, + large_list_by_level, + ) + + +def _elementwise_renest_output(flat_values, list_lengths_by_level, large_list_by_level): + """Re-nest one element-wise UDF output using its first input's recorded shape. + + Levels are rebuilt from innermost to outermost. A ``None`` length creates a null list and + consumes no flat values; zero creates an empty, non-null list. ``large_list_by_level`` preserves + each input level's int32 ``ListArray`` versus int64 ``LargeListArray`` offset width. + """ + import pyarrow as pa + + result = flat_values + for list_lengths, use_large_list in zip( + reversed(list_lengths_by_level), reversed(large_list_by_level) + ): + offsets = [0] + running = 0 + nulls = [] + for length in list_lengths: + nulls.append(length is None) + if length is not None: + running += length + offsets.append(running) + list_type = pa.LargeListArray if use_large_list else pa.ListArray + result = list_type.from_arrays( + pa.array(offsets, type=pa.int64() if use_large_list else pa.int32()), + result, + mask=pa.array(nulls, type=pa.bool_()), + ) + return result + + def _elementwise_flat_batch_to_pandas_or_arrow_udf_inputs( flat_batch, input_schema, is_pandas, runner_conf ): """Adapt one flattened input batch to a pandas or Arrow element-wise UDF's inputs. - ``flat_batch`` contains the aligned leaf Arrays produced once per input batch by - ``ArrowBatchTransformer.flatten_elementwise_inputs``. The Arrow flavor receives those Arrays - unchanged. The pandas flavor converts the entire batch to Series / DataFrames using the leaf - ``input_schema`` and the same options as the ordinary scalar pandas UDF path. + ``flat_batch`` contains one aligned leaf Array per UDF argument. The Arrow flavor receives those + Arrays unchanged. The pandas flavor converts the entire batch to Series / DataFrames using the + leaf ``input_schema`` and the same options as the ordinary scalar pandas UDF path. The scalar pandas / Arrow UDFs and their iterator variants receive whole pandas Series / DataFrames or ``pa.Array`` objects. The row-at-a-time ``SQL_ARROW_ELEMENTWISE_UDF`` follows the same element-wise flatten/invoke/re-nest path, but converts the flat Arrays to Python values and invokes the function once per element tuple, so it does not use this adapter. - ``flatten_elementwise_inputs`` assigns positional ``_N`` field names only to form a - RecordBatch. Clear them from Series results so this batch wrapper does not expose new names to - user code; struct inputs remain DataFrames with their real child-field names. Shared by the - scalar and iterator element-wise worker paths that back Python UDFs in higher-order-function - lambdas. + Positional ``_N`` field names exist only to form the RecordBatch. Clear them from Series results + so this batch wrapper does not expose new names to user code; struct inputs remain DataFrames + with their real child-field names. Shared by the scalar and iterator element-wise worker paths + that back Python UDFs in higher-order-function lambdas. """ if not is_pandas: return flat_batch.columns @@ -1713,10 +1781,9 @@ def _elementwise_pandas_or_arrow_udf_output_to_flat_batch( ``SQL_ARROW_ELEMENTWISE_UDF`` shares their element-wise flattening and re-nesting, but collects ordinary Python result values and converts them with its Python-value converter instead. - Keep each UDF result as a one-column RecordBatch so the worker can pass it directly to - ``ArrowBatchTransformer.renest_elementwise_outputs``. The transformer extracts the flat Array, - rebuilds its list levels, and assembles the final output batch. Shared by the scalar and - iterator pandas / Arrow element-wise worker paths. See ``ExtractPythonUDFFromLambda``. + Keep each UDF result as a one-column RecordBatch so iterator paths can concatenate and slice + output chunks while retaining their schema. Shared by the scalar and iterator pandas / Arrow + element-wise worker paths. See ``ExtractPythonUDFFromLambda``. """ import pyarrow as pa @@ -2763,8 +2830,8 @@ def func(split_index: int, data: Iterator[pa.RecordBatch]) -> Iterator[pa.Record args_kwargs_offsets, depth, arg_converters, - # The UDF returns one value per element, so its Arrow element type is the - # declared return type rather than the surrounding array operator type. + # 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", @@ -2794,7 +2861,7 @@ def func(split_index: int, data: Iterator[pa.RecordBatch]) -> Iterator[pa.Record # 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_batches = [] + output_arrays = [] for info in udf_infos: ( wrapped_func, @@ -2804,10 +2871,10 @@ def func(split_index: int, data: Iterator[pa.RecordBatch]) -> Iterator[pa.Record arrow_element_type, result_conv, ) = info - flat_batch, shape_levels, is_large_levels = ( - ArrowBatchTransformer.flatten_elementwise_inputs( - input_batch, offsets, depth - ) + # Flatten each argument `depth` list levels to its leaves; the first argument's + # per-level shapes drive the re-nest. + flat_batch, list_lengths_by_level, large_list_by_level = ( + _elementwise_flatten_inputs(input_batch, offsets, depth) ) columns = [] for column, conv in zip(flat_batch.columns, arg_converters): @@ -2822,28 +2889,27 @@ def func(split_index: int, data: Iterator[pa.RecordBatch]) -> Iterator[pa.Record results = _evaluate_elementwise_udf(wrapped_func, rows) verify_result_row_count(len(results), total_elements) - # Convert the flat results before re-nesting them with this UDF's own shapes. + # Convert results and re-nest to array 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 SQL_ARROW_BATCHED_UDF, which catches only ArrowInvalid: an - # element-wise UDF commonly returns list/struct values whose mismatches surface - # as ArrowTypeError, so both errors use the explicit cast fallback here. + # 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_batches.append( - ( - pa.RecordBatch.from_arrays([flat_arr], ["_0"]), - shape_levels, - is_large_levels, + output_arrays.append( + _elementwise_renest_output( + flat_arr, list_lengths_by_level, large_list_by_level ) ) - yield ArrowBatchTransformer.renest_elementwise_outputs(output_batches, col_names) + yield pa.RecordBatch.from_arrays(output_arrays, col_names) return func, ser @@ -2915,7 +2981,7 @@ def func(split_index: int, data: Iterator[pa.RecordBatch]) -> Iterator[pa.Record def func(split_index: int, data: Iterator[pa.RecordBatch]) -> Iterator[pa.RecordBatch]: for input_batch in data: - output_batches = [] + output_arrays = [] for ( wrapped_func, offsets, @@ -2928,10 +2994,8 @@ def func(split_index: int, data: Iterator[pa.RecordBatch]) -> Iterator[pa.Record # pandas or Arrow function'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 shapes drive the re-nest). - flat_batch, shape_levels, is_large_levels = ( - ArrowBatchTransformer.flatten_elementwise_inputs( - input_batch, offsets, depth - ) + flat_batch, list_lengths_by_level, large_list_by_level = ( + _elementwise_flatten_inputs(input_batch, offsets, depth) ) total_elements = flat_batch.num_rows udf_inputs = _elementwise_flat_batch_to_pandas_or_arrow_udf_inputs( @@ -2973,9 +3037,13 @@ def func(split_index: int, data: Iterator[pa.RecordBatch]) -> Iterator[pa.Record flat_output_batch = _elementwise_pandas_or_arrow_udf_output_to_flat_batch( result, return_type, udf_output_schema, is_pandas, runner_conf ) - output_batches.append((flat_output_batch, shape_levels, is_large_levels)) + output_arrays.append( + _elementwise_renest_output( + flat_output_batch.column(0), list_lengths_by_level, large_list_by_level + ) + ) - yield ArrowBatchTransformer.renest_elementwise_outputs(output_batches, col_names) + yield pa.RecordBatch.from_arrays(output_arrays, col_names) return func, ser @@ -3043,10 +3111,12 @@ def extract_flat(batch: pa.RecordBatch): # 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. - flat_batch, shape_levels, is_large_levels = ( - ArrowBatchTransformer.flatten_elementwise_inputs(batch, args_offsets, depth) + flat_batch, list_lengths_by_level, large_list_by_level = ( + _elementwise_flatten_inputs(batch, args_offsets, depth) + ) + pending_shapes.append( + (list_lengths_by_level, large_list_by_level, flat_batch.num_rows) ) - pending_shapes.append((shape_levels, is_large_levels, flat_batch.num_rows)) num_input_elements += flat_batch.num_rows udf_inputs = _elementwise_flat_batch_to_pandas_or_arrow_udf_inputs( flat_batch, udf_input_schema, is_pandas, runner_conf @@ -3088,7 +3158,7 @@ def extract_flat(batch: pa.RecordBatch): def emit_ready(): nonlocal pending_chunks, pending_len while pending_shapes: - shape_levels, is_large_levels, needed = pending_shapes[0] + list_lengths_by_level, large_list_by_level, needed = pending_shapes[0] if needed > pending_len: break pending_shapes.popleft() @@ -3100,9 +3170,10 @@ def emit_ready(): remainder = combined.slice(needed) pending_chunks = [remainder] if remainder.num_rows else [] pending_len -= needed - yield ArrowBatchTransformer.renest_elementwise_outputs( - [(flat_batch, shape_levels, is_large_levels)], ["_0"] + nested = _elementwise_renest_output( + flat_batch.column(0), list_lengths_by_level, large_list_by_level ) + yield pa.RecordBatch.from_arrays([nested], ["_0"]) def process_results(): nonlocal pending_chunks, pending_len, empty_schema, num_output_elements From 1b7907df02de598f1f5fee3e734b5470a2b7d003 Mon Sep 17 00:00:00 2001 From: Spenser Sun Date: Mon, 21 Sep 2026 23:58:15 +0000 Subject: [PATCH 3/5] [SPARK-59624][PYTHON] Refine batch helper contracts --- python/pyspark/sql/conversion.py | 4 ++- python/pyspark/sql/tests/test_conversion.py | 4 +-- python/pyspark/worker.py | 28 +++++++++++++++------ 3 files changed, 26 insertions(+), 10 deletions(-) diff --git a/python/pyspark/sql/conversion.py b/python/pyspark/sql/conversion.py index 3633c1d1a69ba..7d99e574f06bf 100644 --- a/python/pyspark/sql/conversion.py +++ b/python/pyspark/sql/conversion.py @@ -24,6 +24,7 @@ TYPE_CHECKING, Any, Callable, + Iterable, Iterator, List, Optional, @@ -181,7 +182,7 @@ def select_columns(cls, batch: "pa.RecordBatch", column_indices: list[int]) -> " ) @classmethod - def concat_batches(cls, batches: Sequence["pa.RecordBatch"]) -> "pa.RecordBatch": + def concat_batches(cls, batches: Iterable["pa.RecordBatch"]) -> "pa.RecordBatch": """Concatenate same-schema RecordBatches by row. A single batch is returned unchanged. PyArrow before 19.0.0 has no ``concat_batches``; @@ -190,6 +191,7 @@ def concat_batches(cls, batches: Sequence["pa.RecordBatch"]) -> "pa.RecordBatch" """ import pyarrow as pa + batches = tuple(batches) assert batches if len(batches) == 1: return batches[0] diff --git a/python/pyspark/sql/tests/test_conversion.py b/python/pyspark/sql/tests/test_conversion.py index 549fa39c58dc5..ba2fb4060a79e 100644 --- a/python/pyspark/sql/tests/test_conversion.py +++ b/python/pyspark/sql/tests/test_conversion.py @@ -122,9 +122,9 @@ def test_concat_batches(self): pa.RecordBatch.from_arrays([pa.array([1, 2])], ["x"]), pa.RecordBatch.from_arrays([pa.array([3])], ["x"]), ] - result = ArrowBatchTransformer.concat_batches(batches) + result = ArrowBatchTransformer.concat_batches(iter(batches)) self.assertEqual(result.column(0).to_pylist(), [1, 2, 3]) - self.assertIs(ArrowBatchTransformer.concat_batches(batches[:1]), batches[0]) + self.assertIs(ArrowBatchTransformer.concat_batches(iter(batches[:1])), batches[0]) def test_wrap_struct_basic(self): """Test wrapping columns into a struct.""" diff --git a/python/pyspark/worker.py b/python/pyspark/worker.py index fc3b27d033229..f5ccee33007d5 100644 --- a/python/pyspark/worker.py +++ b/python/pyspark/worker.py @@ -35,6 +35,7 @@ Callable, Iterable, Optional, + Sequence, Tuple, Union, ) @@ -1644,7 +1645,7 @@ def mapper(_, it): return mapper, ser -def _elementwise_udf_input_type(data_type, depth): +def _elementwise_udf_input_type(data_type: DataType, depth: int) -> DataType: """The element type ``depth`` ``ArrayType`` levels below ``data_type``. A lifted UDF's input arrives as ``array^depth`` (one ``array`` level per enclosing higher- @@ -1656,7 +1657,9 @@ def _elementwise_udf_input_type(data_type, depth): return data_type -def _elementwise_flatten_inputs(batch, input_column_indices, depth): +def _elementwise_flatten_inputs( + batch: "pa.RecordBatch", input_column_indices: Sequence[int], depth: int +) -> tuple["pa.RecordBatch", list[list[Optional[int]]], list[bool]]: """Flatten element-wise UDF inputs and capture the first input's list shape. Returns ``(flat_batch, list_lengths_by_level, large_list_by_level)``. ``flat_batch`` contains @@ -1696,7 +1699,11 @@ def _elementwise_flatten_inputs(batch, input_column_indices, depth): ) -def _elementwise_renest_output(flat_values, list_lengths_by_level, large_list_by_level): +def _elementwise_renest_output( + flat_values: "pa.Array", + list_lengths_by_level: Sequence[Sequence[Optional[int]]], + large_list_by_level: Sequence[bool], +) -> "pa.Array": """Re-nest one element-wise UDF output using its first input's recorded shape. Levels are rebuilt from innermost to outermost. A ``None`` length creates a null list and @@ -1727,8 +1734,11 @@ def _elementwise_renest_output(flat_values, list_lengths_by_level, large_list_by def _elementwise_flat_batch_to_pandas_or_arrow_udf_inputs( - flat_batch, input_schema, is_pandas, runner_conf -): + flat_batch: "pa.RecordBatch", + input_schema: StructType, + is_pandas: bool, + runner_conf: RunnerConf, +) -> list[Union["pd.Series", "pd.DataFrame", "pa.Array"]]: """Adapt one flattened input batch to a pandas or Arrow element-wise UDF's inputs. ``flat_batch`` contains one aligned leaf Array per UDF argument. The Arrow flavor receives those @@ -1767,8 +1777,12 @@ def _elementwise_flat_batch_to_pandas_or_arrow_udf_inputs( def _elementwise_pandas_or_arrow_udf_output_to_flat_batch( - output, return_type, output_schema, is_pandas, runner_conf -): + output: Union["pd.Series", "pd.DataFrame", "pa.Array"], + return_type: DataType, + output_schema: "pa.Schema", + is_pandas: bool, + runner_conf: RunnerConf, +) -> "pa.RecordBatch": """Convert one pandas or Arrow UDF result over flat elements to a one-column Arrow batch. ``output`` is a pandas Series / DataFrame (pandas flavor) or a ``pa.Array`` (Arrow flavor); the From 1673914e64a8b29576677968bf35b4551230c7d4 Mon Sep 17 00:00:00 2001 From: Spenser Sun Date: Wed, 23 Sep 2026 03:08:15 +0000 Subject: [PATCH 4/5] [SPARK-59624][PYTHON] Validate element-wise helper invariants --- python/pyspark/worker.py | 5 ++++- 1 file changed, 4 insertions(+), 1 deletion(-) diff --git a/python/pyspark/worker.py b/python/pyspark/worker.py index f5ccee33007d5..802de7f0dee6a 100644 --- a/python/pyspark/worker.py +++ b/python/pyspark/worker.py @@ -1653,6 +1653,7 @@ def _elementwise_udf_input_type(data_type: DataType, depth: int) -> DataType: ``ExtractPythonUDFFromLambda``. """ for _ in range(depth): + assert isinstance(data_type, ArrayType) data_type = data_type.elementType return data_type @@ -1760,9 +1761,11 @@ def _elementwise_flat_batch_to_pandas_or_arrow_udf_inputs( import pandas as pd + timezone = runner_conf.timezone + assert timezone is not None results = ArrowToPandasConversion.to_pandas( flat_batch, - timezone=runner_conf.timezone, + timezone=timezone, schema=input_schema, struct_in_pandas="dict", ndarray_as_list=False, From 39f77b787a45e7e0f24e81d7d605d823ff6235c3 Mon Sep 17 00:00:00 2001 From: Spenser Sun Date: Wed, 23 Sep 2026 04:45:58 +0000 Subject: [PATCH 5/5] [SPARK-59624][PYTHON] Reuse batch transformer in Arrow handlers --- python/pyspark/eval_handlers/_arrow.py | 17 ++--------------- 1 file changed, 2 insertions(+), 15 deletions(-) diff --git a/python/pyspark/eval_handlers/_arrow.py b/python/pyspark/eval_handlers/_arrow.py index 21e10b68f9413..98ba58d1b061f 100644 --- a/python/pyspark/eval_handlers/_arrow.py +++ b/python/pyspark/eval_handlers/_arrow.py @@ -251,19 +251,6 @@ def run(self, split_index: int, data: Iterator[GroupedBatch]) -> Iterator[pa.Rec yield ArrowBatchTransformer.wrap_struct(batch) -def _concat_group_batches(batch_list: list["pa.RecordBatch"]) -> "pa.RecordBatch": - """Concatenate a group's RecordBatches into a single one, with a fallback for - pyarrow before 19.0.0 (which lacks ``pa.concat_batches``). Remove the fallback - once support for those versions is dropped.""" - import pyarrow as pa - - if hasattr(pa, "concat_batches"): - return pa.concat_batches(batch_list) - return pa.RecordBatch.from_struct_array( - pa.concat_arrays([b.to_struct_array() for b in batch_list]) - ) - - class ArrowGroupedAggUDFHandler(GroupedEvalTypeHandler["pa.RecordBatch"]): """SQL_GROUPED_AGG_ARROW_UDF: each UDF reduces its input columns over the whole group to a single scalar; emit one row per group with one column per UDF, @@ -290,7 +277,7 @@ def run(self, split_index: int, data: Iterator[GroupedBatch]) -> Iterator[pa.Rec batch_list = list(group) if not batch_list: continue - concatenated = _concat_group_batches(batch_list) + concatenated = ArrowBatchTransformer.concat_batches(batch_list) results = [ udf_func( *[concatenated.column(o) for o in args_offsets], @@ -369,7 +356,7 @@ def run(self, split_index: int, data: Iterator[GroupedBatch]) -> Iterator[pa.Rec batch_list = list(group) if not batch_list: continue - concatenated = _concat_group_batches(batch_list) + concatenated = ArrowBatchTransformer.concat_batches(batch_list) num_rows = concatenated.num_rows result_arrays = []