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]