viirya commented on code in PR #58978: URL: https://github.com/apache/spark/pull/58978#discussion_r4101610539
########## sql/core/src/main/scala/org/apache/spark/sql/execution/python/InProcessPythonRuntime.scala: ########## @@ -0,0 +1,258 @@ +/* + * 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.nio.ByteBuffer +import java.util.concurrent.{Callable, ExecutionException, TimeoutException, TimeUnit} + +import scala.jdk.CollectionConverters._ + +import jep.{JepException, SharedInterpreter} + +import org.apache.spark.{TaskContext, TaskKilledException} +import org.apache.spark.api.python.PythonException +import org.apache.spark.internal.Logging +import org.apache.spark.util.{ThreadUtils, Utils} + +/** Owns one interpreter generation per executor plugin lifecycle. */ +private[python] object InProcessPythonRuntime extends Logging { + val SITE_PACKAGES_CONFIG = "spark.inprocess.python.sitePackages" + private val TRACEBACK_SENTINEL = "__INPROCESS_UDF_TRACEBACK__:" + private var active: InterpreterSession = _ + + def initialize(sitePackages: Seq[String] = Seq.empty): Unit = synchronized { + if (active != null && !active.isTerminated) { + require(active.isRunning && active.sitePackages == sitePackages, + "In-process Python is stopping or already initialized with different sitePackages") + } else { + val candidate = new InterpreterSession(sitePackages) + try { + candidate.initialize() + active = candidate + } catch { + case t: Throwable => Utils.tryWithSafeFinally { throw t } { candidate.shutdown() } + } + } + } + + def currentSession: InterpreterSession = synchronized { + checkState(active != null && active.isRunning) + active + } + + def shutdown(): Unit = { + val session = synchronized { active } + if (session != null) session.shutdown() + } + + private def checkState(running: Boolean): Unit = { + checkState(running, "In-process Python is not running; initialize the executor plugin first") + } + + private def checkState(running: Boolean, message: String): Unit = { + if (!running) throw new IllegalStateException(message) + } + + /** + * Tasks retain this generation, so stale tasks cannot enter a later SparkContext's interpreter. + * Lifecycle operations only hold the monitor while enqueueing work, never while running Python. + */ + private[python] class InterpreterSession(val sitePackages: Seq[String] = Seq.empty) { + private val executor = ThreadUtils.newDaemonSingleThreadExecutor("inprocess-python") + @volatile private var running = true + // Accessed only on the owning thread. + private var interp: SharedInterpreter = _ + + def isRunning: Boolean = running + def isTerminated: Boolean = executor.isTerminated + + private[python] def onInterpreterThread[T](body: => T): T = { + val context = Option(TaskContext.get()) + context.foreach(_.killTaskIfInterrupted()) + val gate = new Object + var started = false + var cancelled = false + val future = synchronized { + checkState(running) + executor.submit(new Callable[T] { + override def call(): T = { + gate.synchronized { + if (cancelled) throw new TaskKilledException("Cancelled before Python invocation") + started = true + } + body + } + }) + } + var interrupted = false + try { + while (true) { + val taskCancelled = context.exists(_.isInterrupted()) + if (interrupted || taskCancelled) { + val cancelledBeforeStart = gate.synchronized { + if (started) false else { + cancelled = true + future.cancel(false) + true + } + } + if (cancelledBeforeStart) { + context.foreach(_.killTaskIfInterrupted()) + throw new InterruptedException("Cancelled before Python invocation") + } + } + try { + val result = future.get(100, TimeUnit.MILLISECONDS) + context.foreach(_.killTaskIfInterrupted()) + return result + } catch { + case _: TimeoutException => + case _: InterruptedException => interrupted = true + case e: ExecutionException => throw e.getCause + } + } + throw new IllegalStateException("Unreachable") + } finally { + // Once native work starts, wait for it even after cancellation: the caller still owns + // CDI structs that Python may use. Pending work, however, is safe to cancel immediately. + if (interrupted) Thread.currentThread().interrupt() + } + } + + def initialize(): Unit = onInterpreterThread { + val candidate = new SharedInterpreter() + try { + candidate.set("_site_packages", sitePackages.asJava) + candidate.exec( Review Comment: Both bootstrap scripts now use the same try/except BaseException guard and convert failures to RuntimeError. Added guard tests and an integration test where .pth processing raises KeyboardInterrupt, then verified that initialization can be retried and a UDF can run successfully. ########## sql/core/src/main/scala/org/apache/spark/sql/execution/python/InProcessPythonRuntime.scala: ########## @@ -0,0 +1,258 @@ +/* + * 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.nio.ByteBuffer +import java.util.concurrent.{Callable, ExecutionException, TimeoutException, TimeUnit} + +import scala.jdk.CollectionConverters._ + +import jep.{JepException, SharedInterpreter} + +import org.apache.spark.{TaskContext, TaskKilledException} +import org.apache.spark.api.python.PythonException +import org.apache.spark.internal.Logging +import org.apache.spark.util.{ThreadUtils, Utils} + +/** Owns one interpreter generation per executor plugin lifecycle. */ +private[python] object InProcessPythonRuntime extends Logging { + val SITE_PACKAGES_CONFIG = "spark.inprocess.python.sitePackages" + private val TRACEBACK_SENTINEL = "__INPROCESS_UDF_TRACEBACK__:" + private var active: InterpreterSession = _ + + def initialize(sitePackages: Seq[String] = Seq.empty): Unit = synchronized { + if (active != null && !active.isTerminated) { + require(active.isRunning && active.sitePackages == sitePackages, + "In-process Python is stopping or already initialized with different sitePackages") + } else { + val candidate = new InterpreterSession(sitePackages) + try { + candidate.initialize() + active = candidate + } catch { + case t: Throwable => Utils.tryWithSafeFinally { throw t } { candidate.shutdown() } + } + } + } + + def currentSession: InterpreterSession = synchronized { + checkState(active != null && active.isRunning) + active + } + + def shutdown(): Unit = { + val session = synchronized { active } + if (session != null) session.shutdown() + } + + private def checkState(running: Boolean): Unit = { + checkState(running, "In-process Python is not running; initialize the executor plugin first") + } + + private def checkState(running: Boolean, message: String): Unit = { + if (!running) throw new IllegalStateException(message) + } + + /** + * Tasks retain this generation, so stale tasks cannot enter a later SparkContext's interpreter. + * Lifecycle operations only hold the monitor while enqueueing work, never while running Python. + */ + private[python] class InterpreterSession(val sitePackages: Seq[String] = Seq.empty) { + private val executor = ThreadUtils.newDaemonSingleThreadExecutor("inprocess-python") + @volatile private var running = true + // Accessed only on the owning thread. + private var interp: SharedInterpreter = _ + + def isRunning: Boolean = running + def isTerminated: Boolean = executor.isTerminated + + private[python] def onInterpreterThread[T](body: => T): T = { + val context = Option(TaskContext.get()) + context.foreach(_.killTaskIfInterrupted()) + val gate = new Object + var started = false + var cancelled = false + val future = synchronized { + checkState(running) + executor.submit(new Callable[T] { + override def call(): T = { + gate.synchronized { + if (cancelled) throw new TaskKilledException("Cancelled before Python invocation") + started = true + } + body + } + }) + } + var interrupted = false + try { + while (true) { + val taskCancelled = context.exists(_.isInterrupted()) + if (interrupted || taskCancelled) { + val cancelledBeforeStart = gate.synchronized { + if (started) false else { + cancelled = true + future.cancel(false) + true + } + } + if (cancelledBeforeStart) { + context.foreach(_.killTaskIfInterrupted()) + throw new InterruptedException("Cancelled before Python invocation") + } + } + try { + val result = future.get(100, TimeUnit.MILLISECONDS) + context.foreach(_.killTaskIfInterrupted()) + return result + } catch { + case _: TimeoutException => + case _: InterruptedException => interrupted = true + case e: ExecutionException => throw e.getCause + } + } + throw new IllegalStateException("Unreachable") + } finally { + // Once native work starts, wait for it even after cancellation: the caller still owns + // CDI structs that Python may use. Pending work, however, is safe to cancel immediately. + if (interrupted) Thread.currentThread().interrupt() + } + } + + def initialize(): Unit = onInterpreterThread { + val candidate = new SharedInterpreter() + try { + candidate.set("_site_packages", sitePackages.asJava) + candidate.exec( + """import os, site, sys + |_configured = [os.path.abspath(p) for p in _site_packages] + |_before = set(sys.path) + |for _path in _configured: + | site.addsitedir(_path) + |_added = [p for p in sys.path if p not in _before and p not in _configured] + |_preferred = list(dict.fromkeys(_configured + _added)) + |sys.path[:] = _preferred + [p for p in sys.path if p not in _preferred] + |del _site_packages, _configured, _before, _added, _preferred + |""".stripMargin) + candidate.exec("from pyspark.inprocess.runtime import " + + "_inprocess_invoke, _inprocess_register, _inprocess_release, _udfs") + interp = candidate + } catch { + case t: Throwable => Utils.tryWithSafeFinally { throw t } { candidate.close() } + } + } + + /** Enqueue cleanup after outstanding calls without creating an executor or waiting. */ + def release(handles: Seq[String]): Unit = synchronized { + if (running && handles.nonEmpty) { + executor.submit(new Runnable { + override def run(): Unit = { + if (interp != null) interp.invoke("_inprocess_release", handles.asJava) + } + }) + } + // During shutdown the queued close clears all remaining handles. + } + + /** A timeout bounds plugin stop, not native execution or CDI buffer ownership. */ + def shutdown(waitMillis: Long = 5000L): Unit = { + synchronized { + if (running) { + running = false + executor.submit(new Runnable { + override def run(): Unit = { + if (interp != null) { + try { + interp.exec("_udfs.clear()") + } finally { + try { interp.close() } finally { interp = null } + } + } + } + }) + executor.shutdown() + } + } + try { + if (!executor.awaitTermination(waitMillis, TimeUnit.MILLISECONDS)) { + logWarning("In-process Python is still stopping; native work and its buffers " + + "remain alive until the invocation finishes or the process exits.") + } + } catch { + case _: InterruptedException => Thread.currentThread().interrupt() + } + } + + def register( + handle: String, + serializedUdf: Array[Byte], + returnTypeJson: String, + timeZoneId: String, + pythonVersion: String, + largeVarTypes: Boolean): Unit = { + // Bulk-copy on the task thread. JEP's PyJBuffer supports memoryview without per-byte JNI. + val command = ByteBuffer.allocateDirect(serializedUdf.length) + command.put(serializedUdf).flip() + onInterpreterThread { + withPythonException { + interp.invoke("_inprocess_register", handle, command, returnTypeJson, timeZoneId, + pythonVersion, java.lang.Boolean.valueOf(largeVarTypes)) + } + } + } + + def invoke( + handle: String, + inputArrayPtrs: Array[Long], + inputSchemaPtrs: Array[Long], + outputArrayAddr: Long, + outputSchemaAddr: Long, + expectedRows: Int, + argumentNames: Array[String]): Long = onInterpreterThread { + val start = System.nanoTime() + val arrayPtrs = inputArrayPtrs.map(java.lang.Long.valueOf).toSeq.asJava + val schemaPtrs = inputSchemaPtrs.map(java.lang.Long.valueOf).toSeq.asJava + withPythonException { + interp.invoke("_inprocess_invoke", handle, arrayPtrs, schemaPtrs, + java.lang.Long.valueOf(outputArrayAddr), java.lang.Long.valueOf(outputSchemaAddr), + java.lang.Integer.valueOf(expectedRows), argumentNames.toSeq.asJava) + } + (System.nanoTime() - start) / 1000000 Review Comment: Invocation durations now remain in nanoseconds until the evaluator updates the metric. The timer carries the sub-millisecond remainder across calls, so fast batches accumulate instead of each contributing zero. Initialization uses the same helper. Added a regression test for repeated sub-millisecond durations. ########## python/pyspark/inprocess/udf.py: ########## @@ -0,0 +1,183 @@ +# +# 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 inspect import getfullargspec +from typing import Any, Callable, Optional, Union + +from pyspark import Accumulator, Broadcast, cloudpickle +from pyspark.errors import PySparkValueError +from pyspark.sql.column import Column +from pyspark.sql.types import DataType + + +class _InProcessPickler(cloudpickle.CloudPickler): + def reducer_override(self, obj: Any) -> Any: + 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") + + argspec = getfullargspec(func) + if not argspec.args and argspec.varargs is None and not argspec.kwonlyargs: + raise PySparkValueError( + errorClass="INVALID_PANDAS_UDF", + messageParameters={"detail": "0-arg inprocess_udfs are not supported."}, + ) + self._func = func + self._serialized: Optional[bytes] = None + + def _serialize(self) -> bytes: + if self._serialized is None: + self._serialized = _serialize_udf(self._func) + return self._serialized + + def __call__(self, *cols: Union[Column, str], **kwargs: Union[Column, str]) -> Column: + """ + 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 + + 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 + assert jvm is not None + + # Convert Python Column objects to JVM Column objects + if not cols and not kwargs: + raise PySparkValueError( + errorClass="INVALID_PANDAS_UDF", + messageParameters={"detail": "An inprocess_udf requires at least one argument."}, + ) + jcols = [_to_java_column(c) for c in cols] + jcols.extend( + jvm.PythonSQLUtils.namedArgumentExpression(name, _to_java_column(value)) + for name, value in kwargs.items() + ) + + # 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._serialize(), + self._return_type.json(), Review Comment: Added lazy DDL parsing through _parse_datatype_string, so decoration does not require an active session. Unsupported return-type values now fail at decoration. The wrapper also exposes func, returnType, evalType, deterministic, and asNondeterministic(), and preserves the function name and docstring. Added unit and integration coverage. ########## python/pyspark/inprocess/runtime.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. +# + + +"""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 + +_UDF_TRACEBACK_SENTINEL = "__INPROCESS_UDF_TRACEBACK__:" +_udfs: dict[str, tuple[Callable[..., pa.Array], pa.DataType]] = {} + + +def _inprocess_register( + handle: str, + serialized_udf: Any, + return_type_json: str, + timezone: str, + python_version: str, + large_var_types: 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, + ) + _udfs[handle] = (func, expected_type) + except BaseException: + # In JEP, an uncaught SystemExit can terminate the entire executor JVM. + raise RuntimeError(_UDF_TRACEBACK_SENTINEL + _traceback.format_exc()) 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 _check_nested_nulls(array: pa.Array, expected_type: pa.DataType) -> None: + def check_field(values: pa.Array, field: pa.Field) -> None: + if not field.nullable and values.null_count: + raise ValueError(f"In-process UDF returned nulls in non-nullable field {field.name}") + _check_nested_nulls(values, field.type) + + if pa.types.is_struct(expected_type): + # Children under a null parent do not contribute values to the result. + visible = pc.filter(array, pc.is_valid(array)) Review Comment: Required-field checks are now compiled at registration. Nullable subtrees without required descendants are skipped, and parent filtering only runs when parent nulls are present. Removed the unconditional map concatenation as well. Added checks that the common paths avoid filter/concat calls, while retaining coverage for hidden nulls, required descendants, and sliced maps. ########## python/pyspark/inprocess/bridge.py: ########## @@ -0,0 +1,38 @@ +# +# 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. +# + +""" Review Comment: Removed the unused module. The CDI bridge documentation remains with the Scala implementation. -- 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]
