viirya commented on code in PR #58978: URL: https://github.com/apache/spark/pull/58978#discussion_r4109458114
########## python/pyspark/inprocess/runtime.py: ########## @@ -0,0 +1,261 @@ +# +# Licensed to the Apache Software Foundation (ASF) under one or more +# contributor license agreements. See the NOTICE file distributed with +# this work for additional information regarding copyright ownership. +# The ASF licenses this file to You under the Apache License, Version 2.0 +# (the "License"); you may not use this file except in compliance with +# the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# + + +"""Arrow CDI entry points called on the executor's dedicated JEP interpreter thread. + +Functions are registered once per task and released when that task finishes. Calls +pass only a handle and CDI addresses, so large closures are not copied per batch. +""" + +import sys +import traceback as _traceback +from typing import Any, Callable, Iterable, Optional, Sequence + +import pyarrow as pa +import pyarrow.compute as pc + +from pyspark import cloudpickle +from pyspark.errors import PySparkRuntimeError +from pyspark.sql.pandas.types import to_arrow_type +from pyspark.sql.types import _parse_datatype_json_string +from pyspark.util import try_simplify_traceback + +_UDF_TRACEBACK_SENTINEL = "__INPROCESS_UDF_TRACEBACK__:" +NullChecker = Callable[[pa.Array], None] +_udfs: dict[str, tuple[Callable[..., pa.Array], pa.DataType, NullChecker, bool, bool]] = {} + + +def _format_exception(hide: bool, simplified: bool) -> str: + kind, error, tb = sys.exc_info() + if hide: + return "".join(_traceback.format_exception_only(kind, error)) + if simplified and tb is not None: + simple_tb = try_simplify_traceback(tb) + if simple_tb is not None: + tb = simple_tb + if error is not None: + error.__cause__ = None + return "".join(_traceback.format_exception(kind, error, tb)) + + +def _inprocess_register( + handle: str, + serialized_udf: Any, + return_type_json: str, + timezone: str, + python_version: str, + large_var_types: bool = False, + hide_traceback: bool = False, + simplified_traceback: bool = False, +) -> None: + try: + embedded_version = "%d.%d" % sys.version_info[:2] + if python_version != embedded_version: + raise PySparkRuntimeError( + errorClass="PYTHON_VERSION_MISMATCH", + messageParameters={ + "worker_version": embedded_version, + "driver_version": python_version, + }, + ) + # JEP exposes direct ByteBuffers through the buffer protocol. Unpickle a separate + # function per task without iterating over a PyJArray one JNI call per byte. + func = cloudpickle.loads(memoryview(serialized_udf)) + expected_type = to_arrow_type( + _parse_datatype_json_string(return_type_json), + timezone=timezone, + prefers_large_types=large_var_types, + error_on_duplicated_field_names_in_struct=True, + ) + checker = _null_checker(expected_type) or (lambda array: None) + _udfs[handle] = (func, expected_type, checker, hide_traceback, simplified_traceback) + except BaseException: + # In JEP, an uncaught SystemExit can terminate the entire executor JVM. + raise RuntimeError( + _UDF_TRACEBACK_SENTINEL + _format_exception(hide_traceback, simplified_traceback) + ) from None + + +def _inprocess_release(handles: Iterable[str]) -> None: + for handle in handles: + _udfs.pop(handle, None) + + +def _nullable_type(data_type: pa.DataType) -> pa.DataType: + def nullable_field(field: pa.Field) -> pa.Field: + return pa.field(field.name, _nullable_type(field.type), nullable=True) + + if pa.types.is_struct(data_type): + return pa.struct([nullable_field(field) for field in data_type]) + if pa.types.is_list(data_type): + return pa.list_(nullable_field(data_type.value_field)) + if pa.types.is_large_list(data_type): + return pa.large_list(nullable_field(data_type.value_field)) + if pa.types.is_map(data_type): + return pa.map_( + _nullable_type(data_type.key_type), + nullable_field(data_type.item_field), + keys_sorted=data_type.keys_sorted, + ) + return data_type + + +def _null_checker(expected_type: pa.DataType) -> Optional[NullChecker]: + """Compile checks only for required fields and their ancestors, once per registration.""" + + def field_checker(field: pa.Field) -> Optional[NullChecker]: + nested = _null_checker(field.type) + if field.nullable: + return nested + + def check(values: pa.Array) -> None: + if values.null_count: + raise ValueError( + f"In-process UDF returned nulls in non-nullable field {field.name}" + ) + if nested is not None: + nested(values) + + return check + + if pa.types.is_struct(expected_type): + fields = [(i, field_checker(f)) for i, f in enumerate(expected_type)] + checks = [(i, check) for i, check in fields if check is not None] + if not checks: + return None + + def check_struct(array: pa.Array) -> None: + # Only children of valid parents are logically visible. + visible = pc.filter(array, pc.is_valid(array)) if array.null_count else array Review Comment: Added a fast path when the raw children prove there can be no required-field violation. Struct checks filter only the child that needs checking, preserving sibling payload buffers. Map checks rely on Arrow validation for top-level key nullability while retaining checks on required descendants. Regression tests verify that null-parent struct/map results with null-free children invoke neither filter nor concat, and that checking another child does not copy a sibling payload. ########## sql/core/src/main/scala/org/apache/spark/sql/execution/python/ArrowEvalPythonExec.scala: ########## @@ -102,11 +103,12 @@ case class ArrowEvalPythonExec( // The Arrow FieldVectors are extracted directly from ArrowColumnVector and // serialized to IPC, bypassing the row-based ArrowWriter conversion. override def supportsColumnar: Boolean = - child.supportsColumnar && conf.arrowPySparkUDFColumnarInputEnabled + evalType != PythonEvalType.SQL_SCALAR_ARROW_INPROCESS_UDF && Review Comment: Implemented `InProcessArrowEvalPythonExec extends EvalPythonExec` and the specialized strategy case. It retains the existing evaluator factory and logical planning, and removes all four in-process branches from `ArrowEvalPythonExec`. Tests pass for cached columnar children, both partition-evaluator settings, mixed worker/in-process UDFs, and non-root limit/offset ordering. ########## python/pyspark/sql/udf.py: ########## @@ -856,6 +856,14 @@ def register( [Row(sum_udf(v1)=1), Row(sum_udf(v1)=5)] """ + # Avoid importing the optional PyArrow-backed module for ordinary UDF registration. Review Comment: Removed both explicit guards and their `sys.modules` lookups. The existing eval-type allow-lists still reject the wrapper with `INVALID_UDF_EVAL_TYPE`; the classic and Connect registration regressions pass. ########## docs/sql-pyspark-inprocess-udf.md: ########## @@ -0,0 +1,621 @@ +--- +layout: global +title: In-Process Python UDFs +displayTitle: In-Process Python UDFs +license: | + Licensed to the Apache Software Foundation (ASF) under one or more + contributor license agreements. See the NOTICE file distributed with + this work for additional information regarding copyright ownership. + The ASF licenses this file to You under the Apache License, Version 2.0 + (the "License"); you may not use this file except in compliance with + the License. You may obtain a copy of the License at + + http://www.apache.org/licenses/LICENSE-2.0 + + Unless required by applicable law or agreed to in writing, software + distributed under the License is distributed on an "AS IS" BASIS, + WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + See the License for the specific language governing permissions and + limitations under the License. +--- + +* Table of contents +{:toc} + +## Runtime and result contract + +Each executor owns a dedicated interpreter thread. The plugin initializes the +interpreter on that thread, and task calls and shutdown are dispatched to the +same thread. Calls from concurrent tasks are queued on the interpreter thread. +One task per executor is recommended for throughput, but is not a correctness requirement. +Application-level Python parallelism comes from multiple executor JVMs. +The plugin configures JEP's process-wide interpreter with hash seed `0`, matching +Spark's default Python worker seed. It must initialize before any other JEP user in +the JVM. The seed cannot change between SparkContexts in the same process; a custom +worker `PYTHONHASHSEED` does not override this embedded-runtime setting. + +Task cancellation cannot safely stop arbitrary native Python code. An interrupted +caller waits for the current invocation to finish before freeing the Arrow CDI +structures, then restores its interrupt status. A UDF that never returns can +therefore prevent its task from completing cancellation and block every subsequent +in-process UDF on that executor, including calls from other tasks, jobs, and sessions. +Recovery from a permanently hung invocation requires replacing the executor process. +Plugin shutdown stops accepting new calls and waits up to five seconds for the interpreter thread. If a call is +still running, cleanup stays queued behind it; its memory remains live until the +call returns or the process exits. Shutdown does not forcibly interrupt native +code. A new interpreter cannot start until the previous one has fully stopped. + +A scalar UDF must return a `pyarrow.Array` with exactly one element per input row. +The runtime checks the result type against the declared Spark type, including +nested fields, decimal scale, and timestamp unit/timezone. Value types must match +exactly: use an explicit PyArrow cast in the UDF for numeric or other conversions. +Nested field nullability may differ if the actual values satisfy the declared nullability. Sliced results, including nested +child slices, are copied to remove offsets that Arrow Java's CDI importer cannot +read. Compatible results retain zero-copy transfer. + +The API produces a regular `PythonUDF` expression with an in-process evaluation +type. Spark's existing `ArrowEvalPython` planning rules handle aggregation, +nested calls, nondeterminism, and filter/limit pushdown. `ArrowEvalPythonExec` selects +an in-process evaluator factory for this evaluation type, reusing the projection, +row queue, result join, and partition-evaluator path. Ordinary Python UDFs continue +to use Python workers. + +`maxRecordsPerBatch <= 0` means no row-count limit. The independent +`spark.sql.execution.arrow.maxBytesPerBatch` limit still applies when positive. +Only UDF arguments are converted to Arrow. Other columns stay in Spark rows, +buffered in a spillable queue until the results are joined back. Duplicate nested +field names in UDF arguments or declared results are rejected before Arrow Java +reads their buffers. + +Each batch uses fresh input buffers. A Python function may retain an input array; +later batches do not overwrite it. Retained arrays keep native memory alive, so +functions should release them when no longer needed. JVM input vectors and result +vectors are released on task completion, early termination and failure. + +UDF deserialization uses PySpark's bundled cloudpickle. Each task registers its +own function instance once and passes a small handle for subsequent batches. +Task completion queues release of the registered function and its closure state. Imported +Python modules still share executor-wide state. Extra site-packages paths are +processed with `site.addsitedir` before loading the runtime bridge, including `.pth` +files. Configured directories and newly discovered `.pth` paths precede system +paths. Already imported modules cannot be replaced by changing the search path. + +Spark broadcasts, accumulators, `SparkContext.addPyFile`, and Python `TaskContext` +are not supported by this embedded runtime. Captured broadcast and accumulator +objects are rejected during serialization; functions must not access them through +imported modules either. Install modules on executors before startup, optionally +using `spark.inprocess.python.sitePackages`. SQL registration through +`spark.udf.register` is not supported and is rejected at registration time. +Spark Connect does not support this execution mode; both client SQL registration and +server planning reject it. The decorator accepts a `DataType` or a DDL string; DDL +strings are parsed lazily with the active Spark session. It exposes `func`, `returnType`, +`evalType`, `deterministic`, and `asNondeterministic()` along with the function's name +and docstring. +Functions must receive at least one input column (a literal also works) to determine +the batch length. Positional and keyword arguments are supported. Functions are +serialized on first use, so globals can be defined or rebound after decoration +and before that first call. The driver's Python major.minor +version must match the embedded interpreter; registration checks this before +unpickling. Python exceptions, including `SystemExit` during deserialization or +execution, are converted into task failures. Tracebacks honor the query's +`spark.sql.execution.pyspark.udf.hideTraceback.enabled` and +`spark.sql.execution.pyspark.udf.simplifiedTraceback.enabled` settings. Native process +termination remains outside this exception handling. + +## Overview + +In-process Python UDFs embed CPython directly into the Spark executor JVM using +[jep (Java Embedded Python)](https://github.com/ninia/jep), eliminating the IPC overhead of +standard Python UDFs and pandas UDFs. Data is passed to Python as +[PyArrow](https://arrow.apache.org/docs/python/) arrays via the +[Arrow C Data Interface](https://arrow.apache.org/docs/format/CDataInterface.html) — zero-copy +for compatible input and output buffers. Row-to-Arrow conversion and normalization +of sliced results still copy data. + +**Use `inprocess_udf` when:** +- You are already using `pandas_udf` for vectorized transformations and want lower latency. +- Your UDF operates on Arrow/PyArrow arrays (e.g. using `pyarrow.compute`). +- You can deploy enough executor JVMs for Python parallelism (see [Requirements](#requirements)). + +**Stick with `pandas_udf` or `udf` when:** +- You need pandas Series semantics in your UDF logic. +- You need concurrent Python invocations within a single executor. +- You are not able to install jep on executors. + +--- + +## Quick Start + +### 1. Install dependencies + +```bash +pip install "jep>=4.3.2" pyarrow cloudpickle +``` + +JEP and `org.apache.arrow:arrow-c-data` are provided dependencies and are not +bundled with Spark. Supply their JARs on the driver/executor classpaths before +starting Spark, and make the JEP native library available. Use an `arrow-c-data` +version matching Spark's Arrow Java version. Installing the Python packages alone +does not supply the Arrow Java CDI JAR. + +Building JEP from source requires a JDK, a C compiler, and development headers for +the Python version being embedded (for example, `python3.12-dev` on Ubuntu with +Python 3.12). These headers are build dependencies; running a prebuilt compatible +JEP installation does not require the development package. The corresponding +Python shared library must remain available at runtime. + +### 2. Register the plugin + +```python +spark = SparkSession.builder \ + .config("spark.plugins", + "org.apache.spark.sql.execution.python.InProcessPythonPlugin") \ + .config("spark.executor.cores", "1") \ + .config("spark.task.cpus", "1") \ + .getOrCreate() +``` + +### 3. Write and call a UDF + +```python +import pyarrow.compute as pc +from pyspark.inprocess.udf import inprocess_udf +from pyspark.sql.types import LongType + +@inprocess_udf(return_type=LongType()) +def double(x): + return pc.multiply(x, 2) + +df = spark.range(10) +df.select(double(df["id"])).show() +``` + +The function receives a `pa.Array` for each input column and must return a `pa.Array`. + +--- + +## Examples + +### String transformation + +```python +import pyarrow.compute as pc +from pyspark.inprocess.udf import inprocess_udf +from pyspark.sql.types import StringType + +@inprocess_udf(return_type=StringType()) +def upper(s): + return pc.utf8_upper(s) + +df = spark.createDataFrame([("hello",), ("world",)], ["text"]) +df.select(upper(df["text"])).show() +# +------------+ +# |upper(text) | +# +------------+ +# |HELLO | +# |WORLD | +# +------------+ +``` + +### Multi-column UDF + +A UDF receives one `pa.Array` argument per input column: + +```python +import pyarrow.compute as pc +from pyspark.inprocess.udf import inprocess_udf +from pyspark.sql.types import DoubleType + +@inprocess_udf(return_type=DoubleType()) +def weighted_sum(x, y): + return pc.add(pc.multiply(x, 0.6), pc.multiply(y, 0.4)) + +df = spark.createDataFrame([(1.0, 2.0), (3.0, 4.0)], ["x", "y"]) +df.select(weighted_sum(df["x"], df["y"])).show() +``` + +### Closure capture + +Free variables are captured by cloudpickle and frozen into the serialized UDF. The captured +value is evaluated once at UDF definition time and shipped with the function to every executor: Review Comment: Updated this to say the value is captured at first use and subsequent calls reuse the cached serialization. The first-use regression verifies that rebinding before the first call is picked up and rebinding afterward does not change later calls. -- This is an automated message from the Apache Git Service. To respond to the message, please log on to GitHub and use the URL above to go to the specific comment. To unsubscribe, e-mail: [email protected] For queries about this service, please contact Infrastructure at: [email protected] --------------------------------------------------------------------- To unsubscribe, e-mail: [email protected] For additional commands, e-mail: [email protected]
