viirya commented on code in PR #58978: URL: https://github.com/apache/spark/pull/58978#discussion_r4112404500
########## python/pyspark/inprocess/runtime.py: ########## @@ -0,0 +1,322 @@ +# +# 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 +from typing import Any, Callable, Iterable, Optional, Sequence + +import pyarrow as pa Review Comment: Added `require_minimum_pyarrow_version()` to the driver wrapper, executor bootstrap, and registration path. Tests mock an older PyArrow version to verify the standard `UNSUPPORTED_PACKAGE_VERSION` error; they do not require installing an older PyArrow. ########## sql/core/src/main/scala/org/apache/spark/sql/execution/python/InProcessArrowEvalPythonEvaluatorFactory.scala: ########## @@ -0,0 +1,233 @@ +/* + * 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.api.python.ChainedPythonFunctions +import org.apache.spark.sql.catalyst.InternalRow +import org.apache.spark.sql.catalyst.expressions.{Attribute, PythonUDF} +import org.apache.spark.sql.execution.arrow.ArrowWriter +import org.apache.spark.sql.execution.metric.SQLMetric +import org.apache.spark.sql.execution.python.EvalPythonExec.ArgumentMetadata +import org.apache.spark.sql.types.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. + */ +class InProcessArrowEvalPythonEvaluatorFactory( + childOutput: Seq[Attribute], + udfs: Seq[PythonUDF], + output: Seq[Attribute], + batchSize: Int, + maxBytes: Long, + timeZoneId: String, + largeVarTypes: Boolean, + hideTraceback: Boolean, + simplifiedTraceback: Boolean, + tracebackWithLocals: Boolean, + metrics: Map[String, SQLMetric]) + extends EvalPythonEvaluatorFactory(childOutput, udfs, output) { + + override protected def evaluate( + funcs: Seq[(ChainedPythonFunctions, Long)], + argMetas: Array[Array[ArgumentMetadata]], + rows: Iterator[InternalRow], + inputSchema: StructType, + context: TaskContext): Iterator[InternalRow] = { + ArrowUtils.failDuplicatedFieldNames(inputSchema) + val functions = funcs.map { case (chain, _) => + require(chain.funcs.size == 1, "In-process UDF chains must use separate evaluation nodes") + chain.funcs.head + } + val inputOrdinals = argMetas.map(_.map(_.offset)) + def checkCancellation(): Unit = context.killTaskIfInterrupted() + + val expectedFields = udfs.map { udf => + ArrowUtils.toArrowField("result", udf.dataType, true, timeZoneId, largeVarTypes) + } + val processingTime = new InProcessArrowEvalPythonEvaluatorFactory.NanosecondTimer( + metrics("pythonProcessingTime")) + val initTime = new InProcessArrowEvalPythonEvaluatorFactory.NanosecondTimer( + metrics("pythonInitTime")) + val arrowSchema = ArrowUtils.toArrowSchema(inputSchema, timeZoneId, largeVarTypes) + var runtime: InProcessPythonRuntime.InterpreterSession = null + val handles = functions.map(_ => UUID.randomUUID().toString) + var registered = false + var writer: ArrowWriter = null + val results = ArrayBuffer.empty[ArrowColumnVector] + var closed = false + val startedAt = System.nanoTime() + + def closeBatch(): Unit = { + val resources = ArrayBuffer.empty[AutoCloseable] + resources ++= results + results.clear() + if (writer != null) { + resources += writer.root + writer = null + } + AutoCloseables.close(resources.asJava) + } + + def close(): Unit = { + if (!closed) { + closed = true + metrics("pythonTotalTime") += (System.nanoTime() - startedAt) / 1000000 + Utils.tryWithSafeFinally { + closeBatch() + } { + if (registered) runtime.release(handles) + } + } + } + + context.addTaskCompletionListener[Unit](_ => close()) + + new Iterator[InternalRow] { + private var batchIter: Iterator[InternalRow] = Iterator.empty + + override def hasNext: Boolean = { + checkCancellation() + val available = !closed && (batchIter.hasNext || rows.hasNext) + if (!available) close() + available + } + + override def next(): InternalRow = { + if (!hasNext) throw new NoSuchElementException("End of in-process UDF input") + try { + if (!batchIter.hasNext) { + closeBatch() + if (!registered) { + runtime = InProcessPythonRuntime.currentSession + // Mark before registering so failure after any registration still cleans up. + registered = true + val start = System.nanoTime() Review Comment: Registration timing now starts on the interpreter thread, excluding queue time, using the same timing helper as invocation. Total timing starts when the iterator is first used, so an unconsumed iterator contributes no `pythonTotalTime`. Added regression tests for queued work and an unused iterator. ########## python/pyspark/inprocess/udf.py: ########## @@ -0,0 +1,225 @@ +# +# 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 functools import update_wrapper +from inspect import getfullargspec +from typing import Any, Callable, Optional, Union + +from pyspark import Accumulator, Broadcast, cloudpickle +from pyspark.errors import PySparkTypeError, PySparkValueError +from pyspark.sql.column import Column +from pyspark.sql.types import DataType, _parse_datatype_string +from pyspark.util import PythonEvalType + + +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, return_type: DataType) -> bytes: + buffer = io.BytesIO() + _InProcessPickler(buffer).dump((func, return_type)) + 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: Union[DataType, str], deterministic: bool = True + ) -> None: + if not isinstance(return_type, (DataType, str)): + raise PySparkTypeError( + errorClass="NOT_EXPECTED_TYPE", + messageParameters={ + "expected_type": "DataType or str", + "arg_name": "return_type", + "arg_type": type(return_type).__name__, + }, + ) + self._return_type = return_type + self._parsed_return_type: Optional[DataType] = None + self.evalType = PythonEvalType.SQL_SCALAR_ARROW_INPROCESS_UDF + 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: Review Comment: Switched the check to `inspect.signature(func).parameters`. This accepts `**kwargs`-only functions and correctly rejects callable instances whose bound `__call__` takes no arguments. Added coverage for both, including keyword-only invocation through Spark. ########## sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/PythonUDF.scala: ########## @@ -51,7 +51,8 @@ object PythonUDF { PythonEvalType.SQL_SCALAR_PANDAS_UDF, PythonEvalType.SQL_SCALAR_PANDAS_ITER_UDF, PythonEvalType.SQL_SCALAR_ARROW_UDF, - PythonEvalType.SQL_SCALAR_ARROW_ITER_UDF + PythonEvalType.SQL_SCALAR_ARROW_ITER_UDF, + PythonEvalType.SQL_SCALAR_ARROW_INPROCESS_UDF Review Comment: Excluded eval type 258 from the chaining-driven force-inline condition in `CollapseProject`. Added a regression test that preserves an expensive producer referenced more than once. The existing worker-UDF chaining tests also pass. ########## python/pyspark/inprocess/runtime.py: ########## @@ -0,0 +1,322 @@ +# +# 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 +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.util import _format_exception + +_UDF_TRACEBACK_SENTINEL = "__INPROCESS_UDF_TRACEBACK__:" +NullChecker = Callable[[pa.Array], None] +_udfs: dict[str, tuple[Callable[..., pa.Array], pa.DataType, NullChecker, bool, bool, bool]] = {} + + +def _inprocess_register( + handle: str, + serialized_udf: Any, + timezone: str, + python_version: str, + large_var_types: bool = False, + hide_traceback: bool = False, + simplified_traceback: bool = False, + traceback_with_locals: 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. + # Carry the type with the closure so driver-defined UDTs need no module import. + func, return_type = cloudpickle.loads(memoryview(serialized_udf)) + expected_type = to_arrow_type( + return_type, + timezone=timezone, + prefers_large_types=large_var_types, + error_on_duplicated_field_names_in_struct=True, + ) + if large_var_types: + expected_type = _large_binary_type(expected_type) Review Comment: Adopted this approach. Registration now exports the JVM's expected `Field` through CDI, and Python imports it with `pa.Field._import_from_c`. Removed the pickled return type, the timezone/large-type registration arguments, and the executor-side `to_arrow_type` and `_large_binary_type` conversion. Driver-side return-type validation remains. Output normalization and restoration of declared JVM metadata also remain, since returned arrays still need validation. Tests cover nested types, metadata/nullability, UDTs, temporal and Variant/spatial types, and failure cleanup. ########## sql/core/src/main/scala/org/apache/spark/sql/execution/python/ArrowEvalPythonExec.scala: ########## @@ -66,6 +66,7 @@ private[spark] class BatchIterator[T](iter: Iterator[T], batchSize: Int) * Following eval types are supported: * * <ul> + * <li> SQL_SCALAR_ARROW_INPROCESS_UDF for an embedded scalar Arrow UDF Review Comment: Removed the bullet. Eval type 258 is handled by the dedicated `InProcessArrowEvalPythonExec`. -- 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]
