andygrove commented on code in PR #5051:
URL: https://github.com/apache/datafusion-comet/pull/5051#discussion_r3867744580


##########
spark/src/main/scala/org/apache/spark/sql/comet/execution/arrow/ArrowCachedBatchSerializer.scala:
##########
@@ -0,0 +1,512 @@
+/*
+ * 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.comet.execution.arrow
+
+import scala.collection.JavaConverters._
+
+import org.apache.spark.TaskContext
+import org.apache.spark.rdd.RDD
+import org.apache.spark.sql.catalyst.InternalRow
+import org.apache.spark.sql.catalyst.expressions.{Attribute, Expression, 
GenericInternalRow, IsNotNull, IsNull, UnsafeProjection}
+import org.apache.spark.sql.columnar.{CachedBatch, SimpleMetricsCachedBatch, 
SimpleMetricsCachedBatchSerializer}
+import org.apache.spark.sql.comet.util.Utils
+import org.apache.spark.sql.execution.columnar.DefaultCachedBatchSerializer
+import org.apache.spark.sql.internal.SQLConf
+import org.apache.spark.sql.types._
+import org.apache.spark.sql.vectorized.{ColumnarBatch, ColumnVector}
+import org.apache.spark.storage.StorageLevel
+import org.apache.spark.unsafe.types.{ByteArray, UTF8String}
+import org.apache.spark.util.io.ChunkedByteBuffer
+
+import org.apache.comet.CometArrowAllocator
+
+/**
+ * Cached batch format used when Comet writes Spark in-memory cache data.
+ *
+ * `columns` holds one compressed Arrow stream per cached column, in 
cache-schema order, produced
+ * by `Utils.serializeBatchColumns`. Storing columns separately is what lets a 
scan decode only
+ * the ones it projected; a single stream covering the whole batch would have 
to be inflated in
+ * full before any projection could be applied. The cache manager still owns 
storage and eviction;
+ * this class only changes the cached payload.
+ */
+private case class CometCachedBatch(
+    override val numRows: Int,
+    override val sizeInBytes: Long,
+    override val stats: InternalRow,
+    columns: Array[ChunkedByteBuffer])
+    extends SimpleMetricsCachedBatch
+
+/**
+ * Cache serializer that stores Comet-compatible Arrow batches in Spark's 
in-memory cache.
+ *
+ * The cached payload format is decided by the schema alone. A relation whose 
schema Comet's Arrow
+ * writer supports is stored as `CometCachedBatch`, and every other relation 
is delegated in full
+ * to Spark's `DefaultCachedBatchSerializer`. The format deliberately does not 
depend on any
+ * runtime config: `spark.sql.cache.serializer` is a static conf, so 
installing this serializer is
+ * already a per-application decision, and a relation whose format could flip 
mid-session cannot
+ * be read back reliably. `spark.comet.exec.inMemoryCache.enabled` still 
governs whether a scan
+ * over the cache runs natively, and its value at startup is what makes 
`CometDriverPlugin`
+ * install this serializer in the first place.
+ *
+ * Reads of `CometCachedBatch` keep working when the native scan is disabled, 
because Spark then
+ * reads the same cached data through the SparkToColumnar fallback path.
+ */
+class ArrowCachedBatchSerializer extends SimpleMetricsCachedBatchSerializer {
+
+  import ArrowCachedBatchSerializer.supportsSchema
+
+  private val fallback = new DefaultCachedBatchSerializer()
+
+  // Bounds and null counts per column, gathered before the batch is 
serialized: serializing
+  // clears the batch's vectors, and the per-column byte sizes that complete 
the statistics row
+  // are only known afterwards. See statsRow.
+  private def gatherColumnStats(
+      batch: ColumnarBatch,
+      attrs: Seq[Attribute]): (Array[Any], Array[Any], Array[Int]) = {
+    val numCols = attrs.length
+    val lower = new Array[Any](numCols)
+    val upper = new Array[Any](numCols)
+    val nulls = Array.fill[Int](numCols)(0)
+    val numRows = batch.numRows()
+
+    var c = 0
+    while (c < numCols) {
+      val dt = attrs(c).dataType
+      val col = batch.column(c)
+      var r = 0
+      while (r < numRows) {
+        if (col.isNullAt(r)) {
+          nulls(c) += 1
+        } else if (tracksBounds(dt)) {
+          val value = readValue(col, dt, r)
+          if (lower(c) == null || compare(dt, value, lower(c)) < 0) {
+            lower(c) = value
+          }
+          if (upper(c) == null || compare(dt, value, upper(c)) > 0) {
+            upper(c) = value
+          }
+        }
+        r += 1
+      }
+      c += 1
+    }
+
+    (lower, upper, nulls)
+  }
+
+  // Build the statistics row expected by SimpleMetricsCachedBatchSerializer.
+  // For each cached column Spark expects five values in this order:
+  // lower bound, upper bound, null count, row count, and size in bytes.
+  private def statsRow(
+      lower: Array[Any],
+      upper: Array[Any],
+      nulls: Array[Int],
+      numRows: Int,
+      columnSizes: Array[Long]): InternalRow = {
+    val numCols = lower.length
+    val values = new Array[Any](numCols * 5)
+    var c = 0
+    while (c < numCols) {
+      val base = c * 5
+      values(base) = lower(c)
+      values(base + 1) = upper(c)
+      values(base + 2) = nulls(c)
+      values(base + 3) = numRows
+      // Each column is its own compressed stream, so its size is known 
exactly. Cache pruning
+      // uses bounds/null-count/row-count rather than this field, but Spark 
reserves it and
+      // reports it, so record the real value.
+      values(base + 4) = columnSizes(c)
+      c += 1
+    }
+
+    new GenericInternalRow(values)
+  }
+
+  // Spark can prune cache batches only for types whose bounds can be compared.
+  // Other types still report null count and row count but leave bounds as 
null.
+  private def tracksBounds(dt: DataType): Boolean = dt match {
+    case BooleanType | ByteType | ShortType | IntegerType | LongType | 
FloatType | DoubleType |
+        _: DecimalType | StringType | DateType | TimestampType | 
TimestampNTZType =>
+      true
+    case _ => false
+  }
+
+  // Read a non-null value from a ColumnVector using Spark's internal value 
type
+  // for the corresponding DataType.
+  private def readValue(col: ColumnVector, dt: DataType, rowId: Int): Any = dt 
match {
+    case BooleanType => col.getBoolean(rowId)
+    case ByteType => col.getByte(rowId)
+    case ShortType => col.getShort(rowId)
+    case IntegerType | DateType => col.getInt(rowId)
+    case LongType | TimestampType | TimestampNTZType => col.getLong(rowId)
+    case FloatType => col.getFloat(rowId)
+    case DoubleType => col.getDouble(rowId)
+    case d: DecimalType => col.getDecimal(rowId, d.precision, d.scale)
+    case StringType => col.getUTF8String(rowId).copy()
+    case _ => null
+  }
+
+  // Compare values using the same physical representation used in the stats 
row.
+  private def compare(dt: DataType, left: Any, right: Any): Int = dt match {
+    case BooleanType =>
+      java.lang.Boolean.compare(left.asInstanceOf[Boolean], 
right.asInstanceOf[Boolean])
+    case ByteType =>
+      java.lang.Byte.compare(left.asInstanceOf[Byte], right.asInstanceOf[Byte])
+    case ShortType =>
+      java.lang.Short.compare(left.asInstanceOf[Short], 
right.asInstanceOf[Short])
+    case IntegerType | DateType =>
+      java.lang.Integer.compare(left.asInstanceOf[Int], 
right.asInstanceOf[Int])
+    case LongType | TimestampType | TimestampNTZType =>
+      java.lang.Long.compare(left.asInstanceOf[Long], right.asInstanceOf[Long])
+    case FloatType =>
+      java.lang.Float.compare(left.asInstanceOf[Float], 
right.asInstanceOf[Float])
+    case DoubleType =>
+      java.lang.Double.compare(left.asInstanceOf[Double], 
right.asInstanceOf[Double])
+    case _: DecimalType =>
+      left.asInstanceOf[Decimal].compare(right.asInstanceOf[Decimal])
+    case StringType =>
+      ByteArray.compareBinary(
+        left.asInstanceOf[UTF8String].getBytes,
+        right.asInstanceOf[UTF8String].getBytes)
+    case other =>
+      throw new IllegalStateException(s"compare called for unsupported type 
$other")
+  }
+
+  // Compute Spark-compatible cache stats before serializing each batch to 
Arrow.
+  // The stats are stored beside the Arrow bytes so Spark's cache filter can 
prune
+  // CometCachedBatch without decoding the batch first.
+  //
+  // A columnar input batch is not guaranteed to be Arrow-backed; see 
supportsColumnarInput for
+  // why. Batches that are not get copied into Arrow first, since 
Utils.serializeBatches only
+  // writes CometVector columns.
+  private def encodeBatches(
+      batches: Iterator[ColumnarBatch],
+      attrs: Seq[Attribute]): Iterator[CachedBatch] = {
+    val arrowSchema =
+      Utils.toArrowSchema(Utils.fromAttributes(attrs), 
CometArrowStream.NATIVE_TIMEZONE)
+
+    batches.map { batch =>
+      // Bounds and null counts are read from the input batch, which 
serializing then clears, so
+      // they have to be gathered first. The row is only assembled once the 
per-column sizes are
+      // known.
+      val (lower, upper, nulls) = gatherColumnStats(batch, attrs)
+      val numRows = batch.numRows()
+
+      val columns = if (Utils.isArrowBacked(batch)) {
+        Utils.serializeBatchColumns(batch)
+      } else {
+        val arrowBatch =
+          CometArrowConverters.columnarBatchToArrowBatch(batch, arrowSchema, 
CometArrowAllocator)
+        try Utils.serializeBatchColumns(arrowBatch)
+        finally arrowBatch.close()

Review Comment:
   Filed the durable part of this as #5488.
   
   To be precise about what does and does not outlive this PR, since I called 
it pre-existing above and that was ambiguous. `CometInMemoryTableScanExec` and 
`ArrowCachedBatchSerializer` are both new files in #5051, so the two findings 
in the scan node cannot survive it being dropped. This one is different: 
`getFieldVector` and both `Utils.serializeBatches` consumers, `getByteArrayRdd` 
in `operators.scala` and `CometBroadcastExchangeExec`, are on `main` today.
   
   And the fix here does not close it. Making `isArrowBacked` reject those 
vectors only reroutes the cache write path to conversion; the broadcast and 
collect paths have no conversion fallback and still throw. 
`isSupportedFieldVector` is a useful building block for a real fix, but it goes 
away with this PR too. #5488 has the producer trail through 
`CometMapInBatchExec`, and notes that neither of us ran the mapInArrow path end 
to end.



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