dongjoon-hyun commented on code in PR #58978:
URL: https://github.com/apache/spark/pull/58978#discussion_r4076762056
##########
dev/spark-test-image/python-312/Dockerfile:
##########
@@ -44,6 +44,7 @@ RUN apt-get update && apt-get install -y \
libssl-dev \
openjdk-17-jdk-headless \
python3.12 \
+ python3.12-dev \
Review Comment:
Oh, `jep` requires `python3.12-dev` always?
##########
dev/spark-test-image/python-312/Dockerfile:
##########
@@ -66,3 +67,10 @@ RUN python3.12 -m pip install -U pip
RUN python3.12 -m pip install --group ml_torch --index-url
https://download.pytorch.org/whl/cpu && \
python3.12 -m pip install --group ci_classic_standard --group
ci_connect_standard && \
python3.12 -m pip cache purge
+
+# Exercise in-process UDFs in the existing pyspark-sql job. JEP needs Python
+# headers and a JDK when building its JNI library; no separate test job is
needed.
+RUN JAVA_HOME="$(dirname "$(dirname "$(readlink -f "$(command -v javac)")")")"
\
+ python3.12 -m pip install --no-cache-dir "jep==4.3.1" cffi
Review Comment:
Please start to use the latest one, 4.3.2.
- https://github.com/ninia/jep/releases/tag/v4.3.2 (2026-08-15)
- https://github.com/ninia/jep/releases/tag/v4.3.1 (2025-11-03)
##########
python/pyspark/inprocess/runtime.py:
##########
@@ -0,0 +1,96 @@
+#
+# 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.
+#
+
+
+"""
+In-process Python UDF runtime entry point.
+
+``_inprocess_invoke`` is imported into the jep SharedInterpreter's global
namespace
+during executor initialization (see ``InProcessPythonRuntime.initialize()``),
then called
+directly from the JVM via ``interp.invoke("_inprocess_invoke", ...)``.
+
+Both input and output use the Arrow C Data Interface (CDI). The JVM
pre-allocates
+ArrowArray/ArrowSchema C structs for every input column and for the output,
passing
+their native addresses as Python ints. Input arrays are reconstructed via
+``pa.Array._import_from_c`` (zero-copy). The output is written via
``arr._export_to_c``
+into the JVM-owned structs (zero-copy).
+
+jep type conversions (Java -> Python):
+ byte[] -> bytes (or sequence of signed ints; masked to
unsigned below)
+ List<Long> (boxed) -> list of Python ints
+ Long -> int
+"""
+
+import traceback as _traceback
+from functools import lru_cache
+
+import pyarrow as pa
+
+from pyspark import cloudpickle
+from pyspark.sql.pandas.types import to_arrow_type
+from pyspark.sql.types import _parse_datatype_json_string
+
+_UDF_TRACEBACK_SENTINEL = "__INPROCESS_UDF_TRACEBACK__:"
+
+
+@lru_cache(maxsize=128)
+def _load_udf(serialized_udf: bytes, return_type_json: str, timezone: str):
+ return (
+ cloudpickle.loads(serialized_udf),
+ to_arrow_type(_parse_datatype_json_string(return_type_json),
timezone=timezone),
+ )
+
+
+def _validate_result(result, expected_rows: int, expected_type: pa.DataType)
-> None:
+ if not isinstance(result, pa.Array):
+ raise TypeError(f"In-process UDF must return a pyarrow.Array, got
{type(result).__name__}")
+ if len(result) != expected_rows:
+ raise ValueError(f"In-process UDF returned {len(result)} rows;
expected {expected_rows}")
+ if result.type != expected_type:
+ raise TypeError(f"In-process UDF returned {result.type}; expected
{expected_type}")
+ result.validate()
+
+
+def _inprocess_invoke(
+ serialized_udf,
+ input_array_ptrs,
+ input_schema_ptrs,
+ output_array_ptr: int,
+ output_schema_ptr: int,
+ expected_rows: int,
+ return_type_json: str,
+ timezone: str,
+) -> None:
+ """Consume input CDI structs and export a validated, row-preserving result.
+
+ The caller owns the struct memory and releases any unconsumed exports on
failure.
+ Imported input arrays and exported output buffers follow Arrow's release
callbacks.
+ """
+ udf_key = bytes(b & 0xFF for b in serialized_udf)
+ udf_func, expected_type = _load_udf(udf_key, return_type_json, timezone)
+ if len(input_array_ptrs) != len(input_schema_ptrs):
+ raise ValueError("Mismatched input ArrowArray and ArrowSchema pointer
counts")
+ input_arrays = [
+ pa.Array._import_from_c(int(ap), int(sp))
+ for ap, sp in zip(input_array_ptrs, input_schema_ptrs)
+ ]
+ try:
+ result = udf_func(*input_arrays)
+ _validate_result(result, int(expected_rows), expected_type)
+ result._export_to_c(int(output_array_ptr), int(output_schema_ptr))
Review Comment:
A result with a non-zero offset passes `_validate_result` and is exported
with `ArrowArray.offset > 0`. Arrow Java 19.0.0's `ArrayImporter` ignores
`offset`: it builds `ArrowFieldNode(length, null_count)` and reads buffers from
position 0. Any sliced result, e.g. `pa.array([9, 1, 2, 3],
pa.int64()).slice(1)` for a 3-row batch, is then read as `[9, 1, 2]` without
any error. Validity bits and string offsets shift the same way. Could we reject
or normalize (copy) results with `offset != 0` before `_export_to_c`?
##########
sql/core/src/main/scala/org/apache/spark/sql/execution/python/InProcessPythonUDF.scala:
##########
@@ -0,0 +1,84 @@
+/*
+ * 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 org.apache.spark.sql.catalyst.expressions.{
+ Attribute, AttributeReference, AttributeSet, Expression, ExprId,
NamedExpression, Unevaluable
+}
+import org.apache.spark.sql.catalyst.plans.logical.{LogicalPlan, UnaryNode}
+import org.apache.spark.sql.types.DataType
+
+/**
+ * A cloudpickle-serialized Python function to be executed in-process via jep
+ * (Java Embedded Python).
+ *
+ * Distinct from [[org.apache.spark.sql.catalyst.expressions.PythonUDF]] which
uses an
+ * out-of-process Python worker connected via socket.
+ *
+ * Evaluated by [[InProcessArrowEvalExec]], which passes Arrow column buffers
to CPython
+ * as PyArrow arrays via native memory addresses (zero-copy input), then
imports the
+ * PyArrow result buffers through CDI without copying.
+ *
+ * @param name display name for plan explain output
+ * @param serializedFunc cloudpickle-serialized Python function bytes
+ * @param children input column expressions
+ * @param dataType declared return type (validated against the Arrow
result)
+ * @param udfDeterministic whether the UDF is deterministic
+ * @param resultId unique identifier for this UDF result
+ */
+case class InProcessPythonUDF(
+ name: String,
+ serializedFunc: Array[Byte],
+ children: Seq[Expression],
+ dataType: DataType,
+ udfDeterministic: Boolean = true,
+ resultId: ExprId = NamedExpression.newExprId)
+ extends Expression with Unevaluable {
Review Comment:
Because this is neither a `PythonUDF` nor tagged with the `PYTHON_UDF` tree
pattern, the Python UDF guards in rules that run before extraction don't apply:
- `PruneFileSourcePartitions`, `PruneHiveTablePartitions`, `FileScanBuilder`
and `PushDownUtils` check `!f.exists(_.isInstanceOf[PythonUDF])`. So
`spark.table("partitioned").where(f(F.col("dt")) == 1)` is treated as a
partition filter and evaluated on the driver, which fails with "Cannot evaluate
expression".
- `InjectRuntimeFilter` checks `containsPattern(PYTHON_UDF)`, so a creation
side such as `Filter(f(x) > 0, dim)` can be copied into a bloom-filter
subquery. That subquery is never re-optimized, and `transformUp` does not enter
subqueries, so the UDF is never extracted there and fails at runtime.
Could this extend `PythonFuncExpression` (or at least set `nodePatterns =
Seq(PYTHON_UDF)` and update the guards)?
##########
python/pyspark/inprocess/runtime.py:
##########
@@ -0,0 +1,96 @@
+#
+# 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.
+#
+
+
+"""
+In-process Python UDF runtime entry point.
+
+``_inprocess_invoke`` is imported into the jep SharedInterpreter's global
namespace
+during executor initialization (see ``InProcessPythonRuntime.initialize()``),
then called
+directly from the JVM via ``interp.invoke("_inprocess_invoke", ...)``.
+
+Both input and output use the Arrow C Data Interface (CDI). The JVM
pre-allocates
+ArrowArray/ArrowSchema C structs for every input column and for the output,
passing
+their native addresses as Python ints. Input arrays are reconstructed via
+``pa.Array._import_from_c`` (zero-copy). The output is written via
``arr._export_to_c``
+into the JVM-owned structs (zero-copy).
+
+jep type conversions (Java -> Python):
+ byte[] -> bytes (or sequence of signed ints; masked to
unsigned below)
+ List<Long> (boxed) -> list of Python ints
+ Long -> int
+"""
+
+import traceback as _traceback
+from functools import lru_cache
+
+import pyarrow as pa
+
+from pyspark import cloudpickle
+from pyspark.sql.pandas.types import to_arrow_type
+from pyspark.sql.types import _parse_datatype_json_string
+
+_UDF_TRACEBACK_SENTINEL = "__INPROCESS_UDF_TRACEBACK__:"
+
+
+@lru_cache(maxsize=128)
+def _load_udf(serialized_udf: bytes, return_type_json: str, timezone: str):
+ return (
+ cloudpickle.loads(serialized_udf),
+ to_arrow_type(_parse_datatype_json_string(return_type_json),
timezone=timezone),
+ )
+
+
+def _validate_result(result, expected_rows: int, expected_type: pa.DataType)
-> None:
+ if not isinstance(result, pa.Array):
+ raise TypeError(f"In-process UDF must return a pyarrow.Array, got
{type(result).__name__}")
+ if len(result) != expected_rows:
+ raise ValueError(f"In-process UDF returned {len(result)} rows;
expected {expected_rows}")
+ if result.type != expected_type:
+ raise TypeError(f"In-process UDF returned {result.type}; expected
{expected_type}")
+ result.validate()
+
+
+def _inprocess_invoke(
+ serialized_udf,
+ input_array_ptrs,
+ input_schema_ptrs,
+ output_array_ptr: int,
+ output_schema_ptr: int,
+ expected_rows: int,
+ return_type_json: str,
+ timezone: str,
+) -> None:
+ """Consume input CDI structs and export a validated, row-preserving result.
+
+ The caller owns the struct memory and releases any unconsumed exports on
failure.
+ Imported input arrays and exported output buffers follow Arrow's release
callbacks.
+ """
+ udf_key = bytes(b & 0xFF for b in serialized_udf)
+ udf_func, expected_type = _load_udf(udf_key, return_type_json, timezone)
+ if len(input_array_ptrs) != len(input_schema_ptrs):
+ raise ValueError("Mismatched input ArrowArray and ArrowSchema pointer
counts")
+ input_arrays = [
+ pa.Array._import_from_c(int(ap), int(sp))
+ for ap, sp in zip(input_array_ptrs, input_schema_ptrs)
+ ]
+ try:
+ result = udf_func(*input_arrays)
+ _validate_result(result, int(expected_rows), expected_type)
+ result._export_to_c(int(output_array_ptr), int(output_schema_ptr))
+ except Exception:
Review Comment:
`except Exception` doesn't catch `SystemExit`. In JEP 4.3.1,
`process_py_exception` (`jep_exceptions.c`) calls C `exit()` for `SystemExit`,
so a UDF, or a library it calls, that runs `sys.exit()` terminates the whole
executor JVM (the driver JVM in local mode), possibly with exit code 0.
`_load_udf` (`cloudpickle.loads`) is also outside the `try`. Should this catch
`BaseException` and convert it?
##########
sql/core/src/main/scala/org/apache/spark/sql/execution/python/InProcessPythonUDF.scala:
##########
@@ -0,0 +1,84 @@
+/*
+ * 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 org.apache.spark.sql.catalyst.expressions.{
+ Attribute, AttributeReference, AttributeSet, Expression, ExprId,
NamedExpression, Unevaluable
+}
+import org.apache.spark.sql.catalyst.plans.logical.{LogicalPlan, UnaryNode}
+import org.apache.spark.sql.types.DataType
+
+/**
+ * A cloudpickle-serialized Python function to be executed in-process via jep
+ * (Java Embedded Python).
+ *
+ * Distinct from [[org.apache.spark.sql.catalyst.expressions.PythonUDF]] which
uses an
+ * out-of-process Python worker connected via socket.
+ *
+ * Evaluated by [[InProcessArrowEvalExec]], which passes Arrow column buffers
to CPython
+ * as PyArrow arrays via native memory addresses (zero-copy input), then
imports the
+ * PyArrow result buffers through CDI without copying.
+ *
+ * @param name display name for plan explain output
+ * @param serializedFunc cloudpickle-serialized Python function bytes
+ * @param children input column expressions
+ * @param dataType declared return type (validated against the Arrow
result)
+ * @param udfDeterministic whether the UDF is deterministic
+ * @param resultId unique identifier for this UDF result
+ */
+case class InProcessPythonUDF(
+ name: String,
+ serializedFunc: Array[Byte],
+ children: Seq[Expression],
+ dataType: DataType,
+ udfDeterministic: Boolean = true,
+ resultId: ExprId = NamedExpression.newExprId)
+ extends Expression with Unevaluable {
+
+ override def nullable: Boolean = true
+ override def prettyName: String = name
+
+ override lazy val deterministic: Boolean =
+ udfDeterministic && children.forall(_.deterministic)
+
+ lazy val resultAttribute: Attribute =
+ AttributeReference(name, dataType, nullable)(exprId = resultId)
+
+ override def toString: String = s"$name(${children.mkString(",
")})#${resultId.id}"
+
+ override protected def withNewChildrenInternal(
+ newChildren: IndexedSeq[Expression]): InProcessPythonUDF =
+ copy(children = newChildren)
+}
+
+/**
+ * Logical plan node that evaluates [[InProcessPythonUDF]]s in-process via jep.
+ * Inserted by [[ExtractInProcessPythonUDFs]] during query optimization,
before physical planning.
+ * Planned as [[InProcessArrowEvalExec]] by
[[org.apache.spark.sql.execution.SparkStrategies]].
+ */
+case class InProcessEvalPython(
Review Comment:
`InProcessEvalPython` is not in `PushPredicateThroughNonJoin.canPushThrough`
or `LimitPushDown`, and `ScanOperation` can't see through it. Deterministic
non-UDF conjuncts therefore stay above the UDF node:
- `spark.read.parquet(p).filter((F.col("dt") == "x") & (f(F.col("m")) > 0))`
loses partition pruning and data filters.
- `df.filter(df.d != 0).filter(div(df.n, df.d) > 1)` passes `d == 0` rows to
the UDF, so `pc.divide` raises `ArrowInvalid`. `arrow_udf` pushes the guard
below `ArrowEvalPython` and succeeds.
##########
python/pyspark/inprocess/runtime.py:
##########
@@ -0,0 +1,96 @@
+#
+# 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.
+#
+
+
+"""
+In-process Python UDF runtime entry point.
+
+``_inprocess_invoke`` is imported into the jep SharedInterpreter's global
namespace
+during executor initialization (see ``InProcessPythonRuntime.initialize()``),
then called
+directly from the JVM via ``interp.invoke("_inprocess_invoke", ...)``.
+
+Both input and output use the Arrow C Data Interface (CDI). The JVM
pre-allocates
+ArrowArray/ArrowSchema C structs for every input column and for the output,
passing
+their native addresses as Python ints. Input arrays are reconstructed via
+``pa.Array._import_from_c`` (zero-copy). The output is written via
``arr._export_to_c``
+into the JVM-owned structs (zero-copy).
+
+jep type conversions (Java -> Python):
+ byte[] -> bytes (or sequence of signed ints; masked to
unsigned below)
+ List<Long> (boxed) -> list of Python ints
+ Long -> int
+"""
+
+import traceback as _traceback
+from functools import lru_cache
+
+import pyarrow as pa
+
+from pyspark import cloudpickle
+from pyspark.sql.pandas.types import to_arrow_type
+from pyspark.sql.types import _parse_datatype_json_string
+
+_UDF_TRACEBACK_SENTINEL = "__INPROCESS_UDF_TRACEBACK__:"
+
+
+@lru_cache(maxsize=128)
+def _load_udf(serialized_udf: bytes, return_type_json: str, timezone: str):
+ return (
+ cloudpickle.loads(serialized_udf),
+ to_arrow_type(_parse_datatype_json_string(return_type_json),
timezone=timezone),
+ )
+
+
+def _validate_result(result, expected_rows: int, expected_type: pa.DataType)
-> None:
+ if not isinstance(result, pa.Array):
+ raise TypeError(f"In-process UDF must return a pyarrow.Array, got
{type(result).__name__}")
+ if len(result) != expected_rows:
+ raise ValueError(f"In-process UDF returned {len(result)} rows;
expected {expected_rows}")
+ if result.type != expected_type:
Review Comment:
`!=` also compares nested field nullability, and `to_arrow_type` sets it
from `containsNull` / `nullable`. An identity UDF declared as
`ArrayType(StringType())` over `F.split(...)` input receives `list<element:
string not null>` from the JVM, so it fails with `TypeError ... expected
list<element: string>`. Declaring `containsNull=False` instead rejects ordinary
nullable PyArrow results. Could the nullability check be relaxed, e.g. by
casting when only nullability differs?
##########
sql/core/src/main/scala/org/apache/spark/sql/execution/python/InProcessPythonUDF.scala:
##########
@@ -0,0 +1,84 @@
+/*
+ * 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 org.apache.spark.sql.catalyst.expressions.{
+ Attribute, AttributeReference, AttributeSet, Expression, ExprId,
NamedExpression, Unevaluable
+}
+import org.apache.spark.sql.catalyst.plans.logical.{LogicalPlan, UnaryNode}
+import org.apache.spark.sql.types.DataType
+
+/**
+ * A cloudpickle-serialized Python function to be executed in-process via jep
+ * (Java Embedded Python).
+ *
+ * Distinct from [[org.apache.spark.sql.catalyst.expressions.PythonUDF]] which
uses an
+ * out-of-process Python worker connected via socket.
+ *
+ * Evaluated by [[InProcessArrowEvalExec]], which passes Arrow column buffers
to CPython
+ * as PyArrow arrays via native memory addresses (zero-copy input), then
imports the
+ * PyArrow result buffers through CDI without copying.
+ *
+ * @param name display name for plan explain output
+ * @param serializedFunc cloudpickle-serialized Python function bytes
+ * @param children input column expressions
+ * @param dataType declared return type (validated against the Arrow
result)
+ * @param udfDeterministic whether the UDF is deterministic
+ * @param resultId unique identifier for this UDF result
+ */
+case class InProcessPythonUDF(
+ name: String,
+ serializedFunc: Array[Byte],
+ children: Seq[Expression],
+ dataType: DataType,
+ udfDeterministic: Boolean = true,
+ resultId: ExprId = NamedExpression.newExprId)
+ extends Expression with Unevaluable {
+
+ override def nullable: Boolean = true
+ override def prettyName: String = name
+
+ override lazy val deterministic: Boolean =
Review Comment:
With `deterministic=False`, this expression is neither `Nondeterministic`
nor `UserDefinedExpression`, so `PullOutNondeterministic` ignores it:
- `df.groupBy(nd(df.x)).count()` fails with INTERNAL_ERROR
"Non-deterministic expression ... should not appear in grouping expression".
- `df.orderBy(nd(df.x))` fails with `INVALID_NON_DETERMINISTIC_EXPRESSIONS`.
`udf(...).asNondeterministic()` works in both cases.
##########
sql/core/src/main/scala/org/apache/spark/sql/execution/python/InProcessArrowEvalExec.scala:
##########
@@ -0,0 +1,198 @@
+/*
+ * 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 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, Expression,
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. Rows
(including
+ * computed UDF arguments) are written to Arrow once, then input and output
buffers cross
+ * the JVM/Python boundary without IPC serialization. The runtime owns the JEP
thread.
+ */
+case class InProcessArrowEvalExec(
+ udfs: Seq[InProcessPythonUDF],
+ resultAttrs: Seq[Attribute],
+ child: SparkPlan) extends UnaryExecNode {
+
+ override def output: Seq[Attribute] = child.output ++ resultAttrs
+
+ override protected def doExecute(): RDD[InternalRow] = {
+ val expressions = ArrayBuffer[Expression](child.output: _*)
+ val inputOrdinals = udfs.map { udf =>
+ udf.children.map {
+ case attr: Attribute if child.output.exists(_.exprId == attr.exprId) =>
+ child.output.indexWhere(_.exprId == attr.exprId)
+ case expr =>
+ expressions += expr
+ expressions.size - 1
+ }
+ }
+ // Synthetic names also allow joins with duplicate output column names.
+ val inputSchema = StructType(expressions.zipWithIndex.map { case (expr, i)
=>
+ StructField(s"_input$i", expr.dataType, expr.nullable)
+ }.toSeq)
+ val inputExpressions = expressions.toSeq
+ val childOutput = child.output
+ val resultOutput = output
+ val batchSize = conf.arrowMaxRecordsPerBatch
+ val maxBytes = conf.arrowMaxBytesPerBatch
+ val timeZoneId = conf.sessionLocalTimeZone
+
+ child.execute().mapPartitions { rows =>
+ val context = Option(TaskContext.get())
+ def checkCancellation(): Unit =
context.foreach(_.killTaskIfInterrupted())
+
+ val resultProjection = UnsafeProjection.create(resultOutput,
resultOutput)
+ val projectInput: InternalRow => InternalRow =
+ if (inputExpressions.size == childOutput.size) {
+ identity[InternalRow]
+ } else {
+ val projection = UnsafeProjection.create(inputExpressions,
childOutput)
+ projection.initialize(TaskContext.getPartitionId())
+ projection
+ }
+ val root = VectorSchemaRoot.create(
+ ArrowUtils.toArrowSchema(inputSchema, timeZoneId, false),
ArrowUtils.rootAllocator)
Review Comment:
This call binds to the `toArrowSchema(schema, timeZoneId, largeVarTypes)`
overload, which neither dedups nested field names nor calls
`failDuplicatedFieldNames` (the existing path does, in `PythonArrowInput`).
Arrow Java's default `CONFLICT_REPLACE` then collapses same-named struct
children. Since all child columns go through this schema, a pass-through column
such as `F.struct(df.id, df.id.cast("string").alias("id"))` is enough: the
bigint slot is read as UTF8 (garbage), and the read-back hits
`ArrayIndexOutOfBoundsException`, even though the UDF never reads that column.
##########
sql/core/src/main/scala/org/apache/spark/sql/execution/python/InProcessArrowEvalExec.scala:
##########
@@ -0,0 +1,198 @@
+/*
+ * 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 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, Expression,
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. Rows
(including
+ * computed UDF arguments) are written to Arrow once, then input and output
buffers cross
+ * the JVM/Python boundary without IPC serialization. The runtime owns the JEP
thread.
+ */
+case class InProcessArrowEvalExec(
+ udfs: Seq[InProcessPythonUDF],
+ resultAttrs: Seq[Attribute],
+ child: SparkPlan) extends UnaryExecNode {
+
+ override def output: Seq[Attribute] = child.output ++ resultAttrs
+
+ override protected def doExecute(): RDD[InternalRow] = {
+ val expressions = ArrayBuffer[Expression](child.output: _*)
Review Comment:
Writing every child column, not only the UDF inputs, through Arrow means
columns the UDF never reads can fail the query. A pass-through
`CalendarInterval` with |micros| greater than about 9.2e15 throws
`calendarIntervalArrowNanosOverflowError`, and out-of-range nanos timestamps
throw `timestampNanosEpochNanosOverflowError`. The worker path keeps
pass-through rows in `HybridRowQueue` and succeeds on the same query. It also
avoids the extra copies for wide schemas. Could we export only the UDF
arguments and join the results back to the original rows?
##########
sql/core/src/main/scala/org/apache/spark/sql/execution/python/ExtractInProcessPythonUDFs.scala:
##########
@@ -0,0 +1,101 @@
+/*
+ * 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 scala.collection.mutable
+
+import org.apache.spark.sql.catalyst.expressions.{Attribute,
AttributeReference, Expression, ExternalUserDefinedFunction, NamedExpression,
PythonUDF, WindowExpression}
+import org.apache.spark.sql.catalyst.expressions.aggregate.AggregateExpression
+import org.apache.spark.sql.catalyst.plans.logical.{LogicalPlan, Project}
+import org.apache.spark.sql.catalyst.rules.Rule
+
+/**
+ * Extracts [[InProcessPythonUDF]] expressions from logical plan nodes,
rewriting the plan
+ * so each batch of UDFs is evaluated in a dedicated [[InProcessEvalPython]]
node.
+ *
+ * Simpler than [[ExtractPythonUDFs]]: in-process UDFs are scalar and don't
support
+ * iterator mode, aggregate mode, or nested chaining in Phase 1.
+ *
+ * Example rewrite:
+ * Project [double(a), triple(b)]
+ * Scan
+ * becomes:
+ * Project [inprocessUDF0, inprocessUDF1]
+ * InProcessEvalPython [double(a), triple(b)] to [inprocessUDF0,
inprocessUDF1]
+ * Scan
+ */
+object ExtractInProcessPythonUDFs extends Rule[LogicalPlan] {
+
+ private def hasInProcessUDF(e: Expression): Boolean =
+ e.exists(_.isInstanceOf[InProcessPythonUDF])
+
+ override def apply(plan: LogicalPlan): LogicalPlan = plan.transformUp {
+ // Already extracted - skip to avoid double-wrapping
+ case p: InProcessEvalPython => p
+
+ case node: LogicalPlan if node.expressions.exists(hasInProcessUDF) =>
+ extract(node)
+ }
+
+ private def extract(plan: LogicalPlan): LogicalPlan = {
+ // Collect all distinct InProcessPythonUDFs from this plan's expressions
+ val udfs = plan.expressions
+ .flatMap(_.collect { case u: InProcessPythonUDF => u })
+ .distinct
+
+ if (udfs.isEmpty) return plan
+
+ udfs.foreach { udf =>
+ require(!udf.children.exists(_.exists {
+ case _: InProcessPythonUDF | _: PythonUDF | _:
ExternalUserDefinedFunction |
+ _: AggregateExpression | _: WindowExpression => true
+ case _ => false
+ }), "In-process Python UDFs do not support nested UDF, aggregate or
window arguments")
+ }
+
+ // Map each UDF to a fresh AttributeReference that will hold its result
+ val attributeMap = mutable.LinkedHashMap[InProcessPythonUDF,
NamedExpression]()
+
+ // For each child plan, find UDFs whose inputs are fully satisfied by that
child
+ val newChildren = plan.children.map { child =>
+ val validUdfs = udfs.filter { udf =>
Review Comment:
`extract()` pushes any UDF whose references fit a child below that child,
including an `Aggregate`'s result expressions over grouping keys or constants.
For example, `df.groupBy("k").agg(f(F.col("k")))`,
`df.groupBy("k").count().select("k", f(F.col("k")))` (after `CollapseProject`)
and `df.agg(F.count("*"), f(F.lit(1)))` become `Aggregate([k], [k,
inprocessUDF0 AS ...], InProcessEvalPython(...))`. `inprocessUDF0` is neither a
grouping expression nor an aggregate, so result binding fails with `Couldn't
find inprocessUDF0` (plan-integrity failure in tests). For `PythonUDF`,
`ExtractPythonUDFFromAggregate` lifts these above the `Aggregate`.
##########
sql/core/src/main/scala/org/apache/spark/sql/execution/python/ExtractInProcessPythonUDFs.scala:
##########
@@ -0,0 +1,101 @@
+/*
+ * 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 scala.collection.mutable
+
+import org.apache.spark.sql.catalyst.expressions.{Attribute,
AttributeReference, Expression, ExternalUserDefinedFunction, NamedExpression,
PythonUDF, WindowExpression}
+import org.apache.spark.sql.catalyst.expressions.aggregate.AggregateExpression
+import org.apache.spark.sql.catalyst.plans.logical.{LogicalPlan, Project}
+import org.apache.spark.sql.catalyst.rules.Rule
+
+/**
+ * Extracts [[InProcessPythonUDF]] expressions from logical plan nodes,
rewriting the plan
+ * so each batch of UDFs is evaluated in a dedicated [[InProcessEvalPython]]
node.
+ *
+ * Simpler than [[ExtractPythonUDFs]]: in-process UDFs are scalar and don't
support
+ * iterator mode, aggregate mode, or nested chaining in Phase 1.
+ *
+ * Example rewrite:
+ * Project [double(a), triple(b)]
+ * Scan
+ * becomes:
+ * Project [inprocessUDF0, inprocessUDF1]
+ * InProcessEvalPython [double(a), triple(b)] to [inprocessUDF0,
inprocessUDF1]
+ * Scan
+ */
+object ExtractInProcessPythonUDFs extends Rule[LogicalPlan] {
+
+ private def hasInProcessUDF(e: Expression): Boolean =
+ e.exists(_.isInstanceOf[InProcessPythonUDF])
+
+ override def apply(plan: LogicalPlan): LogicalPlan = plan.transformUp {
+ // Already extracted - skip to avoid double-wrapping
+ case p: InProcessEvalPython => p
+
+ case node: LogicalPlan if node.expressions.exists(hasInProcessUDF) =>
+ extract(node)
+ }
+
+ private def extract(plan: LogicalPlan): LogicalPlan = {
+ // Collect all distinct InProcessPythonUDFs from this plan's expressions
+ val udfs = plan.expressions
+ .flatMap(_.collect { case u: InProcessPythonUDF => u })
+ .distinct
+
+ if (udfs.isEmpty) return plan
+
+ udfs.foreach { udf =>
+ require(!udf.children.exists(_.exists {
Review Comment:
This batch runs after `CollapseProject`, which inlines single-use aliases.
Code that evaluates the UDF in a separate `select`, as the docs suggest, can
still end up as a nested argument and be rejected here:
- `df.groupBy("k").agg(F.sum("v").alias("s")).select("k", f(F.col("s")))`
becomes `Aggregate([k], [k, f(sum(v))])`
- `df.select(g(df.a).alias("b")).select(f(F.col("b")))` becomes `f(g(a))`
Both fail with `IllegalArgumentException: requirement failed: ...`. Regular
Python UDFs handle both shapes: `ExtractPythonUDFFromAggregate` for the first
and the recursive `extract` in `ExtractPythonUDFs` for the second.
##########
sql/core/src/main/scala/org/apache/spark/sql/execution/python/InProcessPythonChecks.scala:
##########
@@ -0,0 +1,63 @@
+/*
+ * 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 org.apache.spark.sql.catalyst.plans.logical.LogicalPlan
+import org.apache.spark.sql.catalyst.rules.Rule
+import org.apache.spark.sql.internal.SQLConf
+
+/**
+ * Validates that in-process Python UDFs are only used when exactly one task
can run per
+ * executor, preventing GIL contention on [[InProcessPythonRuntime]]'s shared
interpreter.
+ *
+ * The constraint: spark.executor.cores / spark.task.cpus == 1
+ *
+ * Typical correct configuration:
+ * spark.executor.cores=1 (one core per executor, parallelism via more
executors)
+ *
+ * Runs after [[ExtractInProcessPythonUDFs]] in the "Extract InProcess Python
UDFs" optimizer
+ * batch, so it sees [[InProcessEvalPython]] nodes.
+ */
+object InProcessPythonChecks extends Rule[LogicalPlan] {
+
+ override def apply(plan: LogicalPlan): LogicalPlan = {
+ plan.foreach {
+ case _: InProcessEvalPython => checkConcurrencyConfig()
+ case _ =>
+ }
+ plan
+ }
+
+ private def checkConcurrencyConfig(): Unit = {
+ val conf = SQLConf.get
+ val executorCores =
+ conf.getConfString("spark.executor.cores", "1").toInt
+ val taskCpus =
+ conf.getConfString("spark.task.cpus", "1").toInt
Review Comment:
`spark.task.cpus` is a `decimalConf`, and fractional values like `0.5` are
documented. `mergeSparkConf` copies the raw string into SQLConf, so
`"0.5".toInt` throws `NumberFormatException` in the optimizer for every query
that uses an in-process UDF. Separately, defaulting an unset
`spark.executor.cores` to `"1"` lets `local[*]` (which the docs recommend) and
standalone defaults pass while N tasks share the interpreter. Consider using
`CPUS_PER_TASK` / `EXECUTOR_CORES` with `ResourceProfile` helpers.
Alternatively, drop the check, since `interpreterLock` already serializes
invocations.
##########
sql/core/src/main/scala/org/apache/spark/sql/execution/python/InProcessArrowEvalExec.scala:
##########
@@ -0,0 +1,198 @@
+/*
+ * 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 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, Expression,
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. Rows
(including
+ * computed UDF arguments) are written to Arrow once, then input and output
buffers cross
+ * the JVM/Python boundary without IPC serialization. The runtime owns the JEP
thread.
+ */
+case class InProcessArrowEvalExec(
+ udfs: Seq[InProcessPythonUDF],
+ resultAttrs: Seq[Attribute],
+ child: SparkPlan) extends UnaryExecNode {
+
+ override def output: Seq[Attribute] = child.output ++ resultAttrs
+
+ override protected def doExecute(): RDD[InternalRow] = {
+ val expressions = ArrayBuffer[Expression](child.output: _*)
+ val inputOrdinals = udfs.map { udf =>
+ udf.children.map {
+ case attr: Attribute if child.output.exists(_.exprId == attr.exprId) =>
+ child.output.indexWhere(_.exprId == attr.exprId)
+ case expr =>
+ expressions += expr
+ expressions.size - 1
+ }
+ }
+ // Synthetic names also allow joins with duplicate output column names.
+ val inputSchema = StructType(expressions.zipWithIndex.map { case (expr, i)
=>
+ StructField(s"_input$i", expr.dataType, expr.nullable)
+ }.toSeq)
+ val inputExpressions = expressions.toSeq
+ val childOutput = child.output
+ val resultOutput = output
+ val batchSize = conf.arrowMaxRecordsPerBatch
+ val maxBytes = conf.arrowMaxBytesPerBatch
+ val timeZoneId = conf.sessionLocalTimeZone
+
+ child.execute().mapPartitions { rows =>
+ val context = Option(TaskContext.get())
+ def checkCancellation(): Unit =
context.foreach(_.killTaskIfInterrupted())
+
+ val resultProjection = UnsafeProjection.create(resultOutput,
resultOutput)
+ val projectInput: InternalRow => InternalRow =
+ if (inputExpressions.size == childOutput.size) {
+ identity[InternalRow]
+ } else {
+ val projection = UnsafeProjection.create(inputExpressions,
childOutput)
+ projection.initialize(TaskContext.getPartitionId())
+ projection
+ }
+ val root = VectorSchemaRoot.create(
+ ArrowUtils.toArrowSchema(inputSchema, timeZoneId, false),
ArrowUtils.rootAllocator)
+ val writer = try {
+ ArrowWriter.create(root)
+ } catch {
+ case t: Throwable => Utils.tryWithSafeFinally { throw t } {
root.close() }
+ }
+ val results = ArrayBuffer.empty[ArrowColumnVector]
+ var closed = false
+
+ def closeResults(): Unit = {
+ val previous = results.toArray
+ results.clear()
+ AutoCloseables.close(previous: _*)
+ }
+
+ def close(): Unit = {
+ if (!closed) {
+ closed = true
+ Utils.tryWithSafeFinally { closeResults() } { writer.root.close() }
+ }
+ }
+
+ context.foreach(_.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) {
+ closeResults()
+ writer.reset()
Review Comment:
The inputs are exported zero-copy, and `writer.reset()` zero-fills and
reuses the same `ArrowBuf`s for the next batch, whether or not Python still
references them. If a stateful UDF keeps an input (for example `state["last"] =
x`, or appends batches to a list), the kept array's values and nulls are
silently replaced by the next batch's data. That breaks Arrow immutability, and
the worker path doesn't have this problem.
##########
python/pyspark/inprocess/udf.py:
##########
@@ -0,0 +1,175 @@
+#
+# 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()
+"""
+
+from typing import Callable
+
+import pyarrow as pa
+
+from pyspark import 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 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 ``InProcessPythonUDF``
+ 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 = cloudpickle.dumps(_wrapped)
Review Comment:
Pickling with `cloudpickle.dumps` directly and loading without the worker
bootstrap skips several things `_prepare_for_python_RDD` and worker.py normally
handle:
- A UDF that uses `sc.broadcast(...)` fails with
`BROADCAST_VARIABLE_NOT_LOADED`, reported as "infrastructure error".
- `acc.add()` inside the UDF is silently lost; it never reaches the driver.
- Modules added with `sc.addPyFile` are not importable.
- There is no `PYTHON_VERSION_MISMATCH` check between the driver's Python
and the embedded one.
At a minimum, these should be rejected or documented.
##########
sql/core/src/main/scala/org/apache/spark/sql/execution/python/InProcessPythonUDF.scala:
##########
@@ -0,0 +1,84 @@
+/*
+ * 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 org.apache.spark.sql.catalyst.expressions.{
+ Attribute, AttributeReference, AttributeSet, Expression, ExprId,
NamedExpression, Unevaluable
+}
+import org.apache.spark.sql.catalyst.plans.logical.{LogicalPlan, UnaryNode}
+import org.apache.spark.sql.types.DataType
+
+/**
+ * A cloudpickle-serialized Python function to be executed in-process via jep
+ * (Java Embedded Python).
+ *
+ * Distinct from [[org.apache.spark.sql.catalyst.expressions.PythonUDF]] which
uses an
+ * out-of-process Python worker connected via socket.
+ *
+ * Evaluated by [[InProcessArrowEvalExec]], which passes Arrow column buffers
to CPython
+ * as PyArrow arrays via native memory addresses (zero-copy input), then
imports the
+ * PyArrow result buffers through CDI without copying.
+ *
+ * @param name display name for plan explain output
+ * @param serializedFunc cloudpickle-serialized Python function bytes
+ * @param children input column expressions
+ * @param dataType declared return type (validated against the Arrow
result)
+ * @param udfDeterministic whether the UDF is deterministic
+ * @param resultId unique identifier for this UDF result
+ */
+case class InProcessPythonUDF(
+ name: String,
+ serializedFunc: Array[Byte],
Review Comment:
`Array[Byte]` compares by reference, and `resultId` isn't normalized in
`canonicalized` (`PythonUDF` resets it to `ExprId(-1)`). Each Python `__call__`
sends a new `byte[]` and a new `resultId`, so two calls to the same UDF are
never `semanticEquals`. As a result, `df.groupBy(f(df.k)).agg(f(df.k))` fails
analysis, and rebuilt DataFrames miss the cache. Also, `expensive` isn't
overridden (`PythonFuncExpression` returns `true`), so
`df.select(f(df.a).alias("x")).filter("x > 0")` inlines the alias into the
pushed filter and evaluates `f` twice.
##########
python/pyspark/inprocess/runtime.py:
##########
@@ -0,0 +1,96 @@
+#
+# 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.
+#
+
+
+"""
+In-process Python UDF runtime entry point.
+
+``_inprocess_invoke`` is imported into the jep SharedInterpreter's global
namespace
+during executor initialization (see ``InProcessPythonRuntime.initialize()``),
then called
+directly from the JVM via ``interp.invoke("_inprocess_invoke", ...)``.
+
+Both input and output use the Arrow C Data Interface (CDI). The JVM
pre-allocates
+ArrowArray/ArrowSchema C structs for every input column and for the output,
passing
+their native addresses as Python ints. Input arrays are reconstructed via
+``pa.Array._import_from_c`` (zero-copy). The output is written via
``arr._export_to_c``
+into the JVM-owned structs (zero-copy).
+
+jep type conversions (Java -> Python):
+ byte[] -> bytes (or sequence of signed ints; masked to
unsigned below)
+ List<Long> (boxed) -> list of Python ints
+ Long -> int
+"""
+
+import traceback as _traceback
+from functools import lru_cache
+
+import pyarrow as pa
+
+from pyspark import cloudpickle
+from pyspark.sql.pandas.types import to_arrow_type
+from pyspark.sql.types import _parse_datatype_json_string
+
+_UDF_TRACEBACK_SENTINEL = "__INPROCESS_UDF_TRACEBACK__:"
+
+
+@lru_cache(maxsize=128)
+def _load_udf(serialized_udf: bytes, return_type_json: str, timezone: str):
+ return (
+ cloudpickle.loads(serialized_udf),
+ to_arrow_type(_parse_datatype_json_string(return_type_json),
timezone=timezone),
+ )
+
+
+def _validate_result(result, expected_rows: int, expected_type: pa.DataType)
-> None:
+ if not isinstance(result, pa.Array):
+ raise TypeError(f"In-process UDF must return a pyarrow.Array, got
{type(result).__name__}")
+ if len(result) != expected_rows:
+ raise ValueError(f"In-process UDF returned {len(result)} rows;
expected {expected_rows}")
+ if result.type != expected_type:
+ raise TypeError(f"In-process UDF returned {result.type}; expected
{expected_type}")
+ result.validate()
+
+
+def _inprocess_invoke(
+ serialized_udf,
+ input_array_ptrs,
+ input_schema_ptrs,
+ output_array_ptr: int,
+ output_schema_ptr: int,
+ expected_rows: int,
+ return_type_json: str,
+ timezone: str,
+) -> None:
+ """Consume input CDI structs and export a validated, row-preserving result.
+
+ The caller owns the struct memory and releases any unconsumed exports on
failure.
+ Imported input arrays and exported output buffers follow Arrow's release
callbacks.
+ """
+ udf_key = bytes(b & 0xFF for b in serialized_udf)
Review Comment:
JEP 4.3.1 passes the Java `byte[]` as a `PyJArray`, which has no buffer
protocol. This line therefore runs an O(pickle size) Python loop, plus a JNI
copy and a fresh hash of the new key, on every batch for every UDF, before the
`lru_cache` lookup. A UDF that captures a model of a few MB pays roughly
0.3-0.75 s per 10k-row batch even on cache hits. Could the JVM compute a digest
once per task and send the bytes only on a cache miss? `udf.dataType.json`
could be computed once per task as well.
--
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]