viirya commented on code in PR #58978: URL: https://github.com/apache/spark/pull/58978#discussion_r4090582848
########## python/pyspark/inprocess/udf.py: ########## @@ -0,0 +1,196 @@ +# +# 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. +# + +""" +Python API for in-process UDF registration. + +Usage:: + + import pyarrow.compute as pc + from pyspark.inprocess import inprocess_udf + from pyspark.sql.types import LongType + + @inprocess_udf(return_type=LongType()) + def double(x): + # x is a pa.Array; return a pa.Array + return pc.multiply(x, 2) + + df.select(double(df.value)).show() +""" + +import io +import sys +from typing import Callable + +import pyarrow as pa + +from pyspark import Accumulator, Broadcast, cloudpickle +from pyspark.sql.types import ( + BooleanType, + ByteType, + DataType, + DoubleType, + FloatType, + IntegerType, + LongType, + ShortType, +) + +# Map from Spark SQL DataType to PyArrow type for output type enforcement. +_SPARK_TO_ARROW: dict = { + LongType(): pa.int64(), + IntegerType(): pa.int32(), + DoubleType(): pa.float64(), + FloatType(): pa.float32(), + BooleanType(): pa.bool_(), + ShortType(): pa.int16(), + ByteType(): pa.int8(), +} + + +class _InProcessPickler(cloudpickle.CloudPickler): + def reducer_override(self, obj): + if isinstance(obj, (Broadcast, Accumulator)): + raise TypeError("In-process UDFs do not support Spark broadcasts or accumulators") + return super().reducer_override(obj) + + +def _serialize_udf(func: Callable) -> bytes: + buffer = io.BytesIO() + _InProcessPickler(buffer).dump(func) + return buffer.getvalue() + + +class InProcessUDFWrapper: + """ + Wraps a Python function as an in-process UDF. + + Returned by ``@inprocess_udf``. Calling an instance with Spark ``Column`` + arguments creates a ``Column`` expression backed by ``PythonUDF`` + on the JVM side. + """ + + def __init__(self, func: Callable, return_type: DataType, deterministic: bool = True) -> None: + self._return_type: DataType = return_type + self._deterministic: bool = deterministic + self._name: str = getattr(func, "__name__", "inprocess_udf") + + # Wrap the function to cast its output to the declared return type. + # This handles the case where the UDF's input column type differs from + # the declared return type (e.g. input is int64, return_type is IntegerType). + arrow_type = _SPARK_TO_ARROW.get(return_type) + if arrow_type is not None: + + def _wrapped(*args, _fn=func, _atype=arrow_type): + result = _fn(*args) + if not isinstance(result, pa.Array): + raise TypeError("In-process UDF must return a pyarrow.Array") + if result.type != _atype: + result = result.cast(_atype) + return result + + self._serialized: bytes = _serialize_udf(_wrapped) + else: + self._serialized = _serialize_udf(func) + + def __call__(self, *cols): + """ + Create a ``Column`` expression invoking this UDF with the given columns. + + Args: + *cols: Spark ``Column`` objects (e.g. ``df.value``, ``col("x")``) + + Returns: + pyspark.sql.Column + """ + from pyspark import SparkContext + from pyspark.sql.classic.column import _to_java_column + from pyspark.sql.column import Column + + sc = SparkContext._active_spark_context + if sc is None: + raise RuntimeError( + "No active SparkContext. Start a SparkSession before calling an inprocess_udf." + ) + + jvm = sc._jvm + + # Convert Python Column objects to JVM Column objects + jcols = [_to_java_column(c) for c in cols] + + # Build a Java ArrayList (py4j vararg spread doesn't work with Arrays.asList) + jlist = jvm.java.util.ArrayList() + for jcol in jcols: + jlist.add(jcol) + + # Use the existing PythonUDF planning contracts with an in-process eval type. + jcol = jvm.org.apache.spark.sql.execution.python.InProcessPythonUDFBuilder.build( + self._name, + self._serialized, + self._return_type.json(), + jlist, + self._deterministic, + "%d.%d" % sys.version_info[:2], + ) + + return Column(jcol) + + +def inprocess_udf(return_type: DataType, deterministic: bool = True) -> Callable: + """ + Decorator to register a Python function as an in-process UDF. + + The decorated function receives one ``pa.Array`` per input column and must + return a single ``pa.Array`` of the declared ``return_type``. + + The result must have the same length as the input batch and its Arrow type + must match the declared Spark type, including nested fields and timestamp + timezone. Nested nullability may be widened, but actual nulls cannot be returned + in non-nullable fields. Numeric and boolean results are cast to the declared + primitive type. Sliced results are copied when required by Arrow Java. + + Spark broadcasts, accumulators, and ``SparkContext.addPyFile`` are unsupported. + Install dependencies on executors before starting Spark. The driver's Python + major.minor version must match the embedded interpreter. + + Args: + return_type: Spark SQL DataType for the UDF return value + deterministic: Whether this UDF produces the same output for the same input. + Set to ``False`` for UDFs that use randomness, external state, + or other sources of non-determinism so the optimizer does not + deduplicate or reorder calls to this UDF. Default: ``True``. + + Returns: + Decorator that wraps the function as an ``InProcessUDFWrapper`` + + Example:: + + @inprocess_udf(return_type=LongType()) + def double(x): + import pyarrow.compute as pc + return pc.multiply(x, 2) + + @inprocess_udf(return_type=LongType(), deterministic=False) + def random_noise(x): + import pyarrow as pa, numpy as np + return pa.array(np.random.randint(0, 100, len(x)), type=pa.int64()) + """ + + def decorator(func: Callable) -> InProcessUDFWrapper: Review Comment: Zero-argument scalar definitions are now rejected, and calling a wrapper without any column arguments is rejected too. Functions must receive an input column or literal to establish the batch length. Updated the tests and documentation accordingly. ########## python/pyspark/inprocess/udf.py: ########## @@ -0,0 +1,196 @@ +# +# 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. +# + +""" +Python API for in-process UDF registration. + +Usage:: + + import pyarrow.compute as pc + from pyspark.inprocess import inprocess_udf + from pyspark.sql.types import LongType + + @inprocess_udf(return_type=LongType()) + def double(x): + # x is a pa.Array; return a pa.Array + return pc.multiply(x, 2) + + df.select(double(df.value)).show() +""" + +import io +import sys +from typing import Callable + +import pyarrow as pa + +from pyspark import Accumulator, Broadcast, cloudpickle +from pyspark.sql.types import ( + BooleanType, + ByteType, + DataType, + DoubleType, + FloatType, + IntegerType, + LongType, + ShortType, +) + +# Map from Spark SQL DataType to PyArrow type for output type enforcement. +_SPARK_TO_ARROW: dict = { + LongType(): pa.int64(), + IntegerType(): pa.int32(), + DoubleType(): pa.float64(), + FloatType(): pa.float32(), + BooleanType(): pa.bool_(), + ShortType(): pa.int16(), + ByteType(): pa.int8(), +} + + +class _InProcessPickler(cloudpickle.CloudPickler): + def reducer_override(self, obj): + if isinstance(obj, (Broadcast, Accumulator)): + raise TypeError("In-process UDFs do not support Spark broadcasts or accumulators") + return super().reducer_override(obj) + + +def _serialize_udf(func: Callable) -> bytes: + buffer = io.BytesIO() + _InProcessPickler(buffer).dump(func) + return buffer.getvalue() + + +class InProcessUDFWrapper: + """ + Wraps a Python function as an in-process UDF. + + Returned by ``@inprocess_udf``. Calling an instance with Spark ``Column`` + arguments creates a ``Column`` expression backed by ``PythonUDF`` + on the JVM side. + """ + + def __init__(self, func: Callable, return_type: DataType, deterministic: bool = True) -> None: + self._return_type: DataType = return_type + self._deterministic: bool = deterministic + self._name: str = getattr(func, "__name__", "inprocess_udf") + + # Wrap the function to cast its output to the declared return type. + # This handles the case where the UDF's input column type differs from + # the declared return type (e.g. input is int64, return_type is IntegerType). + arrow_type = _SPARK_TO_ARROW.get(return_type) + if arrow_type is not None: + + def _wrapped(*args, _fn=func, _atype=arrow_type): + result = _fn(*args) + if not isinstance(result, pa.Array): + raise TypeError("In-process UDF must return a pyarrow.Array") + if result.type != _atype: + result = result.cast(_atype) + return result + + self._serialized: bytes = _serialize_udf(_wrapped) Review Comment: The wrapper now retains the callable and serializes it on first use. Added a regression where the referenced global is defined after decoration. The serialized command remains stable after first use, matching the existing UDF behavior and preserving semantic equality across subsequent calls. ########## python/pyspark/inprocess/udf.py: ########## @@ -0,0 +1,196 @@ +# +# 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. +# + +""" +Python API for in-process UDF registration. + +Usage:: + + import pyarrow.compute as pc + from pyspark.inprocess import inprocess_udf + from pyspark.sql.types import LongType + + @inprocess_udf(return_type=LongType()) + def double(x): + # x is a pa.Array; return a pa.Array + return pc.multiply(x, 2) + + df.select(double(df.value)).show() +""" + +import io +import sys +from typing import Callable + +import pyarrow as pa + +from pyspark import Accumulator, Broadcast, cloudpickle +from pyspark.sql.types import ( + BooleanType, + ByteType, + DataType, + DoubleType, + FloatType, + IntegerType, + LongType, + ShortType, +) + +# Map from Spark SQL DataType to PyArrow type for output type enforcement. +_SPARK_TO_ARROW: dict = { + LongType(): pa.int64(), + IntegerType(): pa.int32(), + DoubleType(): pa.float64(), + FloatType(): pa.float32(), + BooleanType(): pa.bool_(), + ShortType(): pa.int16(), + ByteType(): pa.int8(), +} + + +class _InProcessPickler(cloudpickle.CloudPickler): + def reducer_override(self, obj): + if isinstance(obj, (Broadcast, Accumulator)): + raise TypeError("In-process UDFs do not support Spark broadcasts or accumulators") + return super().reducer_override(obj) + + +def _serialize_udf(func: Callable) -> bytes: + buffer = io.BytesIO() + _InProcessPickler(buffer).dump(func) + return buffer.getvalue() + + +class InProcessUDFWrapper: Review Comment: Added an explicit registration-time rejection for `InProcessUDFWrapper`, so it cannot fall through to the plain-callable `StringType`/`SQL_BATCHED_UDF` path. SQL registration remains unsupported in this revision and is documented as such. Added a regression for the early error; the existing Python UDF suite also passes. ########## python/pyspark/inprocess/runtime.py: ########## @@ -0,0 +1,161 @@ +# +# 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 + +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 + +_UDF_TRACEBACK_SENTINEL = "__INPROCESS_UDF_TRACEBACK__:" +_udfs: dict = {} + + +def _inprocess_register(handle, serialized_udf, return_type_json, timezone, python_version): + 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's PyJArray does not implement the buffer protocol. Convert once per task. + func = cloudpickle.loads(bytes(b & 0xFF for b in serialized_udf)) Review Comment: Registration now bulk-copies the command into a direct ByteBuffer on the task thread, and Python unpickles through `memoryview` using JEP's buffer protocol. This removes the per-byte JNI loop. Each task still deserializes its own function; no shared-function cache was added. A 4 MiB closure regression passes across multiple tasks, but I haven't rerun performance benchmarks yet. ########## sql/core/src/main/scala/org/apache/spark/sql/execution/python/InProcessArrowEvalExec.scala: ########## @@ -0,0 +1,221 @@ +/* + * 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. + */ + +package org.apache.spark.sql.execution.python + +import java.util.UUID + +import scala.collection.mutable.ArrayBuffer +import scala.jdk.CollectionConverters._ + +import org.apache.arrow.c.{ArrowArray, ArrowSchema} +import org.apache.arrow.util.AutoCloseables +import org.apache.arrow.vector.VectorSchemaRoot + +import org.apache.spark.TaskContext +import org.apache.spark.rdd.RDD +import org.apache.spark.sql.catalyst.InternalRow +import org.apache.spark.sql.catalyst.expressions.{Attribute, AttributeSet, Expression, JoinedRow, PythonUDF, UnsafeProjection} +import org.apache.spark.sql.execution.{SparkPlan, UnaryExecNode} +import org.apache.spark.sql.execution.arrow.ArrowWriter +import org.apache.spark.sql.types.{StructField, StructType} +import org.apache.spark.sql.util.ArrowUtils +import org.apache.spark.sql.vectorized.{ArrowColumnVector, ColumnarBatch, ColumnVector} +import org.apache.spark.util.Utils + +/** + * Evaluates scalar Python UDFs using Arrow CDI in the executor process. Only UDF arguments + * are converted to Arrow. Original rows are buffered in a spillable queue and joined with + * the results. Each batch owns its Arrow buffers so Python can safely retain input arrays. + */ +case class InProcessArrowEvalExec( + udfs: Seq[PythonUDF], + resultAttrs: Seq[Attribute], + child: SparkPlan) extends UnaryExecNode { + + override def output: Seq[Attribute] = child.output ++ resultAttrs + + override def producedAttributes: AttributeSet = AttributeSet(resultAttrs) + + override protected def doExecute(): RDD[InternalRow] = { Review Comment: Implemented this direction: `ArrowEvalPythonExec` selects `InProcessArrowEvalPythonEvaluatorFactory`, which inherits the shared projection, queue, join, and argument extraction. Removed the separate strategy branch. Added keyword binding, metric updates, and JVM schema checks. Tests cover both partition-evaluator settings and a columnar-capable cache child; the in-process eval type explicitly uses row execution rather than entering the worker columnar evaluator. ########## sql/core/src/main/scala/org/apache/spark/sql/execution/python/InProcessPythonRuntime.scala: ########## @@ -0,0 +1,214 @@ +/* + * 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. + */ + +package org.apache.spark.sql.execution.python + +import java.util.concurrent.{Callable, ExecutionException, ExecutorService, TimeUnit} +import java.util.concurrent.locks.ReentrantLock + +import scala.jdk.CollectionConverters._ + +import jep.{JepException, SharedInterpreter} + +import org.apache.spark.TaskContext +import org.apache.spark.api.python.PythonException +import org.apache.spark.internal.Logging +import org.apache.spark.util.{ThreadUtils, Utils} + +/** + * Owns one interpreter on a dedicated thread per executor. JEP requires construction, + * invocation and close to happen on the same thread, even when Spark tasks run serially. + */ +private[python] object InProcessPythonRuntime extends Logging { + val SITE_PACKAGES_CONFIG = "spark.inprocess.python.sitePackages" + + // Access to the executor is serialized by onInterpreterThread and shutdown. The interpreter + // itself is accessed only by the executor's thread. + private val interpreterLock = new ReentrantLock() + private var executor: ExecutorService = _ + private var interp: SharedInterpreter = _ + private val TRACEBACK_SENTINEL = "__INPROCESS_UDF_TRACEBACK__:" + + private def withInterpreterLock[T](cancellable: Boolean)(body: => T): T = { + if (cancellable) { + val context = Option(TaskContext.get()) + context.foreach(_.killTaskIfInterrupted()) + // Poll the task state as cancellation need not interrupt the Java thread. + while (!interpreterLock.tryLock(100, TimeUnit.MILLISECONDS)) { + context.foreach(_.killTaskIfInterrupted()) + } + } else { + interpreterLock.lock() + } + try { + if (cancellable) Option(TaskContext.get()).foreach(_.killTaskIfInterrupted()) + body + } finally { + interpreterLock.unlock() + } + } + + /** + * Wait for native code to finish even if the task is interrupted. Returning early would let + * the task free CDI pointers that Python may still be accessing. Restore the interruption + * afterwards so Spark can observe cancellation. Arbitrary Python code cannot be forcibly + * interrupted safely in the executor process. + */ + private[python] def onInterpreterThread[T](body: => T): T = { + runOnInterpreterThread(cancellable = true)(body) + } + + private def runOnInterpreterThread[T](cancellable: Boolean)(body: => T): T = + withInterpreterLock(cancellable) { + if (executor == null) { + executor = ThreadUtils.newDaemonSingleThreadExecutor("inprocess-python") + } + val future = executor.submit(new Callable[T] { + override def call(): T = body + }) + var interrupted = false + try { + var result: Option[T] = None + while (result.isEmpty) { + try { + result = Some(future.get()) + } catch { + case _: InterruptedException => interrupted = true + case e: ExecutionException => throw e.getCause + } + } + if (cancellable) Option(TaskContext.get()).foreach(_.killTaskIfInterrupted()) + result.get + } finally { + if (interrupted) Thread.currentThread().interrupt() + } + } + + private def initializeInterpreter(sitePackages: Seq[String]): Unit = { + if (interp == null) { + val candidate = new SharedInterpreter() + try { + // Configure paths before importing the bridge and its dependencies. + if (sitePackages.nonEmpty) { + candidate.set("_site_packages", sitePackages.asJava) + candidate.eval("import sys; sys.path.extend(list(_site_packages)); del _site_packages") Review Comment: Configured paths are now made absolute and processed with `site.addsitedir` before importing the runtime bridge. Those paths and newly discovered `.pth` paths precede system paths. The integration test loads a helper through a `.pth` file while a conflicting helper is present on the initial search path, and verifies that the configured copy wins. The docs also note that changing paths cannot replace already-imported modules. -- 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]
