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 = [] diff --git a/python/pyspark/sql/conversion.py b/python/pyspark/sql/conversion.py index 55e009ce5c64e..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, @@ -180,6 +181,26 @@ def select_columns(cls, batch: "pa.RecordBatch", column_indices: list[int]) -> " [batch.schema.names[i] for i in column_indices], ) + @classmethod + 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``; + the fallback concatenates the equivalent StructArrays and converts the result back to a + RecordBatch. + """ + import pyarrow as pa + + batches = tuple(batches) + 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 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..ba2fb4060a79e 100644 --- a/python/pyspark/sql/tests/test_conversion.py +++ b/python/pyspark/sql/tests/test_conversion.py @@ -115,6 +115,17 @@ 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(iter(batches)) + self.assertEqual(result.column(0).to_pylist(), [1, 2, 3]) + self.assertIs(ArrowBatchTransformer.concat_batches(iter(batches[:1])), batches[0]) + 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..802de7f0dee6a 100644 --- a/python/pyspark/worker.py +++ b/python/pyspark/worker.py @@ -35,6 +35,7 @@ Callable, Iterable, Optional, + Sequence, Tuple, Union, ) @@ -1644,134 +1645,168 @@ 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: DataType, depth: int) -> DataType: """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``. """ for _ in range(depth): + assert isinstance(data_type, ArrayType) 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. +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 ``(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``. + 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 - 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 + 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_flatten_leaf(col, depth): - """Flatten ``depth`` list levels off ``col`` to its leaf ``pa.Array``, without capturing shape. +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. - 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``. + 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. """ - 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. + import pyarrow as pa - 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) + 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_flatten_column(flat, element_type, is_pandas, runner_conf): - """Adapt one already-flattened ``array`` element column to the vectorized fn's input. +def _elementwise_flat_batch_to_pandas_or_arrow_udf_inputs( + 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`` 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``. + ``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. + + 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 + return flat_batch.columns + + import pandas as pd - return ArrowToPandasConversion._convert_array( - flat, - element_type, - timezone=runner_conf.timezone, + timezone = runner_conf.timezone + assert timezone is not None + results = ArrowToPandasConversion.to_pandas( + flat_batch, + timezone=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, ) + 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_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``. +def _elementwise_pandas_or_arrow_udf_output_to_flat_batch( + 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 + 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 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 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 +1816,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 +2835,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, ) @@ -2858,22 +2890,17 @@ def func(split_index: int, data: Iterator[pa.RecordBatch]) -> Iterator[pa.Record ) = 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, list_lengths_by_level, large_list_by_level = ( + _elementwise_flatten_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) @@ -2894,7 +2921,9 @@ def func(split_index: int, data: Iterator[pa.RecordBatch]) -> Iterator[pa.Record target_type=arrow_element_type, safe=runner_conf.safecheck ) output_arrays.append( - _elementwise_renest_deep(flat_arr, shape_levels, is_large_levels) + _elementwise_renest_output( + flat_arr, list_lengths_by_level, large_list_by_level + ) ) yield pa.RecordBatch.from_arrays(output_arrays, col_names) @@ -2913,9 +2942,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 +2958,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))] @@ -2960,31 +3003,23 @@ def func(split_index: int, data: Iterator[pa.RecordBatch]) -> Iterator[pa.Record 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 + # 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, 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( + flat_batch, udf_input_schema, is_pandas, runner_conf ) - 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) + result = wrapped_func(*udf_inputs) if is_pandas: if not hasattr(result, "__len__"): pd_type = ( @@ -3016,11 +3051,14 @@ 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 + ) + output_arrays.append( + _elementwise_renest_output( + flat_output_batch.column(0), list_lengths_by_level, large_list_by_level + ) ) - nested = _elementwise_renest_deep(flat_arr, shape_levels, is_large_levels) - output_arrays.append(nested) yield pa.RecordBatch.from_arrays(output_arrays, col_names) @@ -3059,11 +3097,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 +3128,17 @@ 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, list_lengths_by_level, large_list_by_level = ( + _elementwise_flatten_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( + (list_lengths_by_level, large_list_by_level, 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,40 +3162,38 @@ 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(): 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() 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) + 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_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 +3202,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 +3217,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.