This is an automated email from the ASF dual-hosted git repository.

Yicong-Huang pushed a commit to branch branch-4.x
in repository https://gitbox.apache.org/repos/asf/spark.git


The following commit(s) were added to refs/heads/branch-4.x by this push:
     new e4e5e1f0e0a4 [SPARK-56757][PYTHON] Refactor SQL_SCALAR_PANDAS_ITER_UDF
e4e5e1f0e0a4 is described below

commit e4e5e1f0e0a432de6394edf4965b47df38d318eb
Author: Yicong Huang <[email protected]>
AuthorDate: Wed Jun 17 07:30:21 2026 +0000

    [SPARK-56757][PYTHON] Refactor SQL_SCALAR_PANDAS_ITER_UDF
    
    ### What changes were proposed in this pull request?
    
    Refactor `SQL_SCALAR_PANDAS_ITER_UDF` to use `ArrowStreamSerializer` as 
pure I/O, moving Arrow-to-Pandas iter conversion logic from 
`ArrowStreamPandasUDFSerializer` into `read_udfs()` in `worker.py`.
    
    With this change every eval type in the Arrow-stream family is on the plain 
`ArrowStreamSerializer`, so the now-dead `ArrowStreamPandasUDFSerializer` 
fallback branch in the serializer selection chain and the unused 
`wrap_pandas_batch_iter_udf` helper are removed as well.
    
    ### Why are the changes needed?
    
    Part of [SPARK-55388](https://issues.apache.org/jira/browse/SPARK-55388). 
Mirrors the pattern applied to `SQL_SCALAR_PANDAS_UDF` 
([SPARK-56648](https://issues.apache.org/jira/browse/SPARK-56648)) and 
`SQL_SCALAR_ARROW_ITER_UDF` 
([SPARK-55577](https://issues.apache.org/jira/browse/SPARK-55577)).
    
    ### Does this PR introduce _any_ user-facing change?
    
    No.
    
    ### How was this patch tested?
    
    Existing tests. No behavior change.
    
    ASV benchmark comparison (`ScalarPandasIterUDF` bench classes, `repeat=3`, 
median). before = `upstream/master`, after = this PR.
    
    **ScalarPandasIterUDFTimeBench** (latency):
    
    ```text
    scenario             udf              before    after     diff
    -------------------  --------------  -------  -------  -------
    sm_batch_few_col     identity_udf      414ms    387ms    -6.6%
    sm_batch_few_col     sort_udf          538ms    516ms    -4.0%
    sm_batch_few_col     nullcheck_udf     444ms    483ms    +8.8%
    sm_batch_many_col    identity_udf      306ms    308ms    +0.6%
    sm_batch_many_col    sort_udf          333ms    314ms    -5.5%
    sm_batch_many_col    nullcheck_udf     301ms    327ms    +8.7%
    lg_batch_few_col     identity_udf      1.13s    1.17s    +3.4%
    lg_batch_few_col     sort_udf          1.63s    1.32s   -19.1%
    lg_batch_few_col     nullcheck_udf     1.19s    1.19s    -0.2%
    lg_batch_many_col    identity_udf      1.59s    1.50s    -5.4%
    lg_batch_many_col    sort_udf          1.58s    1.52s    -4.1%
    lg_batch_many_col    nullcheck_udf     1.49s    1.51s    +1.2%
    pure_ints            identity_udf      196ms    192ms    -2.4%
    pure_ints            sort_udf          270ms    269ms    -0.4%
    pure_ints            nullcheck_udf     220ms    224ms    +1.9%
    pure_floats          identity_udf      209ms    202ms    -3.4%
    pure_floats          sort_udf          337ms    294ms   -12.9%
    pure_floats          nullcheck_udf     233ms    238ms    +1.9%
    pure_strings         identity_udf      1.26s    1.23s    -2.6%
    pure_strings         sort_udf          1.81s    1.65s    -8.4%
    pure_strings         nullcheck_udf     1.25s    1.15s    -7.8%
    pure_ts              identity_udf      464ms    414ms   -10.7%
    pure_ts              sort_udf          488ms    484ms    -0.8%
    pure_ts              nullcheck_udf     424ms    432ms    +2.0%
    mixed_types          identity_udf      719ms    669ms    -6.9%
    mixed_types          sort_udf          737ms    731ms    -0.9%
    mixed_types          nullcheck_udf     693ms    703ms    +1.5%
    ```
    
    The scattered positive cells (up to +8.8%) are benchmark ordering artifacts 
(matrix-mode run-to-run variance is ~+/-10% and these cells flip sign between 
repeated runs): re-running each such scenario in isolation (one scenario per 
fresh Python process, sides alternated, `repeat=5`) shows them flat. This does 
not apply in production where each Spark task runs in a fresh Python worker.
    
    ```text
    scenario (isolated)  udf              before    after     diff
    -------------------  --------------  -------  -------  -------
    sm_batch_few_col     nullcheck_udf     455ms    455ms    -0.1%
    sm_batch_many_col    nullcheck_udf     343ms    335ms    -2.4%
    lg_batch_few_col     identity_udf      1.12s    1.10s    -1.7%
    sm_batch_many_col    identity_udf      298ms    303ms    +1.7%
    lg_batch_many_col    identity_udf      1.19s    1.22s    +2.2%
    ```
    
    **ScalarPandasIterUDFPeakmemBench** (peak memory):
    
    ```text
    scenario             udf              before    after     diff
    -------------------  --------------  -------  -------  -------
    sm_batch_few_col     identity_udf       135M     134M    -0.3%
    sm_batch_few_col     sort_udf           135M     135M    -0.4%
    sm_batch_few_col     nullcheck_udf      133M     133M    +0.2%
    sm_batch_many_col    identity_udf       133M     134M    +0.5%
    sm_batch_many_col    sort_udf           134M     134M    +0.1%
    sm_batch_many_col    nullcheck_udf      134M     135M    +0.1%
    lg_batch_few_col     identity_udf       249M     248M    -0.2%
    lg_batch_few_col     sort_udf           249M     249M    -0.0%
    lg_batch_few_col     nullcheck_udf      249M     249M    -0.0%
    lg_batch_many_col    identity_udf       270M     270M    +0.0%
    lg_batch_many_col    sort_udf           271M     271M    -0.0%
    lg_batch_many_col    nullcheck_udf      271M     270M    -0.1%
    pure_ints            identity_udf       157M     157M    +0.1%
    pure_ints            sort_udf           157M     156M    -0.3%
    pure_ints            nullcheck_udf      157M     156M    -0.2%
    pure_floats          identity_udf       195M     195M    -0.0%
    pure_floats          sort_udf           195M     195M    +0.1%
    pure_floats          nullcheck_udf      195M     195M    +0.2%
    pure_strings         identity_udf       219M     219M    +0.4%
    pure_strings         sort_udf           220M     217M    -1.0%
    pure_strings         nullcheck_udf      213M     213M    -0.0%
    pure_ts              identity_udf       196M     196M    +0.1%
    pure_ts              sort_udf           196M     196M    -0.2%
    pure_ts              nullcheck_udf      196M     196M    -0.2%
    mixed_types          identity_udf       178M     178M    -0.1%
    mixed_types          sort_udf           178M     178M    +0.4%
    mixed_types          nullcheck_udf      178M     178M    -0.0%
    ```
    
    **Summary**: no regression; most scenarios flat to slightly better 
(sort/string-heavy up to -19.1%), the few positive cells are matrix-run 
ordering artifacts (flat when re-run in isolation); peak memory flat.
    
    ### Was this patch authored or co-authored using generative AI tooling?
    
    No
    
    Closes #55756 from Yicong-Huang/SPARK-56757.
    
    Authored-by: Yicong Huang <[email protected]>
    Signed-off-by: Yicong-Huang <[email protected]>
    (cherry picked from commit 0f62d3df0a73385ea7c43c581305480c58d67a88)
    Signed-off-by: Yicong-Huang <[email protected]>
---
 python/pyspark/worker.py | 178 +++++++++++++++++++----------------------------
 1 file changed, 70 insertions(+), 108 deletions(-)

diff --git a/python/pyspark/worker.py b/python/pyspark/worker.py
index e28bbecf5e02..2be915fa358c 100644
--- a/python/pyspark/worker.py
+++ b/python/pyspark/worker.py
@@ -76,7 +76,6 @@ from pyspark.sql.functions import 
SkipRestOfInputTableException
 from pyspark.sql.pandas.serializers import (
     ArrowStreamSerializer,
     ArrowStreamGroupSerializer,
-    ArrowStreamPandasUDFSerializer,
     ArrowStreamPandasUDTFSerializer,
     ArrowStreamCoGroupSerializer,
     ApplyInPandasWithStateSerializer,
@@ -400,43 +399,6 @@ def wrap_udf(f, args_offsets, kwargs_offsets, return_type):
         return args_kwargs_offsets, lambda *a: func(*a)
 
 
-def wrap_pandas_batch_iter_udf(f, return_type, runner_conf):
-    iter_type_label = "pandas.DataFrame" if isinstance(return_type, 
StructType) else "pandas.Series"
-
-    def verify_result(result):
-        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__,
-                },
-            )
-        return result
-
-    def verify_element(elem):
-        import pandas as pd
-
-        if not isinstance(elem, pd.DataFrame if isinstance(return_type, 
StructType) else pd.Series):
-            raise PySparkTypeError(
-                errorClass="UDF_RETURN_TYPE",
-                messageParameters={
-                    "expected": "iterator of {}".format(iter_type_label),
-                    "actual": "iterator of {}".format(type(elem).__name__),
-                },
-            )
-
-        verify_pandas_result(
-            elem, return_type, assign_cols_by_name=True, 
truncate_return_schema=True
-        )
-
-        return elem
-
-    return lambda *iterator: map(
-        lambda res: (res, return_type), map(verify_element, 
verify_result(f(*iterator)))
-    )
-
-
 def _verify_column_schema(
     actual_names: list, expected_names: list, *, assign_cols_by_name: bool
 ) -> None:
@@ -864,7 +826,7 @@ def read_single_udf(pickleSer, udf_info, eval_type, 
runner_conf, udf_index):
     elif eval_type == PythonEvalType.SQL_ARROW_BATCHED_UDF:
         return func, args_offsets, kwargs_offsets, return_type
     elif eval_type == PythonEvalType.SQL_SCALAR_PANDAS_ITER_UDF:
-        return args_offsets, wrap_pandas_batch_iter_udf(func, return_type, 
runner_conf)
+        return func, args_offsets, kwargs_offsets, return_type
     elif eval_type == PythonEvalType.SQL_SCALAR_ARROW_ITER_UDF:
         return func, args_offsets, kwargs_offsets, return_type
     elif eval_type == PythonEvalType.SQL_MAP_PANDAS_ITER_UDF:
@@ -2169,33 +2131,8 @@ def read_udfs(pickleSer, udf_info_list, eval_type, 
runner_conf, eval_conf):
             ser = TransformWithStateInPySparkRowInitStateSerializer(
                 
arrow_max_records_per_batch=runner_conf.arrow_max_records_per_batch
             )
-        elif eval_type in (
-            PythonEvalType.SQL_MAP_ARROW_ITER_UDF,
-            PythonEvalType.SQL_MAP_PANDAS_ITER_UDF,
-            PythonEvalType.SQL_SCALAR_ARROW_UDF,
-            PythonEvalType.SQL_SCALAR_ARROW_ITER_UDF,
-            PythonEvalType.SQL_ARROW_BATCHED_UDF,
-            PythonEvalType.SQL_SCALAR_PANDAS_UDF,
-        ):
-            ser = ArrowStreamSerializer(write_start_stream=True)
         else:
-            # Scalar Pandas UDF handles struct type arguments as pandas 
DataFrames instead of
-            # pandas Series. See SPARK-27240.
-            df_for_struct = eval_type == 
PythonEvalType.SQL_SCALAR_PANDAS_ITER_UDF
-
-            ser = ArrowStreamPandasUDFSerializer(
-                timezone=runner_conf.timezone,
-                safecheck=runner_conf.safecheck,
-                assign_cols_by_name=runner_conf.assign_cols_by_name,
-                df_for_struct=df_for_struct,
-                struct_in_pandas="dict",
-                ndarray_as_list=False,
-                prefer_int_ext_dtype=runner_conf.prefer_int_ext_dtype,
-                arrow_cast=True,
-                input_type=None,
-                
int_to_decimal_coercion_enabled=runner_conf.int_to_decimal_coercion_enabled,
-                prefers_large_types=runner_conf.use_large_var_types,
-            )
+            ser = ArrowStreamSerializer(write_start_stream=True)
     else:
         batch_size = int(os.environ.get("PYTHON_UDF_BATCH_SIZE", "100"))
         ser = BatchedSerializer(CPickleSerializer(), batch_size)
@@ -3355,58 +3292,83 @@ def read_udfs(pickleSer, udf_info_list, eval_type, 
runner_conf, eval_conf):
         return func, None, ser, ser
 
     if eval_type == PythonEvalType.SQL_SCALAR_PANDAS_ITER_UDF:
-        assert num_udfs == 1, "One SCALAR_ITER UDF expected here."
+        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]
 
-        arg_offsets, udf = 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"""
 
-        def func(_, iterator):  # type: ignore[misc]
             num_input_rows = 0
 
-            def map_batch(batch):
+            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
 
-                udf_args = [batch[offset] for offset in arg_offsets]
-                num_input_rows += len(udf_args[0])
-                if len(udf_args) == 1:
-                    return udf_args[0]
-                else:
-                    return tuple(udf_args)
-
-            iterator = map(map_batch, iterator)
-            result_iter = udf(iterator)
-
-            num_output_rows = 0
-            for result_batch, result_type in result_iter:
-                num_output_rows += len(result_batch)
-                # This check is for Scalar Iterator UDF to fail fast.
-                # The length of the entire input can only be explicitly known
-                # by consuming the input iterator in user side. Therefore,
-                # it's very unlikely the output length is higher than
-                # input length.
-                if num_output_rows > num_input_rows:
-                    raise PySparkRuntimeError(
-                        errorClass="OUTPUT_EXCEEDS_INPUT_ROWS", 
messageParameters={}
+            # 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,
                     )
-                yield (result_batch, result_type)
 
-            try:
-                next(iterator)
-            except StopIteration:
-                pass
-            else:
-                raise PySparkRuntimeError(
-                    errorClass="INPUT_NOT_FULLY_CONSUMED",
-                    messageParameters={},
-                )
+            # Apply row limit check (fail-fast)
+            limited = verify_output_row_limit(
+                process_results(),
+                lambda: num_input_rows,
+                error_class="OUTPUT_EXCEEDS_INPUT_ROWS",
+            )
 
-            if num_output_rows != num_input_rows:
-                raise PySparkRuntimeError(
-                    errorClass="RESULT_ROWS_MISMATCH",
-                    messageParameters={
-                        "output_length": str(num_output_rows),
-                        "input_length": str(num_input_rows),
-                    },
-                )
+            # Apply row count match check (final)
+            matched = verify_output_row_count(
+                limited,
+                lambda: num_input_rows,
+                error_class="RESULT_ROWS_MISMATCH",
+            )
+
+            # Yield batches
+            yield from matched
+
+            # Verify iterator consumed
+            verify_iterator_exhausted(
+                args_iter,
+                error_class="INPUT_NOT_FULLY_CONSUMED",
+            )
 
         # profiling is not supported for UDF
         return func, None, ser, ser


---------------------------------------------------------------------
To unsubscribe, e-mail: [email protected]
For additional commands, e-mail: [email protected]

Reply via email to