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

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


The following commit(s) were added to refs/heads/master by this push:
     new 8e422cc75f0b [SPARK-58128][PYTHON] Refactor 
SQL_TRANSFORM_WITH_STATE_PANDAS_INIT_STATE_UDF
8e422cc75f0b is described below

commit 8e422cc75f0b89ed464024ceee3103d6e91f8946
Author: Yicong Huang <[email protected]>
AuthorDate: Fri Jul 17 06:49:22 2026 +0000

    [SPARK-58128][PYTHON] Refactor 
SQL_TRANSFORM_WITH_STATE_PANDAS_INIT_STATE_UDF
    
    ### What changes were proposed in this pull request?
    
    This PR refactors `SQL_TRANSFORM_WITH_STATE_PANDAS_INIT_STATE_UDF` so that 
the worker uses the plain `ArrowStreamSerializer` for pure Arrow stream I/O, 
moving the per-eval-type logic (flattening the `inputData`/`initState` struct 
columns, regrouping rows by grouping key, re-chunking into pandas DataFrames 
bounded by `arrow_max_records_per_batch`/`arrow_max_bytes_per_batch`, splitting 
each group into a data iterator and an init-state iterator, and converting 
result DataFrames back to A [...]
    
    ### Why are the changes needed?
    
    Part of [SPARK-55388](https://issues.apache.org/jira/browse/SPARK-55388). 
Keeping serializers as pure Arrow stream I/O and concentrating 
eval-type-specific logic in `worker.py` makes the per-eval-type data flow 
explicit.
    
    ### Does this PR introduce _any_ user-facing change?
    
    No.
    
    ### How was this patch tested?
    
    Existing tests. No behavior change.
    
    ASV comparison 
(`bench_eval_type.TransformWithStatePandasInitStateUDFTimeBench` / 
`TransformWithStatePandasInitStateUDFPeakmemBench`, `-a repeat=3`): before = 
`upstream/master`, after = this PR. Values are from one representative run per 
side; the conclusion is consistent across runs. All deltas are within 
run-to-run noise (the larger negatives track high before-side variance, e.g. 
`many_groups_lg/count` before `5.05+-0.9s`), showing no regression.
    
    ```text
    time_worker
    
    scenario         udf            before        after         diff
    ---------------  ------------  ------------  ------------  ------
    few_groups_sm    identity_udf   854+-7ms      850+-8ms      -0.5%
    few_groups_sm    sort_udf       875+-10ms     848+-20ms     -3.1%
    few_groups_sm    count_udf      908+-30ms     881+-10ms     -3.0%
    few_groups_lg    identity_udf   7.70+-0.05s   7.78+-0.1s    +1.0%
    few_groups_lg    sort_udf       7.81+-0.04s   7.73+-0.01s   -1.0%
    few_groups_lg    count_udf      7.21+-0.1s    7.26+-0.09s   +0.7%
    many_groups_sm   identity_udf   8.07+-0.1s    8.13+-0.01s   +0.7%
    many_groups_sm   sort_udf       8.36+-0.03s   8.45+-0.02s   +1.1%
    many_groups_sm   count_udf      9.39+-0.09s   9.44+-0.05s   +0.5%
    many_groups_lg   identity_udf   4.21+-0.06s   4.24+-0.06s   +0.7%
    many_groups_lg   sort_udf       4.45+-0.2s    4.31+-0.05s   -3.1%
    many_groups_lg   count_udf      5.05+-0.9s    4.46+-0.1s    -11.7%
    wide_cols        identity_udf   9.14+-0.6s    8.33+-0.07s   -8.9%
    wide_cols        sort_udf       8.70+-0.1s    8.79+-0.3s    +1.0%
    wide_cols        count_udf      8.07+-0.1s    7.92+-0.2s    -1.9%
    mixed_cols       identity_udf   3.76+-0.1s    3.67+-0.08s   -2.4%
    mixed_cols       sort_udf       3.81+-0.2s    3.84+-0.09s   +0.8%
    mixed_cols       count_udf      3.51+-0.2s    3.57+-0.07s   +1.7%
    nested_struct    identity_udf   8.74+-0.2s    8.76+-0.1s    +0.2%
    nested_struct    sort_udf       9.55+-0.8s    8.79+-0.2s    -8.0%
    nested_struct    count_udf      6.33+-0.2s    6.23+-0.1s    -1.6%
    ```
    
    ```text
    peakmem_worker
    
    scenario         udf            before   after
    ---------------  ------------  -------  -------
    few_groups_sm    identity_udf   115M     118M
    few_groups_sm    sort_udf       118M     116M
    few_groups_sm    count_udf      107M     107M
    few_groups_lg    identity_udf   249M     249M
    few_groups_lg    sort_udf       249M     249M
    few_groups_lg    count_udf      249M     249M
    many_groups_sm   identity_udf   176M     176M
    many_groups_sm   sort_udf       178M     180M
    many_groups_sm   count_udf      162M     162M
    many_groups_lg   identity_udf   152M     152M
    many_groups_lg   sort_udf       152M     152M
    many_groups_lg   count_udf      152M     152M
    wide_cols        identity_udf   365M     368M
    wide_cols        sort_udf       369M     377M
    wide_cols        count_udf      343M     343M
    mixed_cols       identity_udf   182M     183M
    mixed_cols       sort_udf       182M     183M
    mixed_cols       count_udf      182M     183M
    nested_struct    identity_udf   211M     211M
    nested_struct    sort_udf       211M     211M
    nested_struct    count_udf      211M     211M
    ```
    
    ### Was this patch authored or co-authored using generative AI tooling?
    
    No.
    
    Closes #57260 from Yicong-Huang/refactor-tws-pandas-init.
    
    Authored-by: Yicong Huang <[email protected]>
    Signed-off-by: Yicong-Huang <[email protected]>
---
 python/pyspark/worker.py | 252 ++++++++++++++++++++++++++++++++++++-----------
 1 file changed, 196 insertions(+), 56 deletions(-)

diff --git a/python/pyspark/worker.py b/python/pyspark/worker.py
index 59c389d8f737..d68a4e7ea8eb 100644
--- a/python/pyspark/worker.py
+++ b/python/pyspark/worker.py
@@ -79,7 +79,6 @@ from pyspark.sql.pandas.serializers import (
     ArrowStreamGroupSerializer,
     ArrowStreamCoGroupSerializer,
     ApplyInPandasWithStateSerializer,
-    TransformWithStateInPandasInitStateSerializer,
     TransformWithStateInPySparkRowSerializer,
     TransformWithStateInPySparkRowInitStateSerializer,
 )
@@ -500,26 +499,6 @@ def verify_arrow_result(
         )
 
 
-def wrap_grouped_transform_with_state_pandas_init_state_udf(f, return_type, 
runner_conf):
-    def wrapped(stateful_processor_api_client, mode, key, value_series_gen):
-        # Split the generator into two using itertools.tee
-        state_values_gen, init_states_gen = itertools.tee(value_series_gen, 2)
-
-        # Extract just the data DataFrames (first element of each tuple)
-        state_values = (data_df for data_df, _ in state_values_gen if not 
data_df.empty)
-
-        # Extract just the init DataFrames (second element of each tuple)
-        init_states = (init_df for _, init_df in init_states_gen if not 
init_df.empty)
-        result_iter = f(stateful_processor_api_client, mode, key, 
state_values, init_states)
-
-        # TODO(SPARK-49100): add verification that elements in result_iter are
-        # indeed of type pd.DataFrame and confirm to assigned cols
-
-        return result_iter
-
-    return lambda p, m, k, v: [(wrapped(p, m, k, v), return_type)]
-
-
 def wrap_grouped_transform_with_state_udf(f, return_type, runner_conf):
     def wrapped(stateful_processor_api_client, mode, key, values):
         result_iter = f(stateful_processor_api_client, mode, key, values)
@@ -817,9 +796,7 @@ def read_single_udf(pickleSer, udf_info, eval_type, 
runner_conf, udf_index):
     elif eval_type == PythonEvalType.SQL_TRANSFORM_WITH_STATE_PANDAS_UDF:
         return func, args_offsets, return_type
     elif eval_type == 
PythonEvalType.SQL_TRANSFORM_WITH_STATE_PANDAS_INIT_STATE_UDF:
-        return args_offsets, 
wrap_grouped_transform_with_state_pandas_init_state_udf(
-            func, return_type, runner_conf
-        )
+        return func, args_offsets, return_type
     elif eval_type == PythonEvalType.SQL_TRANSFORM_WITH_STATE_PYTHON_ROW_UDF:
         return args_offsets, wrap_grouped_transform_with_state_udf(func, 
return_type, runner_conf)
     elif eval_type == 
PythonEvalType.SQL_TRANSFORM_WITH_STATE_PYTHON_ROW_INIT_STATE_UDF:
@@ -2052,16 +2029,6 @@ def read_udfs(pickleSer, udf_info_list, eval_type, 
runner_conf, eval_conf):
                 prefers_large_var_types=runner_conf.use_large_var_types,
                 
int_to_decimal_coercion_enabled=runner_conf.int_to_decimal_coercion_enabled,
             )
-        elif eval_type == 
PythonEvalType.SQL_TRANSFORM_WITH_STATE_PANDAS_INIT_STATE_UDF:
-            ser = TransformWithStateInPandasInitStateSerializer(
-                timezone=runner_conf.timezone,
-                safecheck=runner_conf.safecheck,
-                assign_cols_by_name=runner_conf.assign_cols_by_name,
-                prefer_int_ext_dtype=runner_conf.prefer_int_ext_dtype,
-                
arrow_max_records_per_batch=runner_conf.arrow_max_records_per_batch,
-                
arrow_max_bytes_per_batch=runner_conf.arrow_max_bytes_per_batch,
-                
int_to_decimal_coercion_enabled=runner_conf.int_to_decimal_coercion_enabled,
-            )
         elif eval_type == 
PythonEvalType.SQL_TRANSFORM_WITH_STATE_PYTHON_ROW_UDF:
             ser = TransformWithStateInPySparkRowSerializer(
                 
arrow_max_records_per_batch=runner_conf.arrow_max_records_per_batch
@@ -3504,44 +3471,217 @@ def read_udfs(pickleSer, udf_info_list, eval_type, 
runner_conf, eval_conf):
         return transform_with_state_func, None, ser, ser
 
     if eval_type == 
PythonEvalType.SQL_TRANSFORM_WITH_STATE_PANDAS_INIT_STATE_UDF:
-        # We assume there is only one UDF here because grouped map doesn't
-        # support combining multiple UDFs.
-        assert num_udfs == 1
+        import pyarrow as pa
+        import pandas as pd
+
+        assert num_udfs == 1, "One TRANSFORM_WITH_STATE_PANDAS_INIT_STATE UDF 
expected here."
+        udf, arg_offsets, return_type = udfs[0]
 
         # See TransformWithStateInPandasExec for how arg_offsets are used to
-        # distinguish between grouping attributes and data attributes
-        arg_offsets, f = udfs[0]
+        # distinguish between grouping attributes and data attributes.
         # parsed offsets:
         # [
         #     [groupingKeyOffsets, dedupDataOffsets],
         #     [initStateGroupingOffsets, dedupInitDataOffsets]
         # ]
         parsed_offsets = extract_key_value_indexes(arg_offsets)
-        ser.key_offsets = parsed_offsets[0][0]
-        ser.init_key_offsets = parsed_offsets[1][0]
+        key_offsets = parsed_offsets[0][0]
+        init_key_offsets = parsed_offsets[1][0]
+        output_schema = StructType([StructField("_0", return_type)])
+
         stateful_processor_api_client = StatefulProcessorApiClient(
             eval_conf.state_server_socket_port, eval_conf.grouping_key_schema
         )
 
-        def mapper(a):
-            mode = a[0]
+        arrow_max_records_per_batch = runner_conf.arrow_max_records_per_batch
+        arrow_max_records_per_batch = (
+            arrow_max_records_per_batch if arrow_max_records_per_batch > 0 
else 2**31 - 1
+        )
+        arrow_max_bytes_per_batch = runner_conf.arrow_max_bytes_per_batch
 
-            if mode == TransformWithStateInPandasFuncMode.PROCESS_DATA:
-                key = a[1]
+        def func(
+            split_index: int,
+            data: Iterator[pa.RecordBatch],
+        ) -> Iterator[pa.RecordBatch]:
+            """Apply transformWithStateInPandas UDF with initial state.
 
-                def values_gen():
-                    for x in a[2]:
-                        retVal = x[1]
-                        initVal = x[2]
-                        yield retVal, initVal
+            The input batches carry two struct columns, ``inputData`` and
+            ``initState``; each batch holds one or the other but never both.
+            Rows are flattened out of whichever struct is present, regrouped by
+            grouping key, and re-chunked into pandas DataFrames bounded by
+            arrow_max_records_per_batch and arrow_max_bytes_per_batch. The UDF
+            is invoked once per grouping key with two separate lazy iterators
+            (data DataFrames and init-state DataFrames), then once for
+            PROCESS_TIMER and once for COMPLETE.
+            """
+            total_bytes = 0
+            total_rows = 0
+            average_arrow_row_size = 0.0
 
-                # This must be generator comprehension - do not materialize.
-                return f(stateful_processor_api_client, mode, key, 
values_gen())
-            else:
-                # mode == PROCESS_TIMER or mode == COMPLETE
-                return f(stateful_processor_api_client, mode, None, iter([]))
+            def flatten_columns(cur_batch: "pa.RecordBatch", col_name: str) -> 
"pa.Table":
+                struct_column = 
cur_batch.column(cur_batch.schema.get_field_index(col_name))
+                # Check if the entire column is null: an empty table (no 
columns)
+                # signals the struct is absent from this batch.
+                if struct_column.null_count == len(struct_column):
+                    return pa.Table.from_arrays([], names=[])
+                field_names = [
+                    struct_column.type[i].name for i in 
range(struct_column.type.num_fields)
+                ]
+                field_arrays = [
+                    struct_column.field(i) for i in 
range(struct_column.type.num_fields)
+                ]
+                return pa.Table.from_arrays(field_arrays, names=field_names)
 
-    elif eval_type == PythonEvalType.SQL_TRANSFORM_WITH_STATE_PYTHON_ROW_UDF:
+            def to_pandas(table: "pa.Table") -> list:
+                return ArrowBatchTransformer.to_pandas(
+                    table,
+                    timezone=runner_conf.timezone,
+                    prefer_int_ext_dtype=runner_conf.prefer_int_ext_dtype,
+                )
+
+            def row_stream() -> Iterator[tuple]:
+                nonlocal total_bytes, total_rows, average_arrow_row_size
+                for batch in data:
+                    # Short circuit batch size stats if the batch size is
+                    # unlimited as computing batch size is computationally
+                    # expensive.
+                    if arrow_max_bytes_per_batch != 2**31 - 1 and 
batch.num_rows > 0:
+                        total_bytes += sum(
+                            buf.size
+                            for col in batch.columns
+                            for buf in col.buffers()
+                            if buf is not None
+                        )
+                        total_rows += batch.num_rows
+                        average_arrow_row_size = total_bytes / total_rows
+
+                    data_table = flatten_columns(batch, "inputData")
+                    init_table = flatten_columns(batch, "initState")
+
+                    # Empty table has no columns. Each batch carries either
+                    # input data or init state, never both.
+                    has_data = data_table.num_columns > 0
+                    has_init = init_table.num_columns > 0
+                    assert not (has_data and has_init)
+
+                    if has_data:
+                        for row in pd.concat(to_pandas(data_table), 
axis=1).itertuples(index=False):
+                            batch_key = tuple(row[o] for o in key_offsets)
+                            yield (batch_key, row, None)
+                    elif has_init:
+                        for row in pd.concat(to_pandas(init_table), 
axis=1).itertuples(index=False):
+                            batch_key = tuple(row[o] for o in init_key_offsets)
+                            yield (batch_key, None, row)
+
+            empty_dataframe = pd.DataFrame()
+
+            def generate_data_batches() -> Iterator[tuple]:
+                """
+                Deserialize ArrowRecordBatches and return a generator of
+                (grouping key, data DataFrame, init-state DataFrame) chunks.
+
+                This function must avoid materializing multiple Arrow
+                RecordBatches into memory at the same time. And data chunks
+                from the same grouping key should appear sequentially.
+                """
+                for batch_key, group_rows in itertools.groupby(row_stream(), 
key=lambda x: x[0]):
+                    rows = []
+                    init_state_rows = []
+                    for _, row, init_state_row in group_rows:
+                        if row is not None:
+                            rows.append(row)
+                        if init_state_row is not None:
+                            init_state_rows.append(init_state_row)
+
+                        total_len = len(rows) + len(init_state_rows)
+                        if (
+                            total_len >= arrow_max_records_per_batch
+                            or total_len * average_arrow_row_size >= 
arrow_max_bytes_per_batch
+                        ):
+                            yield (
+                                batch_key,
+                                pd.DataFrame(rows) if rows else 
empty_dataframe.copy(),
+                                (
+                                    pd.DataFrame(init_state_rows)
+                                    if init_state_rows
+                                    else empty_dataframe.copy()
+                                ),
+                            )
+                            rows = []
+                            init_state_rows = []
+                    if rows or init_state_rows:
+                        yield (
+                            batch_key,
+                            pd.DataFrame(rows) if rows else 
empty_dataframe.copy(),
+                            (
+                                pd.DataFrame(init_state_rows)
+                                if init_state_rows
+                                else empty_dataframe.copy()
+                            ),
+                        )
+
+            def convert_results(
+                result_iter: Iterable["pd.DataFrame"],
+            ) -> Iterator["pa.RecordBatch"]:
+                # TODO(SPARK-49100): add verification that elements in 
result_iter are
+                # indeed of type pd.DataFrame and conform to assigned cols
+                for result in result_iter:
+                    if isinstance(return_type, StructType) and not 
isinstance(result, pd.DataFrame):
+                        raise PySparkValueError(
+                            "Invalid return type. Please make sure that the 
UDF returns a "
+                            "pandas.DataFrame when the specified return type 
is StructType."
+                        )
+                    yield PandasToArrowConversion.convert(
+                        [result],
+                        output_schema,
+                        timezone=runner_conf.timezone,
+                        safecheck=runner_conf.safecheck,
+                        arrow_cast=True,
+                        assign_cols_by_name=runner_conf.assign_cols_by_name,
+                        
int_to_decimal_coercion_enabled=runner_conf.int_to_decimal_coercion_enabled,
+                    )
+
+            for key, group in itertools.groupby(generate_data_batches(), 
key=lambda x: x[0]):
+                # These must be generator expressions - do not materialize. The
+                # UDF receives the data and init-state DataFrames as two
+                # separate iterators, with empty chunks filtered out.
+                group_data, group_init = itertools.tee(group, 2)
+                state_values = (data_df for _, data_df, _ in group_data if not 
data_df.empty)
+                init_states = (init_df for _, _, init_df in group_init if not 
init_df.empty)
+                yield from convert_results(
+                    udf(
+                        stateful_processor_api_client,
+                        TransformWithStateInPandasFuncMode.PROCESS_DATA,
+                        key,
+                        state_values,
+                        init_states,
+                    )
+                )
+
+            yield from convert_results(
+                udf(
+                    stateful_processor_api_client,
+                    TransformWithStateInPandasFuncMode.PROCESS_TIMER,
+                    None,
+                    iter([]),
+                    iter([]),
+                )
+            )
+
+            yield from convert_results(
+                udf(
+                    stateful_processor_api_client,
+                    TransformWithStateInPandasFuncMode.COMPLETE,
+                    None,
+                    iter([]),
+                    iter([]),
+                )
+            )
+
+        # profiling is not supported for UDF
+        return func, None, ser, ser
+
+    if eval_type == PythonEvalType.SQL_TRANSFORM_WITH_STATE_PYTHON_ROW_UDF:
         # We assume there is only one UDF here because grouped map doesn't
         # support combining multiple UDFs.
         assert num_udfs == 1


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

Reply via email to