diff --git a/python/pyspark/worker.py b/python/pyspark/worker.py index 558a3d53cd89..cd5b0fb73c32 100644 --- a/python/pyspark/worker.py +++ b/python/pyspark/worker.py @@ -78,8 +78,6 @@ ArrowStreamSerializer, ArrowStreamGroupSerializer, ArrowStreamCoGroupSerializer, - TransformWithStateInPySparkRowSerializer, - TransformWithStateInPySparkRowInitStateSerializer, ) from pyspark.sql.pandas.types import to_arrow_schema, to_arrow_type from pyspark.sql.types import ( @@ -500,37 +498,6 @@ def verify_arrow_result( ) -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) - - # TODO(SPARK-XXXXX): add verification that elements in result_iter are - # indeed of type Row 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_init_state_udf(f, return_type, runner_conf): - def wrapped(stateful_processor_api_client, mode, key, values): - if mode == TransformWithStateInPandasFuncMode.PROCESS_DATA: - values_gen = values[0] - init_states_gen = values[1] - else: - values_gen = iter([]) - init_states_gen = iter([]) - - result_iter = f(stateful_processor_api_client, mode, key, values_gen, init_states_gen) - - # TODO(SPARK-XXXXX): 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_kwargs_support(f, args_offsets, kwargs_offsets): if len(kwargs_offsets): keys = list(kwargs_offsets.keys()) @@ -708,11 +675,9 @@ def read_single_udf(pickleSer, udf_info, eval_type, runner_conf, udf_index): elif eval_type == PythonEvalType.SQL_TRANSFORM_WITH_STATE_PANDAS_INIT_STATE_UDF: 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) + return func, args_offsets, return_type elif eval_type == PythonEvalType.SQL_TRANSFORM_WITH_STATE_PYTHON_ROW_INIT_STATE_UDF: - return args_offsets, wrap_grouped_transform_with_state_init_state_udf( - func, return_type, runner_conf - ) + return func, args_offsets, return_type elif eval_type == PythonEvalType.SQL_COGROUPED_MAP_PANDAS_UDF: argspec = inspect.getfullargspec(chained_func) # signature was lost when wrapping it return func, args_offsets, return_type, len(argspec.args) @@ -1928,14 +1893,6 @@ def read_udfs(pickleSer, udf_info_list, eval_type, runner_conf, eval_conf): ser = ArrowStreamCoGroupSerializer(write_start_stream=True) elif eval_type == PythonEvalType.SQL_COGROUPED_MAP_PANDAS_UDF: ser = ArrowStreamCoGroupSerializer(write_start_stream=True) - 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 - ) - elif eval_type == PythonEvalType.SQL_TRANSFORM_WITH_STATE_PYTHON_ROW_INIT_STATE_UDF: - ser = TransformWithStateInPySparkRowInitStateSerializer( - arrow_max_records_per_batch=runner_conf.arrow_max_records_per_batch - ) else: ser = ArrowStreamSerializer(write_start_stream=True) else: @@ -3303,7 +3260,7 @@ def generate_data_batches(): (grouping key, pandas.DataFrame) chunks. This function must avoid materializing multiple Arrow - RecordBatches into memory at the same time. And data chunks + 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]): @@ -3480,7 +3437,7 @@ def generate_data_batches() -> Iterator[tuple]: (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 + 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]): @@ -3914,64 +3871,283 @@ def result_and_state_stream() -> Iterator[tuple]: 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 + import pyarrow as pa + + assert num_udfs == 1, "One TRANSFORM_WITH_STATE_PYTHON_ROW UDF expected here." + udf, arg_offsets, return_type = udfs[0] # See TransformWithStateInPySparkExec 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 = extract_key_value_indexes(arg_offsets) - ser.key_offsets = parsed_offsets[0][0] + key_offsets = parsed_offsets[0][0] + stateful_processor_api_client = StatefulProcessorApiClient( eval_conf.state_server_socket_port, eval_conf.grouping_key_schema ) - def mapper(a): - mode = a[0] + def func( + split_index: int, + data: Iterator[pa.RecordBatch], + ) -> Iterator[pa.RecordBatch]: + """Apply transformWithStateInPySpark UDF over Row objects. + + Input batches are read row by row without materializing whole + batches. Rows carrying the same grouping key appear sequentially, so + they are regrouped by key and the UDF is invoked once per key with a + lazy iterator of Row objects, then once for PROCESS_TIMER and once + for COMPLETE. The UDF yields (iterator of Row, Spark type) pairs that + are converted back into Arrow RecordBatches wrapped in a single + struct column for the output stream. + """ - if mode == TransformWithStateInPandasFuncMode.PROCESS_DATA: - key = a[1] - values = a[2] + def generate_data_batches() -> Iterator[Tuple[Any, Any]]: + """ + Deserialize ArrowRecordBatches and return a generator of + (grouping key, Row) tuples. - # This must be generator comprehension - do not materialize. - return f(stateful_processor_api_client, mode, key, values) - else: - # mode == PROCESS_TIMER or mode == COMPLETE - return f(stateful_processor_api_client, mode, None, iter([])) + 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 in data: + DataRow = Row(*batch.schema.names) + + # Iterate row by row without converting the whole batch. + num_cols = batch.num_columns + for row_idx in range(batch.num_rows): + row_key = tuple(batch[o][row_idx].as_py() for o in key_offsets) + row = DataRow(*(batch.column(i)[row_idx].as_py() for i in range(num_cols))) + yield row_key, row + + def convert_results(result_rows: Iterable[Any]) -> Iterator["pa.RecordBatch"]: + # TODO(SPARK-XXXXX): add verification that elements in result_rows + # are indeed of type Row and conform to assigned cols + + # Convert spark type to arrow type + # TODO: we need to make this configurable, currently using default values. + arrow_type = to_arrow_type( + return_type, + timezone="UTC", + prefers_large_types=False, + ) - elif eval_type == PythonEvalType.SQL_TRANSFORM_WITH_STATE_PYTHON_ROW_INIT_STATE_UDF: - # We assume there is only one UDF here because grouped map doesn't - # support combining multiple UDFs. - assert num_udfs == 1 + rows_as_dict = [row.asDict(True) for row in result_rows] + + pdf_schema = pa.schema(list(arrow_type)) + record_batch = pa.RecordBatch.from_pylist(rows_as_dict, schema=pdf_schema) + yield ArrowBatchTransformer.wrap_struct(record_batch) + + for key, group in itertools.groupby(generate_data_batches(), key=lambda x: x[0]): + # This must be a generator expression - do not materialize. + values_gen = map(lambda x: x[1], group) + yield from convert_results( + udf( + stateful_processor_api_client, + TransformWithStateInPandasFuncMode.PROCESS_DATA, + key, + values_gen, + ) + ) + + yield from convert_results( + udf( + stateful_processor_api_client, + TransformWithStateInPandasFuncMode.PROCESS_TIMER, + None, + iter([]), + ) + ) + + yield from convert_results( + udf( + stateful_processor_api_client, + TransformWithStateInPandasFuncMode.COMPLETE, + None, + iter([]), + ) + ) + + # profiling is not supported for UDF + return func, None, ser, ser + + if eval_type == PythonEvalType.SQL_TRANSFORM_WITH_STATE_PYTHON_ROW_INIT_STATE_UDF: + import pyarrow as pa + + assert num_udfs == 1, "One TRANSFORM_WITH_STATE_PYTHON_ROW_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] + 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 + ) - if mode == TransformWithStateInPandasFuncMode.PROCESS_DATA: - key = a[1] - values = a[2] + def func( + split_index: int, + data: Iterator[pa.RecordBatch], + ) -> Iterator[pa.RecordBatch]: + """Apply transformWithStateInPySpark UDF with initial state over Rows. - # This must be generator comprehension - do not materialize. - return f(stateful_processor_api_client, mode, key, values) - else: - # mode == PROCESS_TIMER or mode == COMPLETE - return f(stateful_processor_api_client, mode, None, iter([])) + 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 Row lists bounded by + arrow_max_records_per_batch. The UDF is invoked once per chunk with + two separate iterators (data Rows and init-state Rows), then once for + PROCESS_TIMER and once for COMPLETE. The UDF yields (iterator of Row, + Spark type) pairs that are converted back into Arrow RecordBatches + wrapped in a single struct column for the output stream. + """ + + def extract_rows( + cur_batch: "pa.RecordBatch", col_name: str, offsets: list + ) -> Optional[Iterator[Tuple[Any, Any]]]: + data_column = cur_batch.column(cur_batch.schema.get_field_index(col_name)) + + # Check if the entire column is null. + if data_column.null_count == len(data_column): + return None + + data_field_names = [ + data_column.type[i].name for i in range(data_column.type.num_fields) + ] + data_field_arrays = [ + data_column.field(i) for i in range(data_column.type.num_fields) + ] + + DataRow = Row(*data_field_names) + + table = pa.Table.from_arrays(data_field_arrays, names=data_field_names) + + if table.num_rows == 0: + return None + + def row_iterator() -> Iterator[Tuple[Any, Any]]: + for row_idx in range(table.num_rows): + key = tuple(table.column(o)[row_idx].as_py() for o in offsets) + row = DataRow( + *(table.column(i)[row_idx].as_py() for i in range(table.num_columns)) + ) + yield (key, row) + + return row_iterator() + + def row_stream() -> Iterator[Tuple[Any, Optional[Any], Optional[Any]]]: + # The arrow batch is written in the schema: + # schema: StructType = new StructType() + # .add("inputData", dataSchema) + # .add("initState", initStateSchema) + # We parse each batch into tuples of (key, inputData, initState). + # Each batch will have either init_data or input_data, not both. + for batch in data: + input_result = extract_rows(batch, "inputData", key_offsets) + init_result = extract_rows(batch, "initState", init_key_offsets) + + assert not (input_result is not None and init_result is not None) + + if input_result is not None: + for key, input_data_row in input_result: + yield (key, input_data_row, None) + elif init_result is not None: + for key, init_state_row in init_result: + yield (key, None, init_state_row) + + def generate_data_batches() -> Iterator[Tuple[Any, Tuple[Any, Any]]]: + """ + Deserialize ArrowRecordBatches and return a generator of + (grouping key, (data Rows iterator, init-state Rows iterator)) + chunks bounded by arrow_max_records_per_batch. + + 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 k, group_rows in itertools.groupby(row_stream(), key=lambda x: x[0]): + input_rows: list = [] + init_rows: list = [] + + for _, input_row, init_row in group_rows: + if input_row is not None: + input_rows.append(input_row) + if init_row is not None: + init_rows.append(init_row) + + total_len = len(input_rows) + len(init_rows) + if total_len >= arrow_max_records_per_batch: + yield (k, (iter(input_rows), iter(init_rows))) + input_rows = [] + init_rows = [] + + if input_rows or init_rows: + yield (k, (iter(input_rows), iter(init_rows))) + + def convert_results(result_rows: Iterable[Any]) -> Iterator["pa.RecordBatch"]: + # TODO(SPARK-XXXXX): add verification that elements in result_rows + # are indeed of type Row and conform to assigned cols + + # Convert spark type to arrow type + # TODO: we need to make this configurable, currently using default values. + arrow_type = to_arrow_type( + return_type, + timezone="UTC", + prefers_large_types=False, + ) + + rows_as_dict = [row.asDict(True) for row in result_rows] + + pdf_schema = pa.schema(list(arrow_type)) + record_batch = pa.RecordBatch.from_pylist(rows_as_dict, schema=pdf_schema) + yield ArrowBatchTransformer.wrap_struct(record_batch) + + for key, group in itertools.groupby(generate_data_batches(), key=lambda x: x[0]): + # These must be generator expressions - do not materialize. + for _, (values_gen, init_states_gen) in group: + yield from convert_results( + udf( + stateful_processor_api_client, + TransformWithStateInPandasFuncMode.PROCESS_DATA, + key, + values_gen, + init_states_gen, + ) + ) + + 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 else: