viirya commented on code in PR #58978:
URL: https://github.com/apache/spark/pull/58978#discussion_r4233045660


##########
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:
   Filed as SPARK-60071, with the fix in #59282: `HybridRowQueue` itself now 
uses identity equality. Whichever of the two PRs merges second drops the 
override here.



-- 
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]

Reply via email to