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]