sunchao commented on code in PR #5051:
URL: https://github.com/apache/datafusion-comet/pull/5051#discussion_r3867821448
##########
spark/src/main/scala/org/apache/spark/sql/comet/util/Utils.scala:
##########
@@ -398,33 +430,69 @@ object Utils extends CometTypeShim with Logging {
}
}
+ /**
+ * Whether every column in `batch` is an Arrow-backed `CometVector`, so
[[getBatchFieldVectors]]
+ * can hand out its vectors directly. Callers that may receive batches from
a plan they did not
+ * build (e.g. Comet's cache serializer, which Spark hands the cached plan's
columnar output)
+ * use this to convert foreign vectors to Arrow instead of tripping the
exception below.
+ *
+ * Stricter than what [[getBatchFieldVectors]] accepts: a
`ConstantColumnVector` is rejected
+ * here even though that method materializes one, so such a batch takes the
conversion path
+ * rather than being materialized column by column.
+ */
+ def isArrowBacked(batch: ColumnarBatch): Boolean =
+ (0 until batch.numCols()).forall { i =>
+ batch.column(i) match {
+ // Not every CometVector can be handed to getFieldVector: a
CometPlainVector can wrap a
+ // LargeVarCharVector or LargeVarBinaryVector (an accelerated
mapInArrow returning
+ // pa.large_string(), for instance), which it rejects. Answering true
for those would
+ // send a batch down the direct write path that then fails, so check
the vector itself
+ // and let the caller convert instead.
+ case v: CometVector => isSupportedFieldVector(v.getValueVector)
+ case _ => false
+ }
+ }
+
def getBatchFieldVectors(
batch: ColumnarBatch): (Seq[FieldVector], Option[DictionaryProvider]) = {
- var provider: Option[DictionaryProvider] = None
+ val columns = getBatchFieldVectorsWithProviders(batch)
+ (columns.map(_._1), columns.flatMap(_._2).headOption)
Review Comment:
**[P2] Preserve each column's dictionary when broadcasting cached batches**
The per-column cache write fix works, but this whole-batch helper still
discards every provider except the first. A native broadcast can consume
`CometInMemoryTableScanExec` directly after the wrapper is removed, so
`CometBroadcastExchangeExec.getByteArrayRdd` passes a batch with independent
providers to `Utils.serializeBatches`.
I reproduced this with Comet caching, JVM shuffle and AQE disabled: cache a
repartitioned DataFrame with two low-cardinality string columns, then run
`SELECT /*+ BROADCAST(c) */ c.a, c.b FROM range(1) r JOIN dictionary_cache c ON
true`. It fails with `Could not find dictionary with ID 1`. Disabling only
`spark.comet.exec.inMemoryCache.enabled` lets the same query over the same
cache return all 10,000 rows. Please normalize or safely combine the
dictionaries for whole-batch serialization as well as the per-column writer.
##########
spark/src/main/scala/org/apache/spark/sql/comet/execution/arrow/ArrowCachedBatchSerializer.scala:
##########
@@ -0,0 +1,539 @@
+/*
+ * 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 scala.util.control.NonFatal
+
+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()
+ }
+
+ val columnSizes = columns.map(_.size)
+ CometCachedBatch(
+ numRows = numRows,
+ sizeInBytes = columnSizes.sum,
+ stats = statsRow(lower, upper, nulls, numRows, columnSizes),
+ columns = columns)
+ }
+ }
+
+ // Resolve requested columns by exprId, not by name, because aliases may
reuse names.
+ //
+ // An empty selection stays empty rather than expanding to every column.
Spark asks for no
+ // columns when the query only needs the row count (SELECT count(*)), and
since projection now
+ // decides what gets decoded, expanding it would turn the cheapest possible
read into the most
+ // expensive one.
+ private def selectedIndices(
+ cacheAttributes: Seq[Attribute],
+ selectedAttributes: Seq[Attribute]): Array[Int] = {
+ val byExprId = cacheAttributes.zipWithIndex.map { case (attr, idx) =>
+ attr.exprId -> idx
+ }.toMap
+
+ selectedAttributes.map { attr =>
+ byExprId.getOrElse(
+ attr.exprId,
+ throw new IllegalStateException(
+ s"Could not resolve selected attribute ${attr.name} from cache
attributes"))
+ }.toArray
+ }
+
+ // Spark's SimpleMetricsCachedBatchSerializer prunes a batch when the
generated partition filter
+ // does not evaluate to true against the stats row. Bounds are only computed
for the types
+ // tracksBounds accepts, and for every other column the lower and upper
bounds stay null, which
+ // makes a comparison against them evaluate to null and therefore prune the
batch. That would
+ // silently drop rows, so predicates over columns without bounds are not
pushed down at all.
+ // Null counts and row counts are recorded for every column, so IsNull and
IsNotNull stay safe.
+ override def buildFilter(
+ predicates: Seq[Expression],
+ cachedAttributes: Seq[Attribute]): (Int, Iterator[CachedBatch]) =>
Iterator[CachedBatch] = {
+ val prunable = cachedAttributes.collect {
+ case a if tracksBounds(a.dataType) => a.exprId
+ }.toSet
+
+ val prunablePredicates = predicates.filter {
+ case _: IsNull | _: IsNotNull => true
+ case p => p.references.forall(a => prunable.contains(a.exprId))
+ }
+
+ super.buildFilter(prunablePredicates, cachedAttributes)
+ }
+
+ // Comet's Arrow writer only handles the types listed in supportsSchema.
Reporting false here
+ // sends the relation down the row path, where it is delegated to Spark's
default serializer,
+ // instead of failing at cache materialization inside Utils.serializeBatches.
+ //
+ // This answer is schema-only, because attributes are all Spark gives us; it
says nothing about
+ // the vectors. Returning true also makes InMemoryRelation strip the
ColumnarToRow above the
+ // cached plan, so convertColumnarBatchToCachedBatch then receives whatever
that plan produces:
+ // a Comet scan's CometVectors, but equally Spark's vectorized Parquet/ORC
reader or a
+ // connector's own vectors. encodeBatches converts the non-Arrow ones; that
conversion is load
+ // bearing, not defensive.
+ override def supportsColumnarInput(schema: Seq[Attribute]): Boolean =
supportsSchema(schema)
+
+ // A relation Comet stores is always readable as columnar Arrow. Anything
else holds
+ // DefaultCachedBatch, so defer to Spark, which only claims columnar output
for the primitive
+ // types its ColumnAccessor.decompress path can actually decode.
+ override def supportsColumnarOutput(schema: StructType): Boolean = {
+ if (schema.fields.forall(f =>
ArrowCachedBatchSerializer.supportsType(f.dataType))) {
+ true
+ } else {
+ fallback.supportsColumnarOutput(schema)
+ }
+ }
+
+ // Columnar Comet output is stored as compressed Arrow stream bytes. Spark
only calls this when
+ // supportsColumnarInput returned true, so the schema is known to be
Comet-writable here.
+ override def convertColumnarBatchToCachedBatch(
+ input: RDD[ColumnarBatch],
+ schema: Seq[Attribute],
+ storageLevel: StorageLevel,
+ conf: SQLConf): RDD[CachedBatch] = {
+
+ input.mapPartitions { batches =>
+ encodeBatches(batches, schema)
+ }
+ }
+
+ override def convertCachedBatchToColumnarBatch(
+ input: RDD[CachedBatch],
+ cacheAttributes: Seq[Attribute],
+ selectedAttributes: Seq[Attribute],
+ conf: SQLConf): RDD[ColumnarBatch] = {
+ if (!supportsSchema(cacheAttributes)) {
+ return fallback.convertCachedBatchToColumnarBatch(
+ input,
+ cacheAttributes,
+ selectedAttributes,
+ conf)
+ }
+
+ val indices = selectedIndices(cacheAttributes, selectedAttributes)
+
+ input.mapPartitions { it =>
+ // A ColumnReaders closes its readers (releasing the vectors they are
holding) only when the
+ // batch it produced has been consumed. A consumer that stops early --
LIMIT, take(), or a
+ // cancelled task -- leaves the readers for the batch in flight open, so
close them on task
+ // completion. Spark's own ArrowCachedBatchSerializer registers a
listener for the same
+ // reason.
+ //
+ // flatMap consumes each inner iterator fully before building the next,
so at most one batch
+ // is open at a time and tracking the current one is enough. close() is
idempotent, so
+ // closing one that already released itself is a no-op.
+ @volatile var current: ColumnReaders = null
+ Option(TaskContext.get()).foreach { tc =>
+ tc.addTaskCompletionListener[Unit] { _ =>
+ val readers = current
+ current = null
+ if (readers != null) {
+ readers.close()
+ }
+ }
+ }
+
+ it.flatMap {
+ case cb: CometCachedBatch =>
+ if (indices.isEmpty) {
+ // Nothing to decode: the row count is the whole answer, and it is
already here.
+ Iterator.single(new ColumnarBatch(Array.empty[ColumnVector],
cb.numRows))
+ } else {
+ val readers = new ColumnReaders(indices.map(i => cb.columns(i)),
cb.numRows)
+ current = readers
+ readers.batches
+ }
+
+ case other =>
+ throw new IllegalStateException(
+ s"Unsupported cached batch type ${other.getClass.getName}")
+ }
+ }
+ }
+
+ // Decodes one selected column stream apiece and stitches the results back
into a single batch.
+ //
+ // Each stream is self-contained, so the columns a scan did not select are
never inflated. The
+ // decoded vectors stay owned by their readers: closing them releases the
batch, which is why
+ // this yields a single-element iterator that closes on exhaustion, matching
what
+ // ArrowReaderIterator did when the payload was one stream.
+ private class ColumnReaders(buffers: Array[ChunkedByteBuffer], numRows: Int)
{
+ // decodeBatches opens a reader and eagerly decodes its first batch, so it
allocates. If a
+ // later column throws, the readers already opened here are unreachable:
the task-completion
+ // listener cannot release them because `current` is only assigned once
this constructor
+ // returns, so they would leak off-heap for the life of the executor.
+ private val readers: Array[Iterator[ColumnarBatch]] = {
+ val opened = new Array[Iterator[ColumnarBatch]](buffers.length)
+ var i = 0
+ try {
+ while (i < buffers.length) {
+ opened(i) = Utils.decodeBatches(buffers(i), "CometCache")
Review Comment:
**[P2] Close the reader whose own initialization fails**
The new guard correctly closes previously returned readers, but
`decodeBatches` can allocate a dictionary and then throw while loading that
same column's first record batch. In that case it never returns a reader to
`opened(i)`, and neither this catch nor the task-completion listener can close
it.
With one valid dictionary-encoded cached string column, I limited the
allocator to allow its dictionary message but reject the following indices
buffer. After three completed failed tasks, retained Arrow allocations grew
from zero to 64, 128 and 192 bytes. The original two-plain-column reproduction
now returns to zero. Please make reader initialization exception-safe inside
`ArrowReaderIterator`/the reader factory too, so a failed first batch releases
that reader's own allocations.
##########
spark/src/main/scala/org/apache/spark/sql/comet/CometInMemoryTableScanExec.scala:
##########
@@ -0,0 +1,142 @@
+/*
+ * 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
+
+import scala.collection.JavaConverters._
+
+import org.apache.spark.rdd.RDD
+import org.apache.spark.sql.catalyst.expressions.Attribute
+import org.apache.spark.sql.columnar.CachedBatchSerializer
+import org.apache.spark.sql.execution.LeafExecNode
+import org.apache.spark.sql.execution.columnar.{CachedRDDBuilder,
InMemoryTableScanExec}
+import org.apache.spark.sql.execution.metric.{SQLMetric, SQLMetrics}
+import org.apache.spark.sql.vectorized.ColumnarBatch
+
+import org.apache.comet.CometConf
+import org.apache.comet.serde.CometOperatorSerde
+import org.apache.comet.serde.OperatorOuterClass
+import org.apache.comet.serde.OperatorOuterClass.Operator
+import org.apache.comet.serde.QueryPlanSerde.serializeDataType
+
+/**
+ * Reads Spark cached table data when the cache was written by Comet's cache
serializer.
+ *
+ * Spark stores cached data through `CachedBatchSerializer`. This node keeps
the scan inside Comet
+ * by asking the serializer to decode cached batches directly into
`ColumnarBatch` output,
+ * avoiding the extra Spark columnar-to-Comet columnar conversion used by the
default path.
+ *
+ * `relationOutput` is the full schema stored in the cache. `scanOutput` is
the subset requested
+ * by this scan after pruning.
+ */
+case class CometInMemoryTableScanExec(
+ originalPlan: InMemoryTableScanExec,
+ serializer: CachedBatchSerializer,
+ cacheBuilder: CachedRDDBuilder,
Review Comment:
**[P2] Canonicalize the embedded scan to preserve exchange reuse**
This leaf extends `CometExec`, not `CometNativeExec`, so it does not get the
latter's canonicalization of embedded plan fields. `originalPlan` is not a
child and retains its expression IDs even after the Comet scan is
canonicalized. Equivalent references to the same cache therefore fail
`sameResult`.
With AQE disabled, a UNION of two identical grouped aggregates over one
cache produces one shuffle exchange plus `ReusedExchange` when native cache
scanning is disabled, but two independent shuffle exchanges when enabled.
Repeating a broadcast join has the same one-versus-two behavior. In the
canonical scans, `cacheBuilder`, `relationOutput` and `scanOutput` compare
equal; only `originalPlan` differs, and canonicalizing that field restores
equality. Please canonicalize the wrapped scan while preserving its predicates
and source semantics to avoid repeating scans, aggregation, shuffle and
broadcast work.
--
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]