dongjoon-hyun commented on code in PR #58978: URL: https://github.com/apache/spark/pull/58978#discussion_r4209346898
########## sql/core/src/main/scala/org/apache/spark/sql/execution/python/InProcessArrowEvalPythonEvaluatorFactory.scala: ########## @@ -0,0 +1,511 @@ +/* + * 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.io.File +import java.nio.file.Files +import java.util.UUID +import java.util.concurrent.TimeUnit +import java.util.concurrent.atomic.AtomicInteger +import java.util.concurrent.locks.ReentrantLock + +import scala.collection.mutable.ArrayBuffer +import scala.jdk.CollectionConverters._ + +import com.google.common.util.concurrent.Uninterruptibles +import org.apache.arrow.c.{ArrowArray, ArrowSchema, BaseStruct} +import org.apache.arrow.util.AutoCloseables +import org.apache.arrow.vector.VectorSchemaRoot + +import org.apache.spark.{SparkEnv, SparkException, TaskContext} +import org.apache.spark.api.python.ChainedPythonFunctions +import org.apache.spark.sql.catalyst.InternalRow +import org.apache.spark.sql.catalyst.expressions.{Attribute, Expression, JoinedRow, PythonUDF, UnsafeProjection, UnsafeRow} +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._ +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, unless all of them are UDF arguments that read back from Arrow unchanged. + * Each batch owns its Arrow buffers so Python can safely retain input arrays. + * + * The evaluator owns its queue, so that cleanup at task completion is coordinated with a + * consumer on another thread, such as a pipelined Python writer or a TRANSFORM feed thread. + */ +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, + fullValidation: Boolean, + metrics: Map[String, SQLMetric]) + extends EvalPythonEvaluatorFactory(childOutput, udfs, output) { + + private[python] def runtimeSession: InProcessPythonRuntime.InterpreterSession = + InProcessPythonRuntime.currentSession + + /** Unused: `evaluateJoined` always evaluates the UDFs. */ + override protected def evaluate( + funcs: Seq[(ChainedPythonFunctions, Long)], + argMetas: Array[Array[ArgumentMetadata]], + rows: Iterator[InternalRow], + inputSchema: StructType, + context: TaskContext): Iterator[InternalRow] = + throw SparkException.internalError("In-process UDFs are evaluated with their input rows") + + override protected def evaluateJoined( + funcs: Seq[(ChainedPythonFunctions, Long)], + argMetas: Array[Array[ArgumentMetadata]], + rows: Iterator[InternalRow], + inputs: Seq[Expression], + inputSchema: StructType, + context: TaskContext): Option[Iterator[InternalRow]] = { + import InProcessArrowEvalPythonEvaluatorFactory.{Buffered, ReadBack, readsBack} + val inputColumns = inputs.length == childOutput.length && inputs.zip(childOutput).forall { + case (a: Attribute, c) => a.exprId == c.exprId + case _ => false + } + // If all input columns are UDF arguments, they are written to Arrow regardless. Read them + // back from the exported input vectors instead of buffering every input row, if their + // values read back from Arrow exactly as written and as fast as an unsafe row copy. + val joinInput = if (inputColumns && inputSchema.forall(f => readsBack(f.dataType))) { + ReadBack + } else if (inputColumns) { + Buffered(None) + } else { + // Each projected row is written to Arrow before the next input row is pulled, so the + // arguments go into a reused buffer rather than being copied value by value. + val projection = UnsafeProjection.create(inputs, childOutput) + projection.initialize(context.partitionId()) + Buffered(Some(projection)) + } + Some(evaluateBatches(funcs, argMetas, rows, inputSchema, context, joinInput)) + } + + private[python] def evaluateBatches( + funcs: Seq[(ChainedPythonFunctions, Long)], + argMetas: Array[Array[ArgumentMetadata]], + rows: Iterator[InternalRow], + inputSchema: StructType, + context: TaskContext, + joinInput: InProcessArrowEvalPythonEvaluatorFactory.JoinInput): Iterator[InternalRow] = { + import InProcessArrowEvalPythonEvaluatorFactory.{Buffered, ReadBack} + ArrowUtils.failDuplicatedFieldNames(inputSchema) + val functions = funcs.map { case (chain, _) => + if (chain.funcs.size != 1) { + throw SparkException.internalError( + "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) + // Capture before consuming input: an old task must never join a later context's session. + val runtime = runtimeSession + // Rows are copied out of the queue and Arrow vectors before they are returned, so they + // remain valid after task completion releases those, on whichever thread consumes them. + val resultProj = UnsafeProjection.create(output, output) + // Spill files go into a directory of the queue's own, created on the first spill, so that + // task completion can delete them when it cannot close the queue. + @volatile var spillDir: File = null + val (queue, projection) = joinInput match { + case Buffered(projection) => + val localDir = new File(Utils.getLocalDir(SparkEnv.get.conf)) + val serializerManager = SparkEnv.get.serializerManager + // Only the consumer holding the iterator's lock adds and removes rows. + val queue = new HybridRowQueue(context.taskMemoryManager(), localDir, + childOutput.length, serializerManager, lockFree = true) { + override protected def createDiskQueue(): RowQueue = synchronized { + if (spillDir == null) { + spillDir = Files.createTempDirectory(localDir.toPath, "inprocess-udf-").toFile + } + DiskRowQueue(Files.createTempFile(spillDir.toPath, "buffer", "").toFile, + childOutput.length, serializerManager) + } + } + (queue, projection.orNull) + case ReadBack => (null, null) + } + val joined = new JoinedRow + val handles = functions.map(_ => UUID.randomUUID().toString) + var registered = false + var writer: ArrowWriter = null + val results = ArrayBuffer.empty[ArrowColumnVector] + var startedAt = 0L + + def closeBatch(): Unit = { + val resources = ArrayBuffer.empty[AutoCloseable] + resources ++= results + results.clear() + if (writer != null) { + resources += writer.root + writer = null + } + AutoCloseables.close(resources.asJava) + } + + val resources = new InProcessArrowEvalPythonEvaluatorFactory.IteratorResources( + // Closing the queue deletes its spill files, which leaves an empty directory. + releaseTaskMemory = () => if (queue != null) { + Utils.tryWithSafeFinally(queue.close()) { if (spillDir != null) spillDir.delete() } + }, + // Neither starts a process nor throws, also on an interrupted thread. + abandonTaskMemory = () => if (spillDir != null) Utils.deleteQuietly(spillDir), Review Comment: **[Low] Follow-up on R11-3: a spill after abandonment creates a new spill directory, which nothing deletes.** The last paragraph of my round-11 comment (https://github.com/apache/spark/pull/58978#discussion_r4189042205) noted that an abandoned queue stays a `TaskMemoryManager` consumer until the executor cleans up. With the lazy directory, that now leaks. If the queue has not spilled yet, `spillDir` is null here, so nothing is deleted and nothing stops a later spill. When another consumer of the task then allocates, `HybridQueue.spill` calls `createDiskQueue` (`HybridQueue.scala` L85-86), which creates a new `inprocess-udf-*` directory and writes the buffered rows into it (L154-160). Since the task memory is Abandoned, `releaseTaskMemory` never runs, so the directory and its file stay until the executor exits. For example, take stacked `f(g(x))` consumed by a TRANSFORM feed thread or a pipelined writer that is inside g's fill when the task ends early. f's listener abandons after 1 s, and g's listener then waits for g's `queue.add` (L471-473), whose `allocatePage` can spill f's queue under memory pressure. Round 11 created the directory up front and deleted it here, so the same spill failed with `NoSuchFileException` and an ERROR log instead of leaking. The new Buffered tests use `limit(0)`, which takes the disk fallback of `createNewQueue` and never `HybridQueue.spill`. Suggestion: mark the abandonment under the queue's monitor and let the queue stop spilling afterwards, e.g. a `var abandoned = false` guarded by the queue, `override def spill(size: Long, trigger: MemoryConsumer): Long = synchronized { if (abandoned) 0L else super.spill(size, trigger) }`, and `queue.synchronized { abandoned = true; Utils.deleteQuietly(spillDir) }` here. Returning 0 is better than throwing from `createDiskQueue`, because `HybridQueue.spill` is not exception-safe and `TaskMemoryManager` would log an ERROR. The listener takes only the queue's monitor, so this adds no lock-order cycle. ########## sql/core/src/test/scala/org/apache/spark/sql/execution/python/InProcessPythonUDFSuite.scala: ########## @@ -0,0 +1,406 @@ +/* + * 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.io.File +import java.util.Properties +import java.util.concurrent.{CountDownLatch, TimeUnit} +import java.util.concurrent.atomic.{AtomicInteger, AtomicReference} + +import scala.jdk.CollectionConverters._ + +import org.apache.spark.{SparkEnv, SparkException, TaskContextImpl} +import org.apache.spark.api.python.PythonEvalType +import org.apache.spark.internal.config.PLUGINS +import org.apache.spark.memory.{TaskMemoryManager, TestMemoryManager} +import org.apache.spark.sql.{AnalysisException, Column, QueryTest} +import org.apache.spark.sql.api.python.PythonSQLUtils +import org.apache.spark.sql.catalyst.InternalRow +import org.apache.spark.sql.catalyst.expressions.{AttributeReference, PythonUDF, UnsafeProjection} +import org.apache.spark.sql.catalyst.plans.logical.{Aggregate, ArrowEvalPython, Filter, LocalLimit} +import org.apache.spark.sql.execution.{GlobalLimitExec, ProjectExec, SortExec} +import org.apache.spark.sql.execution.metric.SQLMetric +import org.apache.spark.sql.functions._ +import org.apache.spark.sql.internal.SQLConf +import org.apache.spark.sql.test.SharedSparkSession +import org.apache.spark.sql.types.{DataType, LongType, StructField, StructType} +import org.apache.spark.util.Utils + +/** + * Planning regressions, and evaluator tests that need no Python; runtime coverage lives in + * the PySpark integration suite. + */ +class InProcessPythonUDFSuite extends QueryTest with SharedSparkSession { + + import testImplicits._ + + private val plugin = "org.apache.spark.sql.execution.python.InProcessPythonPlugin" + + override def beforeEach(): Unit = { + super.beforeEach() + // These tests plan queries without loading a native interpreter. Advertise the plugin + // after context creation; actual plugin initialization is covered by integration tests. + SparkEnv.get.conf.set(PLUGINS, Seq(plugin)) + } + + override def afterEach(): Unit = { + try { SparkEnv.get.conf.remove(PLUGINS) } finally { super.afterEach() } + } + + private def makeUDF( + name: String, + input: Column, + deterministic: Boolean = true): Column = { + // Each call creates fresh bytes, as Py4J does. Semantic equality must compare their contents. + InProcessPythonUDFBuilder.build( + name, Array[Byte](1, 2), LongType.json, Seq(input).asJava, deterministic, "3.11") + } + + test("in-process UDFs use PythonUDF and ArrowEvalPython planning contracts") { + val df = spark.range(10) + val doubled = makeUDF("double", df("id")) + val expr = doubled.expr.asInstanceOf[PythonUDF] + assert(expr.evalType == PythonEvalType.SQL_SCALAR_ARROW_INPROCESS_UDF) + assert(expr.expensive) + assert(expr.semanticEquals(makeUDF("double", df("id")).expr)) + + val query = df.select(doubled) + val eval = query.queryExecution.optimizedPlan.collect { case p: ArrowEvalPython => p } + assert(eval.size == 1) + assert(eval.head.evalType == PythonEvalType.SQL_SCALAR_ARROW_INPROCESS_UDF) + val physical = query.queryExecution.executedPlan.collect { + case p: InProcessArrowEvalPythonExec => p + } + assert(physical.size == 1) + assert(physical.head.producedAttributes == + (physical.head.outputSet -- physical.head.child.outputSet)) + assert(physical.head.missingInput.isEmpty) + } + + /** Spill directories of in-process evaluators under the executor's local directory. */ + private def spillDirs(): Set[String] = + Option(new File(Utils.getLocalDir(SparkEnv.get.conf)).listFiles()).toSeq.flatten + .map(_.getName).filter(_.startsWith("inprocess-udf-")).toSet + + /** Review Comment: **[Low, test] `BufferedInput` copies the fixtures of `InProcessPythonRuntimeSuite`, and this block splits the planning tests.** `BufferedInput` repeats `BlockingInput` (`InProcessPythonRuntimeSuite.scala` L313-341) except for the join input, the row count and the context. The metrics at L129-131 repeat `allMetrics()` there (L49-51), and L173-176 repeat its `thread` helper. The copies have already drifted: L188 checks only the exception type, while the RuntimeSuite test now also checks the message, as R11-5 asked. The block also sits between the first planning test (L74-93) and the remaining ones (L191 onward), although the class doc lists the evaluator tests second. Two smaller points: `spillDirs()` lists only the root that `Utils.getLocalDir` picks at random, so it can miss the queue's directory when there are several local directories, and `cleanUpAllAllocatedMemory()` at L182 frees nothing with `limit(0)`. Suggestion: share `allMetrics()` and a blocking input that takes the `JoinInput` and the context through a small `private[python]` test helper object in this package, move this block to the end of the class, add the message check at L188, and list every root of `Utils.getOrCreateLocalRootDirs` in `spillDirs()`. ########## sql/core/src/main/scala/org/apache/spark/sql/execution/python/InProcessArrowEvalPythonEvaluatorFactory.scala: ########## @@ -0,0 +1,511 @@ +/* + * 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.io.File +import java.nio.file.Files +import java.util.UUID +import java.util.concurrent.TimeUnit +import java.util.concurrent.atomic.AtomicInteger +import java.util.concurrent.locks.ReentrantLock + +import scala.collection.mutable.ArrayBuffer +import scala.jdk.CollectionConverters._ + +import com.google.common.util.concurrent.Uninterruptibles +import org.apache.arrow.c.{ArrowArray, ArrowSchema, BaseStruct} +import org.apache.arrow.util.AutoCloseables +import org.apache.arrow.vector.VectorSchemaRoot + +import org.apache.spark.{SparkEnv, SparkException, TaskContext} +import org.apache.spark.api.python.ChainedPythonFunctions +import org.apache.spark.sql.catalyst.InternalRow +import org.apache.spark.sql.catalyst.expressions.{Attribute, Expression, JoinedRow, PythonUDF, UnsafeProjection, UnsafeRow} +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._ +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, unless all of them are UDF arguments that read back from Arrow unchanged. + * Each batch owns its Arrow buffers so Python can safely retain input arrays. + * + * The evaluator owns its queue, so that cleanup at task completion is coordinated with a + * consumer on another thread, such as a pipelined Python writer or a TRANSFORM feed thread. + */ +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, + fullValidation: Boolean, + metrics: Map[String, SQLMetric]) + extends EvalPythonEvaluatorFactory(childOutput, udfs, output) { + + private[python] def runtimeSession: InProcessPythonRuntime.InterpreterSession = + InProcessPythonRuntime.currentSession + + /** Unused: `evaluateJoined` always evaluates the UDFs. */ + override protected def evaluate( + funcs: Seq[(ChainedPythonFunctions, Long)], + argMetas: Array[Array[ArgumentMetadata]], + rows: Iterator[InternalRow], + inputSchema: StructType, + context: TaskContext): Iterator[InternalRow] = + throw SparkException.internalError("In-process UDFs are evaluated with their input rows") + + override protected def evaluateJoined( + funcs: Seq[(ChainedPythonFunctions, Long)], + argMetas: Array[Array[ArgumentMetadata]], + rows: Iterator[InternalRow], + inputs: Seq[Expression], + inputSchema: StructType, + context: TaskContext): Option[Iterator[InternalRow]] = { + import InProcessArrowEvalPythonEvaluatorFactory.{Buffered, ReadBack, readsBack} + val inputColumns = inputs.length == childOutput.length && inputs.zip(childOutput).forall { + case (a: Attribute, c) => a.exprId == c.exprId + case _ => false + } + // If all input columns are UDF arguments, they are written to Arrow regardless. Read them + // back from the exported input vectors instead of buffering every input row, if their + // values read back from Arrow exactly as written and as fast as an unsafe row copy. + val joinInput = if (inputColumns && inputSchema.forall(f => readsBack(f.dataType))) { + ReadBack + } else if (inputColumns) { + Buffered(None) + } else { + // Each projected row is written to Arrow before the next input row is pulled, so the + // arguments go into a reused buffer rather than being copied value by value. + val projection = UnsafeProjection.create(inputs, childOutput) + projection.initialize(context.partitionId()) + Buffered(Some(projection)) + } + Some(evaluateBatches(funcs, argMetas, rows, inputSchema, context, joinInput)) + } + + private[python] def evaluateBatches( + funcs: Seq[(ChainedPythonFunctions, Long)], + argMetas: Array[Array[ArgumentMetadata]], + rows: Iterator[InternalRow], + inputSchema: StructType, + context: TaskContext, + joinInput: InProcessArrowEvalPythonEvaluatorFactory.JoinInput): Iterator[InternalRow] = { + import InProcessArrowEvalPythonEvaluatorFactory.{Buffered, ReadBack} + ArrowUtils.failDuplicatedFieldNames(inputSchema) + val functions = funcs.map { case (chain, _) => + if (chain.funcs.size != 1) { + throw SparkException.internalError( + "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) + // Capture before consuming input: an old task must never join a later context's session. + val runtime = runtimeSession + // Rows are copied out of the queue and Arrow vectors before they are returned, so they + // remain valid after task completion releases those, on whichever thread consumes them. + val resultProj = UnsafeProjection.create(output, output) + // Spill files go into a directory of the queue's own, created on the first spill, so that + // task completion can delete them when it cannot close the queue. + @volatile var spillDir: File = null + val (queue, projection) = joinInput match { + case Buffered(projection) => + val localDir = new File(Utils.getLocalDir(SparkEnv.get.conf)) + val serializerManager = SparkEnv.get.serializerManager + // Only the consumer holding the iterator's lock adds and removes rows. + val queue = new HybridRowQueue(context.taskMemoryManager(), localDir, Review Comment: **[Low] In-process queues of the same width in a task are now equal, so `TaskMemoryManager` keeps only one of them.** This anonymous class inherits the case-class `equals` and `hashCode` of `HybridRowQueue` (`RowQueue.scala` L185), which compare `(memManager, tempDir, numFields, serMgr, lockFree)`. Every queue now gets the same `localDir` as `tempDir`, while round 11 gave each one its own directory. `TaskMemoryManager.consumers` is a `HashSet` (`TaskMemoryManager.java` L119, L252), so an equal second queue is never registered and never offered for spilling (L227-233), and `HybridQueue.spill` returns 0 for an equal trigger (`HybridQueue.scala` L73, where `==` calls `equals`). For example, in `SELECT a, f(g(b))` with both UDFs in-process, column pruning leaves two buffered columns in both nodes. So the upper queue can never be spilled for a sort or an aggregation, the lower queue refuses to spill for the upper one, and the OOM breakdown does not attribute the upper queue's memory. The same applies to the partitions that `coalesce` evaluates in one task, and to the two sides of a sort-merge join. Suggestion: identity equality in this anonymous class, e.g. `override def equals(o: Any): Boolean = this eq o.asInstanceOf[AnyRef]`, `override def hashCode(): Int = System.identityHashCode(this)` and `override def canEqual(o: Any): Boolean = false`. The regular queue in `EvalPythonEvaluatorFactory.scala` L121-125 has the same latent issue, so `HybridRowQueue` itself may deserve a separate JIRA. ########## sql/core/src/main/scala/org/apache/spark/sql/execution/python/InProcessArrowEvalPythonEvaluatorFactory.scala: ########## @@ -0,0 +1,511 @@ +/* + * 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.io.File +import java.nio.file.Files +import java.util.UUID +import java.util.concurrent.TimeUnit +import java.util.concurrent.atomic.AtomicInteger +import java.util.concurrent.locks.ReentrantLock + +import scala.collection.mutable.ArrayBuffer +import scala.jdk.CollectionConverters._ + +import com.google.common.util.concurrent.Uninterruptibles +import org.apache.arrow.c.{ArrowArray, ArrowSchema, BaseStruct} +import org.apache.arrow.util.AutoCloseables +import org.apache.arrow.vector.VectorSchemaRoot + +import org.apache.spark.{SparkEnv, SparkException, TaskContext} +import org.apache.spark.api.python.ChainedPythonFunctions +import org.apache.spark.sql.catalyst.InternalRow +import org.apache.spark.sql.catalyst.expressions.{Attribute, Expression, JoinedRow, PythonUDF, UnsafeProjection, UnsafeRow} +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._ +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, unless all of them are UDF arguments that read back from Arrow unchanged. + * Each batch owns its Arrow buffers so Python can safely retain input arrays. + * + * The evaluator owns its queue, so that cleanup at task completion is coordinated with a + * consumer on another thread, such as a pipelined Python writer or a TRANSFORM feed thread. + */ +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, + fullValidation: Boolean, + metrics: Map[String, SQLMetric]) + extends EvalPythonEvaluatorFactory(childOutput, udfs, output) { + + private[python] def runtimeSession: InProcessPythonRuntime.InterpreterSession = + InProcessPythonRuntime.currentSession + + /** Unused: `evaluateJoined` always evaluates the UDFs. */ + override protected def evaluate( + funcs: Seq[(ChainedPythonFunctions, Long)], + argMetas: Array[Array[ArgumentMetadata]], + rows: Iterator[InternalRow], + inputSchema: StructType, + context: TaskContext): Iterator[InternalRow] = + throw SparkException.internalError("In-process UDFs are evaluated with their input rows") + + override protected def evaluateJoined( + funcs: Seq[(ChainedPythonFunctions, Long)], + argMetas: Array[Array[ArgumentMetadata]], + rows: Iterator[InternalRow], + inputs: Seq[Expression], + inputSchema: StructType, + context: TaskContext): Option[Iterator[InternalRow]] = { + import InProcessArrowEvalPythonEvaluatorFactory.{Buffered, ReadBack, readsBack} + val inputColumns = inputs.length == childOutput.length && inputs.zip(childOutput).forall { + case (a: Attribute, c) => a.exprId == c.exprId + case _ => false + } + // If all input columns are UDF arguments, they are written to Arrow regardless. Read them + // back from the exported input vectors instead of buffering every input row, if their + // values read back from Arrow exactly as written and as fast as an unsafe row copy. + val joinInput = if (inputColumns && inputSchema.forall(f => readsBack(f.dataType))) { + ReadBack + } else if (inputColumns) { + Buffered(None) + } else { + // Each projected row is written to Arrow before the next input row is pulled, so the + // arguments go into a reused buffer rather than being copied value by value. + val projection = UnsafeProjection.create(inputs, childOutput) + projection.initialize(context.partitionId()) + Buffered(Some(projection)) + } + Some(evaluateBatches(funcs, argMetas, rows, inputSchema, context, joinInput)) + } + + private[python] def evaluateBatches( + funcs: Seq[(ChainedPythonFunctions, Long)], + argMetas: Array[Array[ArgumentMetadata]], + rows: Iterator[InternalRow], + inputSchema: StructType, + context: TaskContext, + joinInput: InProcessArrowEvalPythonEvaluatorFactory.JoinInput): Iterator[InternalRow] = { + import InProcessArrowEvalPythonEvaluatorFactory.{Buffered, ReadBack} + ArrowUtils.failDuplicatedFieldNames(inputSchema) + val functions = funcs.map { case (chain, _) => + if (chain.funcs.size != 1) { + throw SparkException.internalError( + "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) + // Capture before consuming input: an old task must never join a later context's session. + val runtime = runtimeSession + // Rows are copied out of the queue and Arrow vectors before they are returned, so they + // remain valid after task completion releases those, on whichever thread consumes them. + val resultProj = UnsafeProjection.create(output, output) + // Spill files go into a directory of the queue's own, created on the first spill, so that + // task completion can delete them when it cannot close the queue. + @volatile var spillDir: File = null + val (queue, projection) = joinInput match { + case Buffered(projection) => + val localDir = new File(Utils.getLocalDir(SparkEnv.get.conf)) + val serializerManager = SparkEnv.get.serializerManager + // Only the consumer holding the iterator's lock adds and removes rows. + val queue = new HybridRowQueue(context.taskMemoryManager(), localDir, + childOutput.length, serializerManager, lockFree = true) { + override protected def createDiskQueue(): RowQueue = synchronized { + if (spillDir == null) { + spillDir = Files.createTempDirectory(localDir.toPath, "inprocess-udf-").toFile + } + DiskRowQueue(Files.createTempFile(spillDir.toPath, "buffer", "").toFile, + childOutput.length, serializerManager) + } + } + (queue, projection.orNull) + case ReadBack => (null, null) + } + val joined = new JoinedRow + val handles = functions.map(_ => UUID.randomUUID().toString) + var registered = false + var writer: ArrowWriter = null + val results = ArrayBuffer.empty[ArrowColumnVector] + var startedAt = 0L + + def closeBatch(): Unit = { + val resources = ArrayBuffer.empty[AutoCloseable] + resources ++= results + results.clear() + if (writer != null) { + resources += writer.root + writer = null + } + AutoCloseables.close(resources.asJava) + } + + val resources = new InProcessArrowEvalPythonEvaluatorFactory.IteratorResources( + // Closing the queue deletes its spill files, which leaves an empty directory. + releaseTaskMemory = () => if (queue != null) { + Utils.tryWithSafeFinally(queue.close()) { if (spillDir != null) spillDir.delete() } + }, + // Neither starts a process nor throws, also on an interrupted thread. + abandonTaskMemory = () => if (spillDir != null) Utils.deleteQuietly(spillDir), + releaseOthers = () => { + if (startedAt != 0L) { + metrics("pythonTotalTime") += (System.nanoTime() - startedAt) / 1000000 + } + Utils.tryWithSafeFinally { + closeBatch() + } { + if (registered) runtime.release(handles) + } + }) + + context.addTaskCompletionListener[Unit](_ => resources.close()) + + new Iterator[InternalRow] { + private var batchIter: Iterator[InternalRow] = Iterator.empty + + private def endOfInput: Nothing = + throw new NoSuchElementException("End of in-process UDF input") + + // Releases the resources on failure without replacing its exception. + private def fail(t: Throwable): Nothing = + Utils.tryWithSafeFinally { throw t } { resources.close() } + + // Called with the lock held. + private def hasNextLocked: Boolean = { + if (startedAt == 0L) startedAt = System.nanoTime() + checkCancellation() + val available = batchIter.hasNext || rows.hasNext + if (!available) resources.close() + available + } + + // Each call takes the lock without allocating a closure per row. Within a batch, only + // the consumer advances `batchIter`, whose `hasNext` compares row indexes. + override def hasNext: Boolean = { + if (batchIter.hasNext && !resources.isClosed) return true + if (!resources.enter()) return false + val available = try { + try hasNextLocked catch { case t: Throwable => fail(t) } + } finally { + resources.exit() + } + // Task completion may have closed the iterator while the input was being read. + available && !resources.isClosed + } + + override def next(): InternalRow = { + if (!resources.enter()) endOfInput + try { + try { + if (!hasNextLocked) endOfInput + if (!batchIter.hasNext) nextBatch() + val result = batchIter.next() + resultProj(if (queue != null) joined(queue.remove(), result) else result) + } catch { + case t: Throwable => fail(t) + } + } finally { + resources.exit() + } + } + + // Runs Python without the lock, unless task completion happened before or meanwhile. + private def python[T](body: => T): T = { + if (resources.isClosed) endOfInput + val result = resources.withoutLock(body) + if (resources.isClosed) endOfInput + result + } + + /** + * Writes the next input row to the batch, returning false at the end of input. Task + * completion can happen while the input is read; then the row is not written. + */ + private def pullRow(): Boolean = rows.hasNext && !resources.isClosed && { + val row = rows.next() + if (queue != null) { + // Adding can wait for memory beyond the listener's wait. Announce it before the + // check, so that either the add is skipped or the listener waits for it. + resources.usingTaskMemory = true Review Comment: **[Low, performance] Refining my round-11 suggestion: the per-row flag adds two StoreLoad barriers to every buffered row.** As I suggested in https://github.com/apache/spark/pull/58978#discussion_r4189042183, the consumer now sets and clears the volatile `usingTaskMemory` around every `queue.add`. On x86, each volatile store is followed by a `lock addl`, so a buffered row now pays two more full barriers, on top of the `writeOffset` store in `InMemoryRowQueue.doAdd` (`RowQueue.scala` L92) and the lock and unlock in `next()`. My rough estimate, not a measurement, is 5-15% for narrow buffered rows, e.g. `df.withColumn("y", f("x"))` with other columns kept. The benchmarks in the PR description read back (`bench_inprocess_udf.py` L127 passes every column to the UDF), so they do not show it. Suggestion: could you measure a Buffered case on x86 against round 11? If it matters, the flag could be raised only when the add allocates a page, since only page allocation can wait for other tasks' memory: e.g. override the protected `MemoryConsumer.allocatePage(long)` in this anonymous queue to set the flag and check `isClosed`, and clear the flag when `queue.add` returns. The trade-off is that a non-blocking add stalled for over 1 s, e.g. by a long GC pause, would no longer hold off the listener. ########## sql/core/src/test/scala/org/apache/spark/sql/execution/python/InProcessPythonRuntimeSuite.scala: ########## @@ -0,0 +1,577 @@ +/* + * 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.Collections +import java.util.concurrent.{CountDownLatch, TimeUnit} +import java.util.concurrent.atomic.{AtomicBoolean, AtomicInteger, AtomicReference} + +import org.mockito.Mockito.{mock, when} + +import org.apache.spark.{SparkConf, SparkFunSuite, SparkIllegalArgumentException, TaskContext, TaskKilledException} +import org.apache.spark.api.plugin.PluginContext +import org.apache.spark.api.python.{ChainedPythonFunctions, PythonEvalType, SimplePythonFunction} +import org.apache.spark.internal.config.Python.{IN_PROCESS_PATH_RULE, IN_PROCESS_SITE_PACKAGES} +import org.apache.spark.sql.catalyst.InternalRow +import org.apache.spark.sql.catalyst.expressions.{AttributeReference, PythonUDF} +import org.apache.spark.sql.execution.metric.SQLMetric +import org.apache.spark.sql.types.{LongType, StructField, StructType} +import org.apache.spark.sql.util.ArrowUtils + +class InProcessPythonRuntimeSuite extends SparkFunSuite { + private var runtime: InProcessPythonRuntime.InterpreterSession = _ + + override def beforeEach(): Unit = { + super.beforeEach() + runtime = new InProcessPythonRuntime.InterpreterSession() + } + + override def afterEach(): Unit = { + try { runtime.shutdown() } finally { super.afterEach() } + } + + /** Every metric that an evaluator may update, as `PythonSQLMetrics` defines them. */ + private def allMetrics(): Map[String, SQLMetric] = + (PythonSQLMetrics.pythonSizeMetricsDesc ++ PythonSQLMetrics.pythonTimingMetricsDesc ++ + PythonSQLMetrics.pythonOtherMetricsDesc).keys.map(_ -> new SQLMetric("sum", 0L)).toMap + + test("site-packages config validates JEP include paths") { + val conf = new SparkConf(false) + assert(conf.get(IN_PROCESS_SITE_PACKAGES).isEmpty) + conf.set(IN_PROCESS_SITE_PACKAGES.key, " /opt/venv/lib, /opt/extra ") + assert(conf.get(IN_PROCESS_SITE_PACKAGES) == Seq("/opt/venv/lib", "/opt/extra")) + conf.set(IN_PROCESS_SITE_PACKAGES.key, "back\\slash") + assert(conf.get(IN_PROCESS_SITE_PACKAGES) == Seq("back\\slash")) + Seq("bad'path", "bad\npath", "bad\rpath", "bad\u0000path", + "bad" + new String(Character.toChars(0x1f600)), "bad" + 0xd800.toChar, + s"bad${java.io.File.pathSeparator}path") + .foreach { path => + conf.set(IN_PROCESS_SITE_PACKAGES.key, path) + intercept[IllegalArgumentException] { conf.get(IN_PROCESS_SITE_PACKAGES) } + intercept[IllegalArgumentException] { + InProcessPythonRuntime.InterpreterConfiguration.interpreterConfig(Seq(path)) + } + } + } + + test("registration failure frees its temporary native command buffer") { + val before = ArrowUtils.rootAllocator.getAllocatedMemory + val field = ArrowUtils.toArrowField("result", LongType, true, "UTC") + intercept[NullPointerException] { + // This session deliberately has no interpreter, so invocation fails after allocation. + runtime.register( + "failed", new Array[Byte](1024 * 1024), field, "3.12", false, false, false, true) + } + assert(ArrowUtils.rootAllocator.getAllocatedMemory == before) + runtime.shutdown(waitMillis = 20) + assert(!runtime.isTerminated) + runtime.release(Seq("failed")) + runtime.shutdown() + assert(runtime.isTerminated) + } + + test("plugin reports invalid sitePackages without the installation checklist") { + val ctx = mock(classOf[PluginContext]) + when(ctx.conf()).thenReturn(new SparkConf().set(IN_PROCESS_SITE_PACKAGES.key, "/a'b")) + val e = intercept[SparkIllegalArgumentException] { + new InProcessPythonExecutorPlugin().init(ctx, Collections.emptyMap()) + } + assert(e.getCondition == "INVALID_CONF_VALUE.REQUIREMENT") + assert(e.getMessage.contains(IN_PROCESS_PATH_RULE) && !e.getMessage.contains("libjep")) + } + + test("task-side calls after shutdown report the shutdown") { + runtime.shutdown() + val field = ArrowUtils.toArrowField("result", LongType, true, "UTC") + Seq( + () => runtime.onInterpreterThread(()), + () => runtime.register("stopped", Array.emptyByteArray, field, "3.12", + false, false, false, true) + ).foreach { call => + val e = intercept[IllegalStateException] { call() } + assert(e.getMessage.contains("has been stopped")) + } + } + + test("lifecycle errors distinguish configuration mismatch from stopping") { + val mismatch = intercept[InProcessPythonRuntime.LifecycleException] { + runtime.requireCompatible(Seq("different")) + } + assert(mismatch.getMessage.contains("different sitePackages")) + runtime.shutdown() + val stopping = intercept[InProcessPythonRuntime.LifecycleException] { + runtime.requireCompatible(Seq.empty) + } + assert(stopping.getMessage.contains("still stopping")) + } + + test("sub-millisecond invocations accumulate in processing metrics") { + val metric = new SQLMetric("timing", 0L) + val timer = new InProcessArrowEvalPythonEvaluatorFactory.NanosecondTimer(metric) + (1 to 25).foreach(_ => timer.add(100000L)) + assert(metric.value == 2L) + timer.add(500000L) + assert(metric.value == 3L) + } + + test("unused evaluator iterators do not charge Python total time") { + val metrics = allMetrics() + val context = TaskContext.empty() + class TestEvaluator extends InProcessArrowEvalPythonEvaluatorFactory( + Seq.empty, Seq.empty, Seq.empty, 10, 0L, "UTC", false, false, false, false, true, metrics) { + override private[python] def runtimeSession: InProcessPythonRuntime.InterpreterSession = + runtime + + def createUnusedIterator(): Unit = { + evaluateBatches(Seq.empty, Array.empty, Iterator.empty, new StructType, context, + InProcessArrowEvalPythonEvaluatorFactory.ReadBack) + } + } + new TestEvaluator().createUnusedIterator() + Thread.sleep(20) + context.markTaskCompleted(None) + assert(metrics("pythonTotalTime").value == 0L) + } + + test("evaluators retain the generation captured before consuming any input") { + val metrics = allMetrics() + val context = TaskContext.empty() + val function = SimplePythonFunction( + Seq.empty, Collections.emptyMap[String, String](), Collections.emptyList[String](), + "", "3.12", Collections.emptyList(), null) + val udf = PythonUDF("identity", function, LongType, Seq.empty, + PythonEvalType.SQL_SCALAR_ARROW_INPROCESS_UDF, true) + var lookups = 0 + class TestEvaluator extends InProcessArrowEvalPythonEvaluatorFactory( + Seq.empty, Seq(udf), Seq.empty, 10, 0L, "UTC", false, false, false, false, true, metrics) { + override private[python] def runtimeSession: InProcessPythonRuntime.InterpreterSession = { + lookups += 1 + runtime + } + + def iterator(): Iterator[InternalRow] = evaluateBatches( + Seq((ChainedPythonFunctions(Seq(function)), 0L)), Array(Array.empty), + Iterator.single(InternalRow.empty), new StructType, context, + InProcessArrowEvalPythonEvaluatorFactory.ReadBack) + } + val iterator = new TestEvaluator().iterator() + assert(lookups == 1) + runtime.shutdown() + runtime = new InProcessPythonRuntime.InterpreterSession() + try { + val error = intercept[IllegalStateException] { iterator.next() } + assert(error.getMessage.contains("has been stopped")) + assert(lookups == 1) + } finally { + context.markTaskCompleted(None) + } + } + + private class Releases { + val taskMemory = new AtomicInteger() + val abandoned = new AtomicInteger() + val others = new AtomicInteger() + + def resources(lockWaitMillis: Long = 10000L) + : InProcessArrowEvalPythonEvaluatorFactory.IteratorResources = + new InProcessArrowEvalPythonEvaluatorFactory.IteratorResources( + () => taskMemory.incrementAndGet(), + () => abandoned.incrementAndGet(), + () => others.incrementAndGet(), + lockWaitMillis) + } + + private def thread(body: => Unit): Thread = { + val t = new Thread(() => body) + t.start() + t + } + + /** + * Runs `test` while a consumer on another thread is inside a call, optionally running + * Python, until `test` returns. The consumer is released and joined even if `test` fails. + */ + private def withConsumer( + resources: InProcessArrowEvalPythonEvaluatorFactory.IteratorResources, + inPython: Boolean = false)(test: => Unit): Boolean = { + val entered = new CountDownLatch(1) + val finish = new CountDownLatch(1) + val closedAfterCall = new AtomicBoolean() + val consumer = thread { + assert(resources.enter()) + try { + if (inPython) { + resources.withoutLock { entered.countDown(); finish.await(10, TimeUnit.SECONDS) } + } else { + entered.countDown() + finish.await(10, TimeUnit.SECONDS) + } + closedAfterCall.set(resources.isClosed) + } finally { + resources.exit() + } + } + try { + assert(entered.await(10, TimeUnit.SECONDS)) + test + } finally { + finish.countDown() + consumer.join(10000) + } + assert(!consumer.isAlive) + closedAfterCall.get + } + + test("task completion waits for the consumer's lock and stops later calls") { + val releases = new Releases + val resources = releases.resources() + var closing: Thread = null + withConsumer(resources) { + closing = thread(resources.close()) + closing.join(200) + // Nothing is released while the consumer reads input, the queue or Arrow vectors. + assert(closing.isAlive && resources.isClosed && releases.taskMemory.get == 0) + } + closing.join(10000) + assert(!closing.isAlive && releases.taskMemory.get == 1 && releases.others.get == 1) + assert(!resources.enter()) + assert(releases.taskMemory.get == 1 && releases.others.get == 1) + } + + test("task completion releases task memory at once while Python runs") { + val releases = new Releases + val resources = releases.resources() + val closedAfterPython = withConsumer(resources, inPython = true) { + resources.close() + // The listener does not wait for Python, but keeps the Arrow vectors Python may use. + assert(releases.taskMemory.get == 1 && releases.others.get == 0) + } + assert(closedAfterPython && releases.taskMemory.get == 1 && releases.others.get == 1) + } + + test("task completion waits only briefly for a consumer blocked on its input") { + val releases = new Releases + val resources = releases.resources(lockWaitMillis = 50L) + withConsumer(resources) { + resources.close() + // The executor frees the task memory, after the listener deletes what lives outside it. + assert(releases.taskMemory.get == 0 && releases.abandoned.get == 1) + assert(releases.others.get == 0) + } + assert(releases.taskMemory.get == 0 && releases.others.get == 1) + } + + test("an interrupted completion listener still waits for the consumer's lock") { + val releases = new Releases + val resources = releases.resources() + val interrupted = new AtomicBoolean() + var closing: Thread = null + withConsumer(resources) { + closing = thread { + Thread.currentThread().interrupt() + resources.close() + interrupted.set(Thread.currentThread().isInterrupted) + } + closing.join(200) + assert(closing.isAlive && releases.taskMemory.get == 0) + } + closing.join(10000) + assert(!closing.isAlive && interrupted.get) + assert(releases.taskMemory.get == 1 && releases.others.get == 1) + } + + test("exhausted iterators close within a call and return no more rows") { + val releases = new Releases + val resources = releases.resources() + assert(resources.enter()) + resources.close() + resources.exit() + assert(!resources.enter()) + resources.close() + assert(releases.taskMemory.get == 1 && releases.others.get == 1) + } + + /** + * An evaluator without UDFs, which reads its single input column back from Arrow, so that + * its iterator runs without Python. `rows` blocks on `gate` before reading row `blockAt`. + */ + private class BlockingInput(blockAt: Int) { + val reached = new CountDownLatch(1) + val gate = new CountDownLatch(1) + val pulled = new AtomicInteger() + val context = TaskContext.empty() + private val column = AttributeReference("x", LongType)() + + val rows: Iterator[InternalRow] = new Iterator[InternalRow] { + private def block(): Unit = if (pulled.get == blockAt) { + reached.countDown() + gate.await(10, TimeUnit.SECONDS) + } + override def hasNext: Boolean = { block(); true } + override def next(): InternalRow = { + block() + InternalRow(pulled.incrementAndGet().toLong) + } + } + + def iterator(): Iterator[InternalRow] = { + val metrics = allMetrics() + new InProcessArrowEvalPythonEvaluatorFactory(Seq(column), Seq.empty, Seq(column), 10, + 0L, "UTC", false, false, false, false, true, metrics) { + override private[python] def runtimeSession = runtime + }.evaluateBatches(Seq.empty, Array.empty, rows, + StructType(Seq(StructField("x", LongType))), context, + InProcessArrowEvalPythonEvaluatorFactory.ReadBack) + } + } + + test("task completion stops a batch fill within one input row") { + val input = new BlockingInput(blockAt = 3) + val iterator = input.iterator() + val error = new AtomicReference[Throwable]() + val consumer = thread { + try iterator.next() catch { case t: Throwable => error.set(t) } + } + try { + assert(input.reached.await(10, TimeUnit.SECONDS)) + val closing = thread(input.context.markTaskCompleted(None)) + closing.join(200) + assert(closing.isAlive) + input.gate.countDown() + closing.join(10000) + assert(!closing.isAlive) + } finally { + input.gate.countDown() + consumer.join(10000) + } + // The fill stops at the row it was waiting for, without reading it. + assert(error.get.isInstanceOf[NoSuchElementException]) + assert(error.get.getMessage == "End of in-process UDF input" && input.pulled.get == 3) Review Comment: **[Low, test] No test runs the new `usingTaskMemory` branch or the checks after `rows.next()`.** No test refers to `usingTaskMemory`, so removing the branch at L471-473 of the evaluator, or moving it after the CAS, would still pass every suite and bring back the race of R11-2. Also, with the new check before `rows.next()`, this test now stops before it reads the row, so nothing runs the checks after `rows.next()` any more: the consumer side of the handshake (L271 of the evaluator) and the ReadBack check (L276-277). Round 11's `pulled == 4` ran the latter. Suggestion: a `Releases`-based test next to L267-277, in which the consumer sets the flag inside `withConsumer` and `close()` with `lockWaitMillis = 50L` must wait instead of abandoning (`abandoned == 0` while the consumer runs, `taskMemory == 1` afterwards), and an input variant that blocks only in `next()`, checking `pulled == blockAt + 1` and the "End of in-process UDF input" message for both ReadBack and Buffered. ########## sql/core/src/main/scala/org/apache/spark/sql/execution/python/InProcessArrowEvalPythonEvaluatorFactory.scala: ########## @@ -0,0 +1,511 @@ +/* + * 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.io.File +import java.nio.file.Files +import java.util.UUID +import java.util.concurrent.TimeUnit +import java.util.concurrent.atomic.AtomicInteger +import java.util.concurrent.locks.ReentrantLock + +import scala.collection.mutable.ArrayBuffer +import scala.jdk.CollectionConverters._ + +import com.google.common.util.concurrent.Uninterruptibles +import org.apache.arrow.c.{ArrowArray, ArrowSchema, BaseStruct} +import org.apache.arrow.util.AutoCloseables +import org.apache.arrow.vector.VectorSchemaRoot + +import org.apache.spark.{SparkEnv, SparkException, TaskContext} +import org.apache.spark.api.python.ChainedPythonFunctions +import org.apache.spark.sql.catalyst.InternalRow +import org.apache.spark.sql.catalyst.expressions.{Attribute, Expression, JoinedRow, PythonUDF, UnsafeProjection, UnsafeRow} +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._ +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, unless all of them are UDF arguments that read back from Arrow unchanged. + * Each batch owns its Arrow buffers so Python can safely retain input arrays. + * + * The evaluator owns its queue, so that cleanup at task completion is coordinated with a + * consumer on another thread, such as a pipelined Python writer or a TRANSFORM feed thread. + */ +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, + fullValidation: Boolean, + metrics: Map[String, SQLMetric]) + extends EvalPythonEvaluatorFactory(childOutput, udfs, output) { + + private[python] def runtimeSession: InProcessPythonRuntime.InterpreterSession = + InProcessPythonRuntime.currentSession + + /** Unused: `evaluateJoined` always evaluates the UDFs. */ + override protected def evaluate( + funcs: Seq[(ChainedPythonFunctions, Long)], + argMetas: Array[Array[ArgumentMetadata]], + rows: Iterator[InternalRow], + inputSchema: StructType, + context: TaskContext): Iterator[InternalRow] = + throw SparkException.internalError("In-process UDFs are evaluated with their input rows") + + override protected def evaluateJoined( + funcs: Seq[(ChainedPythonFunctions, Long)], + argMetas: Array[Array[ArgumentMetadata]], + rows: Iterator[InternalRow], + inputs: Seq[Expression], + inputSchema: StructType, + context: TaskContext): Option[Iterator[InternalRow]] = { + import InProcessArrowEvalPythonEvaluatorFactory.{Buffered, ReadBack, readsBack} + val inputColumns = inputs.length == childOutput.length && inputs.zip(childOutput).forall { + case (a: Attribute, c) => a.exprId == c.exprId + case _ => false + } + // If all input columns are UDF arguments, they are written to Arrow regardless. Read them + // back from the exported input vectors instead of buffering every input row, if their + // values read back from Arrow exactly as written and as fast as an unsafe row copy. + val joinInput = if (inputColumns && inputSchema.forall(f => readsBack(f.dataType))) { + ReadBack + } else if (inputColumns) { + Buffered(None) + } else { + // Each projected row is written to Arrow before the next input row is pulled, so the + // arguments go into a reused buffer rather than being copied value by value. + val projection = UnsafeProjection.create(inputs, childOutput) + projection.initialize(context.partitionId()) + Buffered(Some(projection)) + } + Some(evaluateBatches(funcs, argMetas, rows, inputSchema, context, joinInput)) + } + + private[python] def evaluateBatches( + funcs: Seq[(ChainedPythonFunctions, Long)], + argMetas: Array[Array[ArgumentMetadata]], + rows: Iterator[InternalRow], + inputSchema: StructType, + context: TaskContext, + joinInput: InProcessArrowEvalPythonEvaluatorFactory.JoinInput): Iterator[InternalRow] = { + import InProcessArrowEvalPythonEvaluatorFactory.{Buffered, ReadBack} + ArrowUtils.failDuplicatedFieldNames(inputSchema) + val functions = funcs.map { case (chain, _) => + if (chain.funcs.size != 1) { + throw SparkException.internalError( + "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) + // Capture before consuming input: an old task must never join a later context's session. + val runtime = runtimeSession + // Rows are copied out of the queue and Arrow vectors before they are returned, so they + // remain valid after task completion releases those, on whichever thread consumes them. + val resultProj = UnsafeProjection.create(output, output) + // Spill files go into a directory of the queue's own, created on the first spill, so that + // task completion can delete them when it cannot close the queue. + @volatile var spillDir: File = null + val (queue, projection) = joinInput match { + case Buffered(projection) => + val localDir = new File(Utils.getLocalDir(SparkEnv.get.conf)) + val serializerManager = SparkEnv.get.serializerManager + // Only the consumer holding the iterator's lock adds and removes rows. + val queue = new HybridRowQueue(context.taskMemoryManager(), localDir, + childOutput.length, serializerManager, lockFree = true) { + override protected def createDiskQueue(): RowQueue = synchronized { + if (spillDir == null) { + spillDir = Files.createTempDirectory(localDir.toPath, "inprocess-udf-").toFile + } + DiskRowQueue(Files.createTempFile(spillDir.toPath, "buffer", "").toFile, + childOutput.length, serializerManager) + } + } + (queue, projection.orNull) + case ReadBack => (null, null) + } + val joined = new JoinedRow + val handles = functions.map(_ => UUID.randomUUID().toString) + var registered = false + var writer: ArrowWriter = null + val results = ArrayBuffer.empty[ArrowColumnVector] + var startedAt = 0L + + def closeBatch(): Unit = { + val resources = ArrayBuffer.empty[AutoCloseable] + resources ++= results + results.clear() + if (writer != null) { + resources += writer.root + writer = null + } + AutoCloseables.close(resources.asJava) + } + + val resources = new InProcessArrowEvalPythonEvaluatorFactory.IteratorResources( + // Closing the queue deletes its spill files, which leaves an empty directory. + releaseTaskMemory = () => if (queue != null) { + Utils.tryWithSafeFinally(queue.close()) { if (spillDir != null) spillDir.delete() } + }, + // Neither starts a process nor throws, also on an interrupted thread. + abandonTaskMemory = () => if (spillDir != null) Utils.deleteQuietly(spillDir), + releaseOthers = () => { + if (startedAt != 0L) { + metrics("pythonTotalTime") += (System.nanoTime() - startedAt) / 1000000 + } + Utils.tryWithSafeFinally { + closeBatch() + } { + if (registered) runtime.release(handles) + } + }) + + context.addTaskCompletionListener[Unit](_ => resources.close()) + + new Iterator[InternalRow] { + private var batchIter: Iterator[InternalRow] = Iterator.empty + + private def endOfInput: Nothing = + throw new NoSuchElementException("End of in-process UDF input") + + // Releases the resources on failure without replacing its exception. + private def fail(t: Throwable): Nothing = + Utils.tryWithSafeFinally { throw t } { resources.close() } + + // Called with the lock held. + private def hasNextLocked: Boolean = { + if (startedAt == 0L) startedAt = System.nanoTime() + checkCancellation() + val available = batchIter.hasNext || rows.hasNext + if (!available) resources.close() + available + } + + // Each call takes the lock without allocating a closure per row. Within a batch, only + // the consumer advances `batchIter`, whose `hasNext` compares row indexes. + override def hasNext: Boolean = { + if (batchIter.hasNext && !resources.isClosed) return true + if (!resources.enter()) return false + val available = try { + try hasNextLocked catch { case t: Throwable => fail(t) } + } finally { + resources.exit() + } + // Task completion may have closed the iterator while the input was being read. + available && !resources.isClosed + } + + override def next(): InternalRow = { + if (!resources.enter()) endOfInput + try { + try { + if (!hasNextLocked) endOfInput + if (!batchIter.hasNext) nextBatch() + val result = batchIter.next() + resultProj(if (queue != null) joined(queue.remove(), result) else result) + } catch { + case t: Throwable => fail(t) + } + } finally { + resources.exit() + } + } + + // Runs Python without the lock, unless task completion happened before or meanwhile. Review Comment: **[Low, cleanup] A few comments no longer match the code.** - Here, "unless task completion happened before or meanwhile": if it happens meanwhile, Python has already run, and L256 ends the input instead of returning the result. - L260-262: `pullRow` now also returns false once task completion happened (L264), not only at the end of input. - L144-145, "created on the first spill": the disk fallback of `createNewQueue` (`HybridQueue.scala` L115-120), which the new tests use, creates it too. - L423-424, "while it may allocate task memory": the flag is set only around `queue.add`, not while the consumer reads input that may allocate. Suggestion: e.g. "Runs Python without the lock unless task completion already happened, and ends the input instead of returning the result if it happens meanwhile", "returning false at the end of input or once task completion happened", "created with the first disk queue", and "while it adds a row to the queue, which may allocate a page and wait for other tasks' memory". ########## sql/core/src/main/scala/org/apache/spark/sql/execution/python/InProcessArrowEvalPythonEvaluatorFactory.scala: ########## @@ -0,0 +1,511 @@ +/* + * 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.io.File +import java.nio.file.Files +import java.util.UUID +import java.util.concurrent.TimeUnit +import java.util.concurrent.atomic.AtomicInteger +import java.util.concurrent.locks.ReentrantLock + +import scala.collection.mutable.ArrayBuffer +import scala.jdk.CollectionConverters._ + +import com.google.common.util.concurrent.Uninterruptibles +import org.apache.arrow.c.{ArrowArray, ArrowSchema, BaseStruct} +import org.apache.arrow.util.AutoCloseables +import org.apache.arrow.vector.VectorSchemaRoot + +import org.apache.spark.{SparkEnv, SparkException, TaskContext} +import org.apache.spark.api.python.ChainedPythonFunctions +import org.apache.spark.sql.catalyst.InternalRow +import org.apache.spark.sql.catalyst.expressions.{Attribute, Expression, JoinedRow, PythonUDF, UnsafeProjection, UnsafeRow} +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._ +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, unless all of them are UDF arguments that read back from Arrow unchanged. + * Each batch owns its Arrow buffers so Python can safely retain input arrays. + * + * The evaluator owns its queue, so that cleanup at task completion is coordinated with a + * consumer on another thread, such as a pipelined Python writer or a TRANSFORM feed thread. + */ +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, + fullValidation: Boolean, + metrics: Map[String, SQLMetric]) + extends EvalPythonEvaluatorFactory(childOutput, udfs, output) { + + private[python] def runtimeSession: InProcessPythonRuntime.InterpreterSession = + InProcessPythonRuntime.currentSession + + /** Unused: `evaluateJoined` always evaluates the UDFs. */ + override protected def evaluate( + funcs: Seq[(ChainedPythonFunctions, Long)], + argMetas: Array[Array[ArgumentMetadata]], + rows: Iterator[InternalRow], + inputSchema: StructType, + context: TaskContext): Iterator[InternalRow] = + throw SparkException.internalError("In-process UDFs are evaluated with their input rows") + + override protected def evaluateJoined( + funcs: Seq[(ChainedPythonFunctions, Long)], + argMetas: Array[Array[ArgumentMetadata]], + rows: Iterator[InternalRow], + inputs: Seq[Expression], + inputSchema: StructType, + context: TaskContext): Option[Iterator[InternalRow]] = { + import InProcessArrowEvalPythonEvaluatorFactory.{Buffered, ReadBack, readsBack} + val inputColumns = inputs.length == childOutput.length && inputs.zip(childOutput).forall { + case (a: Attribute, c) => a.exprId == c.exprId + case _ => false + } + // If all input columns are UDF arguments, they are written to Arrow regardless. Read them + // back from the exported input vectors instead of buffering every input row, if their + // values read back from Arrow exactly as written and as fast as an unsafe row copy. + val joinInput = if (inputColumns && inputSchema.forall(f => readsBack(f.dataType))) { + ReadBack + } else if (inputColumns) { + Buffered(None) + } else { + // Each projected row is written to Arrow before the next input row is pulled, so the + // arguments go into a reused buffer rather than being copied value by value. + val projection = UnsafeProjection.create(inputs, childOutput) + projection.initialize(context.partitionId()) + Buffered(Some(projection)) + } + Some(evaluateBatches(funcs, argMetas, rows, inputSchema, context, joinInput)) + } + + private[python] def evaluateBatches( + funcs: Seq[(ChainedPythonFunctions, Long)], + argMetas: Array[Array[ArgumentMetadata]], + rows: Iterator[InternalRow], + inputSchema: StructType, + context: TaskContext, + joinInput: InProcessArrowEvalPythonEvaluatorFactory.JoinInput): Iterator[InternalRow] = { + import InProcessArrowEvalPythonEvaluatorFactory.{Buffered, ReadBack} + ArrowUtils.failDuplicatedFieldNames(inputSchema) + val functions = funcs.map { case (chain, _) => + if (chain.funcs.size != 1) { + throw SparkException.internalError( + "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) + // Capture before consuming input: an old task must never join a later context's session. + val runtime = runtimeSession + // Rows are copied out of the queue and Arrow vectors before they are returned, so they + // remain valid after task completion releases those, on whichever thread consumes them. + val resultProj = UnsafeProjection.create(output, output) + // Spill files go into a directory of the queue's own, created on the first spill, so that + // task completion can delete them when it cannot close the queue. + @volatile var spillDir: File = null + val (queue, projection) = joinInput match { + case Buffered(projection) => + val localDir = new File(Utils.getLocalDir(SparkEnv.get.conf)) + val serializerManager = SparkEnv.get.serializerManager + // Only the consumer holding the iterator's lock adds and removes rows. + val queue = new HybridRowQueue(context.taskMemoryManager(), localDir, + childOutput.length, serializerManager, lockFree = true) { + override protected def createDiskQueue(): RowQueue = synchronized { + if (spillDir == null) { + spillDir = Files.createTempDirectory(localDir.toPath, "inprocess-udf-").toFile + } + DiskRowQueue(Files.createTempFile(spillDir.toPath, "buffer", "").toFile, + childOutput.length, serializerManager) + } + } + (queue, projection.orNull) + case ReadBack => (null, null) + } + val joined = new JoinedRow + val handles = functions.map(_ => UUID.randomUUID().toString) + var registered = false + var writer: ArrowWriter = null + val results = ArrayBuffer.empty[ArrowColumnVector] + var startedAt = 0L + + def closeBatch(): Unit = { + val resources = ArrayBuffer.empty[AutoCloseable] + resources ++= results + results.clear() + if (writer != null) { + resources += writer.root + writer = null + } + AutoCloseables.close(resources.asJava) + } + + val resources = new InProcessArrowEvalPythonEvaluatorFactory.IteratorResources( + // Closing the queue deletes its spill files, which leaves an empty directory. + releaseTaskMemory = () => if (queue != null) { + Utils.tryWithSafeFinally(queue.close()) { if (spillDir != null) spillDir.delete() } Review Comment: **[Low, cleanup] Correcting my round-11 suggestion: `delete()` leaves the directory when a file remains that `queue.close()` does not track.** I suggested `delete()` here because `queue.close()` deletes each spill file, but it deletes only the files of the disk queues it still tracks. If `createTempFile` at L158 succeeds and the `DiskRowQueue` constructor then fails (`RowQueue.scala` L130-131, e.g. with EMFILE), or if `HybridQueue.spill` fails midway (e.g. with ENOSPC) and drops the disk queues it already created, the file stays, `delete()` returns false, and the directory stays until the executor exits. Round 11's recursive deletion removed them. Suggestion: `Utils.deleteQuietly(spillDir)` on both paths. It costs the same for an empty directory, and since it accepts null, the null check at L189 can go too. ########## docs/sql-pyspark-inprocess-udf.md: ########## @@ -0,0 +1,727 @@ +--- +layout: global +title: In-Process Python UDFs +displayTitle: In-Process Python UDFs +license: | + Licensed to the Apache Software Foundation (ASF) under one or more + contributor license agreements. See the NOTICE file distributed with + this work for additional information regarding copyright ownership. + The ASF licenses this file to You under the Apache License, Version 2.0 + (the "License"); you may not use this file except in compliance with + the License. You may obtain a copy of the License at + + http://www.apache.org/licenses/LICENSE-2.0 + + Unless required by applicable law or agreed to in writing, software + distributed under the License is distributed on an "AS IS" BASIS, + WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + See the License for the specific language governing permissions and + limitations under the License. +--- + +* Table of contents +{:toc} + +## Runtime and result contract + +Each executor owns a dedicated interpreter thread. The plugin initializes the +interpreter on that thread, and task calls and shutdown are dispatched to the +same thread. The JVM is asked to allocate an 8 MiB stack for this thread; the +actual size is platform-dependent. Calls from concurrent tasks are queued on the +interpreter thread. +One task per executor is recommended for throughput, but is not a correctness requirement. +Application-level Python parallelism comes from multiple executor JVMs. +The plugin configures JEP's process-wide interpreter with hash seed `0`, matching +Spark's default Python worker seed. It must initialize before any other JEP user in +the JVM. The seed cannot change between SparkContexts in the same process; a custom +worker `PYTHONHASHSEED` does not override this embedded-runtime setting. + +Task cancellation cannot safely stop arbitrary native Python code. An interrupted +caller waits for the current invocation to finish before freeing the Arrow CDI +structures, then restores its interrupt status. A UDF that never returns can +therefore prevent its task from completing cancellation and block every subsequent +in-process UDF on that executor, including calls from other tasks, jobs, and sessions. +Recovery from a permanently hung invocation requires replacing the executor process. +Plugin shutdown stops accepting new calls and waits up to five seconds for the interpreter thread. If a call is +still running or a task still owns exported results, cleanup waits for that task to release +its CDI references; the memory remains live until cleanup completes or the process exits. Shutdown does not forcibly interrupt native +code. A new interpreter cannot start until the previous one has fully stopped. + +A scalar UDF must return a `pyarrow.Array` with exactly one element per input row. +The runtime checks the result type against the declared Spark type, including +nested fields, decimal scale, and timestamp unit. Timezone-aware timestamps are relabeled +to `spark.sql.session.timeZone` without changing their UTC instants or copying their buffers. +Timezone-naive and timezone-aware timestamps are not interchangeable. String and binary +offset widths, including nested values, are converted as needed to match +`spark.sql.execution.arrow.useLargeVarTypes`. Large, fixed-size and dictionary-encoded +representations of the declared types (`large_list`, `fixed_size_list`, `string_view`, +`binary_view`, `fixed_size_binary` and dictionary arrays) are cast to the declared type. +These conversions can allocate new buffers. Other value types must match exactly: use an +explicit PyArrow cast for numeric conversions. +Map `keys_sorted` metadata is normalized to Spark's declared map type. +Nested field nullability may differ if the actual values satisfy the declared nullability. Sliced results, including nested +child slices, are copied to remove offsets that Arrow Java's CDI importer cannot +read. Zero-length levels without a usable offsets buffer, which Arrow permits, are given +one. Compatible results retain zero-copy transfer. +Before exporting a result, the runtime performs full Arrow validation, including interior +offsets, because the JVM reads result buffers without bounds checks: a malformed result, +such as one built from raw buffers, could otherwise produce wrong values or crash the +executor. It does not validate UTF-8 in string results, because Spark strings may contain +invalid UTF-8 (for example, `CAST(X'FF' AS STRING)`). Worker-based Arrow UDFs do not +validate their results. To skip the full validation, set +`spark.sql.execution.pythonUDF.inProcess.fullValidation.enabled` to `false`; Arrow's +constant-time validation and the conversions above still apply. + +The API produces a regular `PythonUDF` expression with an in-process evaluation +type. Spark's existing `ArrowEvalPython` planning rules handle aggregation, +nested calls, nondeterminism, and filter/limit pushdown. A dedicated +`InProcessArrowEvalPythonExec` extends `EvalPythonExec`, reusing its argument extraction +and partition-evaluator path, while its evaluator buffers and joins input rows itself. +Ordinary Python UDFs continue to use Python workers. + +`maxRecordsPerBatch <= 0` means no row-count limit. The independent +`spark.sql.execution.arrow.maxBytesPerBatch` limit still applies when positive. +Only UDF arguments are converted to Arrow. Other columns stay in Spark rows, +buffered in a spillable queue until the results are joined back. When every input +column is a UDF argument and its type, other than an array or a map, reads back from Arrow +unchanged, the output reads those columns from the Arrow input vectors instead of buffering +the rows. +Duplicate nested field names in UDF arguments or declared results are rejected before +Arrow Java reads their buffers. + +Each batch uses fresh input buffers. A Python function may retain an input array; +later batches do not overwrite it. Retained arrays keep native memory alive, so +functions should release them when no longer needed. JVM input vectors and result +vectors are released on task completion, early termination and failure. The runtime retains +each exported result until the next invocation for that task or task cleanup, after the JVM +has released its references. The runtime drops its Python references on the interpreter +thread, so releasing JVM results does not trigger Python finalizers on Spark task threads. +Cleanup can remain queued behind another task's invocation. The rows can also be consumed +on another thread, such as a pipelined Python worker's writer. Task completion then stops +that consumer after the input row it is reading, and waits for it, but not for this +operator's Python: it releases the buffered rows at once, and the Arrow vectors when Python +returns. Reading one row can take longer when the input is another in-process UDF, whose +next row may need a batch of Python, or a blocked upstream operator; task completion waits +for at most one second, and then leaves the buffered rows to the executor and deletes Review Comment: **[Low, docs] Follow-up on R11-7: "at most one second" no longer always holds, and the leak sentence lost its condition.** Since this round, while the consumer buffers a row that waits for memory, the listener waits for it without a bound (L471-473 of the evaluator), e.g. until another task frees memory when the consumer is a TRANSFORM feed thread, which nothing interrupts at task completion. The class doc of `IteratorResources` ("waits for the lock only briefly") has the same gap. Also, the executor logs "Managed memory leak detected" only for a task that otherwise succeeds and still holds pages (`Executor.scala` L919). My round-11 wording ("For a successful task") kept that condition, but this sentence drops it. Finally, L107 is 131 characters, while the rest of the paragraph wraps at about 90. Suggestion: e.g. "..., unless the consumer is buffering a row that waits for memory" and "For a task that otherwise succeeds, the executor then logs ...", and rewrap the paragraph. ########## sql/core/src/main/scala/org/apache/spark/sql/execution/python/InProcessArrowEvalPythonEvaluatorFactory.scala: ########## @@ -0,0 +1,511 @@ +/* + * 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.io.File +import java.nio.file.Files +import java.util.UUID +import java.util.concurrent.TimeUnit +import java.util.concurrent.atomic.AtomicInteger +import java.util.concurrent.locks.ReentrantLock + +import scala.collection.mutable.ArrayBuffer +import scala.jdk.CollectionConverters._ + +import com.google.common.util.concurrent.Uninterruptibles +import org.apache.arrow.c.{ArrowArray, ArrowSchema, BaseStruct} +import org.apache.arrow.util.AutoCloseables +import org.apache.arrow.vector.VectorSchemaRoot + +import org.apache.spark.{SparkEnv, SparkException, TaskContext} +import org.apache.spark.api.python.ChainedPythonFunctions +import org.apache.spark.sql.catalyst.InternalRow +import org.apache.spark.sql.catalyst.expressions.{Attribute, Expression, JoinedRow, PythonUDF, UnsafeProjection, UnsafeRow} +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._ +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, unless all of them are UDF arguments that read back from Arrow unchanged. + * Each batch owns its Arrow buffers so Python can safely retain input arrays. + * + * The evaluator owns its queue, so that cleanup at task completion is coordinated with a + * consumer on another thread, such as a pipelined Python writer or a TRANSFORM feed thread. + */ +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, + fullValidation: Boolean, + metrics: Map[String, SQLMetric]) + extends EvalPythonEvaluatorFactory(childOutput, udfs, output) { + + private[python] def runtimeSession: InProcessPythonRuntime.InterpreterSession = + InProcessPythonRuntime.currentSession + + /** Unused: `evaluateJoined` always evaluates the UDFs. */ + override protected def evaluate( + funcs: Seq[(ChainedPythonFunctions, Long)], + argMetas: Array[Array[ArgumentMetadata]], + rows: Iterator[InternalRow], + inputSchema: StructType, + context: TaskContext): Iterator[InternalRow] = + throw SparkException.internalError("In-process UDFs are evaluated with their input rows") + + override protected def evaluateJoined( + funcs: Seq[(ChainedPythonFunctions, Long)], + argMetas: Array[Array[ArgumentMetadata]], + rows: Iterator[InternalRow], + inputs: Seq[Expression], + inputSchema: StructType, + context: TaskContext): Option[Iterator[InternalRow]] = { + import InProcessArrowEvalPythonEvaluatorFactory.{Buffered, ReadBack, readsBack} + val inputColumns = inputs.length == childOutput.length && inputs.zip(childOutput).forall { + case (a: Attribute, c) => a.exprId == c.exprId + case _ => false + } + // If all input columns are UDF arguments, they are written to Arrow regardless. Read them + // back from the exported input vectors instead of buffering every input row, if their + // values read back from Arrow exactly as written and as fast as an unsafe row copy. + val joinInput = if (inputColumns && inputSchema.forall(f => readsBack(f.dataType))) { + ReadBack + } else if (inputColumns) { + Buffered(None) + } else { + // Each projected row is written to Arrow before the next input row is pulled, so the + // arguments go into a reused buffer rather than being copied value by value. + val projection = UnsafeProjection.create(inputs, childOutput) + projection.initialize(context.partitionId()) + Buffered(Some(projection)) + } + Some(evaluateBatches(funcs, argMetas, rows, inputSchema, context, joinInput)) + } + + private[python] def evaluateBatches( + funcs: Seq[(ChainedPythonFunctions, Long)], + argMetas: Array[Array[ArgumentMetadata]], + rows: Iterator[InternalRow], + inputSchema: StructType, + context: TaskContext, + joinInput: InProcessArrowEvalPythonEvaluatorFactory.JoinInput): Iterator[InternalRow] = { + import InProcessArrowEvalPythonEvaluatorFactory.{Buffered, ReadBack} + ArrowUtils.failDuplicatedFieldNames(inputSchema) + val functions = funcs.map { case (chain, _) => + if (chain.funcs.size != 1) { + throw SparkException.internalError( + "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) + // Capture before consuming input: an old task must never join a later context's session. + val runtime = runtimeSession + // Rows are copied out of the queue and Arrow vectors before they are returned, so they + // remain valid after task completion releases those, on whichever thread consumes them. + val resultProj = UnsafeProjection.create(output, output) + // Spill files go into a directory of the queue's own, created on the first spill, so that + // task completion can delete them when it cannot close the queue. + @volatile var spillDir: File = null + val (queue, projection) = joinInput match { + case Buffered(projection) => + val localDir = new File(Utils.getLocalDir(SparkEnv.get.conf)) + val serializerManager = SparkEnv.get.serializerManager + // Only the consumer holding the iterator's lock adds and removes rows. + val queue = new HybridRowQueue(context.taskMemoryManager(), localDir, + childOutput.length, serializerManager, lockFree = true) { + override protected def createDiskQueue(): RowQueue = synchronized { + if (spillDir == null) { + spillDir = Files.createTempDirectory(localDir.toPath, "inprocess-udf-").toFile + } + DiskRowQueue(Files.createTempFile(spillDir.toPath, "buffer", "").toFile, + childOutput.length, serializerManager) + } + } + (queue, projection.orNull) + case ReadBack => (null, null) + } + val joined = new JoinedRow + val handles = functions.map(_ => UUID.randomUUID().toString) + var registered = false + var writer: ArrowWriter = null + val results = ArrayBuffer.empty[ArrowColumnVector] + var startedAt = 0L + + def closeBatch(): Unit = { + val resources = ArrayBuffer.empty[AutoCloseable] + resources ++= results + results.clear() + if (writer != null) { + resources += writer.root + writer = null + } + AutoCloseables.close(resources.asJava) + } + + val resources = new InProcessArrowEvalPythonEvaluatorFactory.IteratorResources( + // Closing the queue deletes its spill files, which leaves an empty directory. + releaseTaskMemory = () => if (queue != null) { + Utils.tryWithSafeFinally(queue.close()) { if (spillDir != null) spillDir.delete() } + }, + // Neither starts a process nor throws, also on an interrupted thread. + abandonTaskMemory = () => if (spillDir != null) Utils.deleteQuietly(spillDir), + releaseOthers = () => { + if (startedAt != 0L) { + metrics("pythonTotalTime") += (System.nanoTime() - startedAt) / 1000000 + } + Utils.tryWithSafeFinally { + closeBatch() + } { + if (registered) runtime.release(handles) + } + }) + + context.addTaskCompletionListener[Unit](_ => resources.close()) + + new Iterator[InternalRow] { + private var batchIter: Iterator[InternalRow] = Iterator.empty + + private def endOfInput: Nothing = + throw new NoSuchElementException("End of in-process UDF input") + + // Releases the resources on failure without replacing its exception. + private def fail(t: Throwable): Nothing = + Utils.tryWithSafeFinally { throw t } { resources.close() } + + // Called with the lock held. + private def hasNextLocked: Boolean = { + if (startedAt == 0L) startedAt = System.nanoTime() + checkCancellation() + val available = batchIter.hasNext || rows.hasNext + if (!available) resources.close() + available + } + + // Each call takes the lock without allocating a closure per row. Within a batch, only + // the consumer advances `batchIter`, whose `hasNext` compares row indexes. + override def hasNext: Boolean = { + if (batchIter.hasNext && !resources.isClosed) return true + if (!resources.enter()) return false + val available = try { + try hasNextLocked catch { case t: Throwable => fail(t) } + } finally { + resources.exit() + } + // Task completion may have closed the iterator while the input was being read. + available && !resources.isClosed + } + + override def next(): InternalRow = { + if (!resources.enter()) endOfInput + try { + try { + if (!hasNextLocked) endOfInput + if (!batchIter.hasNext) nextBatch() + val result = batchIter.next() + resultProj(if (queue != null) joined(queue.remove(), result) else result) + } catch { + case t: Throwable => fail(t) + } + } finally { + resources.exit() + } + } + + // Runs Python without the lock, unless task completion happened before or meanwhile. + private def python[T](body: => T): T = { + if (resources.isClosed) endOfInput + val result = resources.withoutLock(body) + if (resources.isClosed) endOfInput + result + } + + /** + * Writes the next input row to the batch, returning false at the end of input. Task + * completion can happen while the input is read; then the row is not written. + */ + private def pullRow(): Boolean = rows.hasNext && !resources.isClosed && { + val row = rows.next() + if (queue != null) { + // Adding can wait for memory beyond the listener's wait. Announce it before the + // check, so that either the add is skipped or the listener waits for it. + resources.usingTaskMemory = true + try { + if (resources.isClosed) endOfInput + queue.add(row.asInstanceOf[UnsafeRow]) + } finally { + resources.usingTaskMemory = false + } + } else if (resources.isClosed) { + endOfInput + } + writer.write(if (projection != null) projection(row) else row) + true + } + + // Called with the lock held. + private def nextBatch(): Unit = { + closeBatch() + val root = VectorSchemaRoot.create(arrowSchema, ArrowUtils.rootAllocator) + writer = try { + ArrowWriter.create(root) + } catch { + case t: Throwable => Utils.tryWithSafeFinally { throw t } { root.close() } + } + // Task completion stops the fill within a row, and Python never sees a partial batch. + var count = 0 + while (!resources.isClosed && (batchSize <= 0 || count < batchSize) && + (count == 0 || maxBytes <= 0 || writer.sizeInBytes() < maxBytes) && { + checkCancellation() + pullRow() + }) { + count += 1 + } + if (resources.isClosed) endOfInput + if (!registered) { + // Mark before registering so failure after any registration still cleans up. + registered = true + functions.indices.foreach { i => + val func = functions(i) + initTime.add(python(runtime.register(handles(i), func.command.toArray, + expectedFields(i), func.pythonVer, hideTraceback, simplifiedTraceback, + tracebackWithLocals, fullValidation))) + } + } + writer.finish() + metrics("pythonDataSent") += writer.sizeInBytes() + + handles.indices.foreach { udfIndex => + val handle = handles(udfIndex) + val ordinals = inputOrdinals(udfIndex) + checkCancellation() + // Register each acquired resource immediately, including partially exported + // inputs and results of earlier UDFs if a later UDF throws. + val structs = ArrayBuffer.empty[AutoCloseable] + def track[S <: BaseStruct](struct: S): S = { + val closer: AutoCloseable = () => InProcessArrowBridge.closeStruct(struct) + structs += closer + struct + } + def array(): ArrowArray = track(ArrowArray.allocateNew(ArrowUtils.rootAllocator)) + def schema(): ArrowSchema = track(ArrowSchema.allocateNew(ArrowUtils.rootAllocator)) + Utils.tryWithSafeFinally { + val inArrays = ordinals.map(_ => array()) + val inSchemas = ordinals.map(_ => schema()) + val outArray = array() + val outSchema = schema() + ordinals.indices.foreach { i => + InProcessArrowBridge.exportColumn( + writer.root.getVector(ordinals(i)), inArrays(i), inSchemas(i)) + } + processingTime.add(python(runtime.invoke( + handle, + inArrays.map(_.memoryAddress()).toArray, + inSchemas.map(_.memoryAddress()).toArray, + outArray.memoryAddress(), outSchema.memoryAddress(), + count, argMetas(udfIndex).map(_.name.getOrElse(""))))) + results += InProcessArrowBridge.cdiToColumn( + outArray, outSchema, Some(expectedFields(udfIndex))) + metrics("pythonDataReceived") += results.last.getValueVector.getBufferSize + } { + AutoCloseables.close(structs.asJava) + } + } + + metrics("pythonNumRowsReceived") += count + // Input vectors are closed with the writer's root, not with the results. + val inputs = if (joinInput == ReadBack) { + writer.root.getFieldVectors.asScala.map(new ArrowColumnVector(_)) + } else { + Nil + } + val columns = (inputs ++ results).toArray[ColumnVector] + batchIter = new ColumnarBatch(columns, count).rowIterator().asScala + } + } + } +} + +private[python] object InProcessArrowEvalPythonEvaluatorFactory { + /** How the evaluator joins input rows with their results. */ + sealed trait JoinInput + /** Read the input columns back from the exported Arrow input vectors. */ + case object ReadBack extends JoinInput + /** Buffer the input rows, writing their arguments, projected if needed, to Arrow. */ + case class Buffered(projection: Option[UnsafeProjection]) extends JoinInput + + /** + * Whether `ArrowColumnVector` returns exactly the values `ArrowWriter` wrote for this type, + * and an unsafe projection copies them about as fast as an unsafe row. Types with derived + * Arrow representations, such as intervals, nanosecond timestamps, TIME, Variant, geospatial + * types and UDTs, keep the original rows instead. So do arrays and maps, which a projection + * copies element by element out of Arrow, but with a single copy out of an unsafe row. + */ + def readsBack(dataType: DataType): Boolean = dataType match { + case NullType | BooleanType | ByteType | ShortType | IntegerType | LongType | + FloatType | DoubleType | BinaryType | DateType | TimestampType | TimestampNTZType => true + case _: DecimalType => true + case _: StringType => true + case StructType(fields) => fields.forall(f => readsBack(f.dataType)) + case _ => false + } + + /** + * Coordinates cleanup at task completion with the consumer of the evaluator's iterator. The + * consumer can run on another thread, e.g. a pipelined Python writer or a TRANSFORM feed + * thread, and the completion listener cannot tell, since a lazily computing parent (such as + * `coalesce`) can create the iterator on that thread too. + * + * The consumer holds the lock while it reads input, the row queue or Arrow vectors, and + * releases it only while this evaluator's Python runs. The listener (`close`) first requests + * closing, which the consumer checks after each input row, so the listener waits for at most + * one row before it releases task memory (the row queue), ahead of the executor. It releases + * the other resources (Arrow vectors and Python handles) too, unless Python is running; then + * the consumer releases them when Python returns. + * + * Reading one row can take long: the input can be another in-process evaluator, whose next + * row may need a batch of Python, or an upstream operator that only a later listener + * unblocks. So the listener waits for the lock only briefly. Then it leaves the task memory + * to the executor, deleting what lives outside it, and the consumer releases the other + * resources once its row returns, without touching the task memory again. + */ + class IteratorResources( + releaseTaskMemory: () => Unit, + abandonTaskMemory: () => Unit, + releaseOthers: () => Unit, + lockWaitMillis: Long = 1000L) { + private val lock = new ReentrantLock() + @volatile private var closeRequested = false + // Task memory is released by whichever of the consumer and the listener gets here first, + // or abandoned to the executor if the listener gives up on the lock. + private val taskMemory = new AtomicInteger(TaskMemoryHeld) + // Guarded by the lock. + private var inPython = false + private var othersReleased = false + + /** + * Set by the consumer while it may allocate task memory, e.g. for a queue page, which can + * wait for other tasks' memory. The listener then waits for the lock instead of leaving + * the task memory to the executor, since that wait never depends on a later listener. + */ + @volatile var usingTaskMemory = false Review Comment: **[Low, cleanup] The `usingTaskMemory` handshake is split across two classes, through a public `var`.** The consumer half (set, check, add, clear at L269-275) lives in the evaluator and the listener half in `close()` (L471-473), so only the comment at L267-268 keeps the ordering that makes it correct. Also, the `releaseAll()` in the new branch of `close()` is always a no-op: once the listener reads the flag as set, the consumer sees the close at its next check and releases everything before it unlocks, so the branch behaves like the final wait. Suggestion: a non-allocating pair next to `enter()`/`exit()` with a private flag, e.g. `def enterTaskMemory(): Boolean = { usingTaskMemory = true; !closeRequested }` and `def exitTaskMemory(): Unit = usingTaskMemory = false`, used as `try { if (!resources.enterTaskMemory()) endOfInput; queue.add(...) } finally resources.exitTaskMemory()`. Then `close()` could use `else if (!usingTaskMemory && taskMemory.compareAndSet(TaskMemoryHeld, TaskMemoryAbandoned))`, so that a set flag falls through to the final `lock.lock(); lock.unlock()`. -- 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]
