diff --git a/.github/workflows/pr_build_linux.yml b/.github/workflows/pr_build_linux.yml index 3917bbc3f39..3549dc9d50b 100644 --- a/.github/workflows/pr_build_linux.yml +++ b/.github/workflows/pr_build_linux.yml @@ -579,6 +579,7 @@ jobs: org.apache.spark.sql.CometCollationSuite org.apache.comet.CometFuzzAggregateSuite org.apache.spark.sql.comet.execution.arrow.CometArrowStreamSuite + org.apache.spark.sql.comet.execution.arrow.CachedBatchRowIteratorSuite org.apache.spark.sql.CometSparkInternalFunctionsSuite - name: "expressions" value: | diff --git a/.github/workflows/pr_build_macos.yml b/.github/workflows/pr_build_macos.yml index 03326cccce0..f3c209dc366 100644 --- a/.github/workflows/pr_build_macos.yml +++ b/.github/workflows/pr_build_macos.yml @@ -227,6 +227,7 @@ jobs: org.apache.spark.sql.CometCollationSuite org.apache.comet.CometFuzzAggregateSuite org.apache.spark.sql.comet.execution.arrow.CometArrowStreamSuite + org.apache.spark.sql.comet.execution.arrow.CachedBatchRowIteratorSuite org.apache.spark.sql.CometSparkInternalFunctionsSuite - name: "expressions" value: | diff --git a/docs/source/user-guide/latest/in-memory-cache.md b/docs/source/user-guide/latest/in-memory-cache.md index c37fa2dc32d..57d09d900eb 100644 --- a/docs/source/user-guide/latest/in-memory-cache.md +++ b/docs/source/user-guide/latest/in-memory-cache.md @@ -53,7 +53,8 @@ relation whose format could change mid-session could not be read back reliably. codec is a runtime config, but each batch records the codec it was written with, so data cached under one setting stays readable after the setting changes. Turning `spark.comet.exec.inMemoryCache.enabled` off at runtime only sends cached scans back to Spark's -execution path; the cached data stays readable either way. +execution path, where Spark operators read them through the row reader described under +[Limitations](#limitations); the cached data stays readable either way. ## Storage format @@ -175,19 +176,45 @@ registrator. ## Limitations -Reads that feed **Spark** operators rather than Comet ones are slower than Spark's own cache -format, and the narrower the read, the wider the gap. Measured by the same benchmark over the same -5M-row relation, with Comet off so that Spark operators consume the cached data: - -| Read shape | Spark's cache format | Comet's cache format | Slowdown | -| ----------------------- | -------------------: | -------------------: | -------: | -| Row count only (0 of 6) | 35 ms | 183 ms | 5.2x | -| 1 of 6 columns | 54 ms | 257 ms | 4.8x | -| 3 of 6 columns | 98 ms | 331 ms | 3.4x | -| 6 of 6 columns | 410 ms | 623 ms | 1.5x | - -This is why the feature is off by default. The cause is not yet established; -[#5485](https://github.com/apache/datafusion-comet/issues/5485) tracks it. +Reads that feed **Spark** operators rather than Comet ones can be slower than Spark's own cache +format, because every cached batch is decoded from Arrow before Spark reads it. How a Spark +operator reads a relation cached in Comet's format depends on the scan below it: + +- With native execution enabled (`spark.comet.exec.enabled=true`), the cache is scanned by + `CometInMemoryTableScan`, and a Spark operator above it reads the scan's batches through + `CometColumnarToRow`, as it would above any other Comet operator. +- With Comet enabled but native execution disabled, the cache is scanned by Spark's + `InMemoryTableScanExec`. When the operator directly above the scan takes part in whole-stage code + generation, as filters, projections and aggregates do, Comet puts Spark's `ColumnarToRowExec` + between the two, and the generated code reads the cached Arrow vectors directly, with no + intermediate row. The plan shows this as a `ColumnarToRow` above the `InMemoryTableScan`. It + needs `spark.sql.inMemoryColumnarStorage.enableVectorizedReader` (on by default) and whole-stage + code generation, and applies to relations of at most `spark.sql.codegen.maxFields` fields (100 by + default, counting nested fields), beyond which Spark reads a cached relation only as rows. It is + not applied in plan-only mode (`spark.comet.explain.planOnly.enabled`), where Spark executes its + own plan unchanged. +- Otherwise the scan's row reader decodes each batch and writes its rows into one reused + `UnsafeRow`. That covers Comet or `spark.comet.exec.inMemoryCache.enabled` turned off at runtime, + and operators that do not take part in code generation, such as exchanges and limits, or a query + that returns the cached rows as they are. + +Measured by the same benchmark over the same 5M-row relation, with native execution off so that +Spark operators consume the cached data, Comet disabled for the row reader and enabled for the fused +reader (Apple M4, JDK 17, Spark 4.1; the average of two runs): + +| Read shape | Spark's cache format | Comet's format, row reader | Comet's format, fused reader | +| ----------------------- | -------------------: | -------------------------: | ---------------------------: | +| Row count only (0 of 6) | 63 ms | 57 ms | 35 ms | +| 1 of 6 columns | 63 ms | 77 ms | 54 ms | +| 3 of 6 columns | 113 ms | 176 ms | 133 ms | +| 6 of 6 columns | 306 ms | 500 ms | 334 ms | + +The fused reader is faster than Spark's own format for the narrowest reads and within 20% of it for +the others. The row reader takes up to 1.6 times as long, and it is the only reader for relations +wider than `spark.sql.codegen.maxFields`: reading every column of relations of 100, 200 and 1500 +nullable `bigint` columns took 2.2 to 2.5 times as long as from Spark's format. These gaps are why +the feature is still off by default; +[#5485](https://github.com/apache/datafusion-comet/issues/5485) tracks them. Comet's serializer exists because Spark's own Arrow cache format ([SPARK-57268](https://issues.apache.org/jira/browse/SPARK-57268)) is only available from Spark diff --git a/spark/src/main/scala/org/apache/comet/CometConf.scala b/spark/src/main/scala/org/apache/comet/CometConf.scala index f0187c8cf02..5d10dae39d9 100644 --- a/spark/src/main/scala/org/apache/comet/CometConf.scala +++ b/spark/src/main/scala/org/apache/comet/CometConf.scala @@ -263,20 +263,25 @@ object CometConf extends ShimCometConf { val COMET_EXEC_IN_MEMORY_CACHE_ENABLED: ConfigEntry[Boolean] = conf("spark.comet.exec.inMemoryCache.enabled") .category(CATEGORY_EXEC) - .doc("Whether to enable Comet native execution for in-memory cached tables. Its value at " + - "startup also decides whether CometDriverPlugin installs Comet's cache serializer, " + - "which stores cached data in Arrow format. Because spark.sql.cache.serializer is a " + - "static config, the cached format is fixed for the application, and disabling this " + - "at runtime only sends cached scans back to Spark's execution path. Relations whose " + - "schema Comet's Arrow writer does not support are always cached in Spark's default " + - "format. Each cached batch is stored as one Arrow IPC record batch with per-buffer " + - "zstd compression, and a scan copies out only the buffers of the columns it projected, " + - "so the unselected ones are never decompressed. Reads that feed Spark operators rather " + - "than Comet ones still pay a row conversion the default format avoids, and can be " + - "slower than Spark's cache. With spark.kryo.registrationRequired=true, also set " + - "spark.kryo.registrator=org.apache.comet.CometKryoRegistrator before creating the " + - "SparkContext, otherwise caching fails as soon as a block is serialized, including " + - "the disk half of the default MEMORY_AND_DISK storage level.") + .doc( + "Whether to enable Comet native scans and fused Spark reads of in-memory cached tables. " + + "Requires spark.comet.enabled=true. At startup, this setting also decides whether " + + "CometDriverPlugin installs Comet's cache serializer, which stores cached data in " + + "Arrow format. Because spark.sql.cache.serializer is a " + + "static config, the cached format is fixed for the application, and disabling this " + + "or spark.comet.enabled at runtime sends cached scans back to Spark's execution path " + + "without the fused reader. Relations whose schema Comet's Arrow writer does not " + + "support are always cached in Spark's default " + + "format. Each cached batch is stored as one Arrow IPC record batch with per-buffer " + + "zstd compression, and a scan copies out only the buffers of the columns it " + + "projected, so the unselected ones are never decompressed. Eligible Spark " + + "whole-stage codegen consumers read cached vectors directly when vectorized cache " + + "reading is enabled; other Spark row consumers use a reusable row buffer. Decoding " + + "costs can still make wide numeric reads slower than Spark's default cache. With " + + "spark.kryo.registrationRequired=true, also set " + + "spark.kryo.registrator=org.apache.comet.CometKryoRegistrator before creating the " + + "SparkContext, otherwise caching fails as soon as a block is serialized, including " + + "the disk half of the default MEMORY_AND_DISK storage level.") .booleanConf .createWithDefault(false) diff --git a/spark/src/main/scala/org/apache/comet/CometSparkSessionExtensions.scala b/spark/src/main/scala/org/apache/comet/CometSparkSessionExtensions.scala index 7da6ed5813d..b93d951ceb9 100644 --- a/spark/src/main/scala/org/apache/comet/CometSparkSessionExtensions.scala +++ b/spark/src/main/scala/org/apache/comet/CometSparkSessionExtensions.scala @@ -54,7 +54,7 @@ import org.apache.comet.shims.ShimCometSparkSessionExtensions * CometSubqueryBroadcastExec for exchange reuse with Comet broadcasts * b. insertTransitions: ColumnarToRow/RowToColumnar added * c. postColumnarTransitions: RevertNativeForTransitionHeavyStages, - * EliminateRedundantTransitions + * EliminateRedundantTransitions, CometCacheColumnarRule * 5. ReuseExchangeAndSubquery -- Spark deduplicates subqueries (sees Comet nodes) * }}} * @@ -78,7 +78,7 @@ import org.apache.comet.shims.ShimCometSparkSessionExtensions * a. preColumnarTransitions: CometRule (no-op, already converted) * b. insertTransitions * c. postColumnarTransitions: RevertNativeForTransitionHeavyStages, - * EliminateRedundantTransitions + * EliminateRedundantTransitions, CometCacheColumnarRule * }}} * * On Spark 3.4, injectQueryStageOptimizerRule is unavailable. CometExecRule does not wrap SABs, diff --git a/spark/src/main/scala/org/apache/comet/rules/CometCacheColumnarRule.scala b/spark/src/main/scala/org/apache/comet/rules/CometCacheColumnarRule.scala new file mode 100644 index 00000000000..df35763820a --- /dev/null +++ b/spark/src/main/scala/org/apache/comet/rules/CometCacheColumnarRule.scala @@ -0,0 +1,100 @@ +/* + * 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.comet.rules + +import org.apache.spark.sql.catalyst.expressions.LeafExpression +import org.apache.spark.sql.catalyst.expressions.codegen.CodegenFallback +import org.apache.spark.sql.catalyst.rules.Rule +import org.apache.spark.sql.comet.execution.arrow.ArrowCachedBatchSerializer +import org.apache.spark.sql.execution.{CodegenSupport, ColumnarToRowExec, ColumnarToRowTransition, SparkPlan, WholeStageCodegenExec} +import org.apache.spark.sql.execution.adaptive.QueryStageExec +import org.apache.spark.sql.execution.columnar.InMemoryTableScanExec +import org.apache.spark.sql.internal.SQLConf + +import org.apache.comet.CometConf.COMET_EXEC_IN_MEMORY_CACHE_ENABLED +import org.apache.comet.CometSparkSessionExtensions.isCometLoaded + +/** + * Lets Spark's generated consumers read cached Arrow vectors without an intermediate UnsafeRow. + * + * Data flows upward. Spark's InputAdapter/whole-stage wrappers and an optional AQE cache stage + * are omitted: + * {{{ + * Before After + * +------------------------+ +------------------------+ + * | Spark codegen consumer | | Spark codegen consumer | + * +------------------------+ +------------------------+ + * ^ ^ + * | UnsafeRow | column values + * +------------------------+ +------------------------+ + * | InMemoryTableScanExec | | ColumnarToRowExec | + * | row iterator | | fused with consumer | + * +------------------------+ +------------------------+ + * ^ + * | ColumnarBatch + * +------------------------+ + * | InMemoryTableScanExec | + * | Arrow vectors | + * +------------------------+ + * }}} + * + * @param preview + * true in the plan-only preview, which shows the plan Comet would execute. Otherwise the rule + * leaves plans alone in plan-only mode, where Spark executes each query unchanged. + */ +case class CometCacheColumnarRule(preview: Boolean = false) extends Rule[SparkPlan] { + override def apply(plan: SparkPlan): SparkPlan = { + if (!isCometLoaded(conf) || !COMET_EXEC_IN_MEMORY_CACHE_ENABLED.get(conf)) return plan + if (!preview && CometRule.planOnlyApplies(conf, plan)) return plan + if (!conf.wholeStageEnabled) return plan + if (conf.getConf(SQLConf.CODEGEN_FACTORY_MODE).toString == "NO_CODEGEN") return plan + + plan.transformUp { + case parent: CodegenSupport + if parent.supportCodegen && !parent.supportsColumnar && + !parent.isInstanceOf[ColumnarToRowTransition] && + !WholeStageCodegenExec.isTooManyFields(conf, parent.schema) && + !parent.children.exists(p => WholeStageCodegenExec.isTooManyFields(conf, p.schema)) && + !parent.expressions.exists(_.exists { + case _: LeafExpression => false + case _: CodegenFallback => true + case _ => false + }) => + // Match the consuming edge rather than every scan: an existing columnar consumer (or a + // cache stage being materialized by AQE) must keep receiving batches. Spark inserts an + // InputAdapter around the scan later, while this transition fuses with the row consumer. + parent.withNewChildren(parent.children.map { + case child if isColumnarCometCache(child) => ColumnarToRowExec(child) + case child => child + }) + } + } + + private def isColumnarCometCache(plan: SparkPlan): Boolean = { + plan.supportsColumnar && (plan match { + case scan: InMemoryTableScanExec => + // The serializer delegates unsupported schemas to Spark, whose cache keeps its own reader. + scan.relation.cacheBuilder.serializer.isInstanceOf[ArrowCachedBatchSerializer] && + ArrowCachedBatchSerializer.supportsSchema(scan.relation.output) + case stage: QueryStageExec => isColumnarCometCache(stage.plan) + case _ => false + }) + } +} diff --git a/spark/src/main/scala/org/apache/comet/rules/CometRule.scala b/spark/src/main/scala/org/apache/comet/rules/CometRule.scala index fa2fb7dc32e..17ef4a6ce1b 100644 --- a/spark/src/main/scala/org/apache/comet/rules/CometRule.scala +++ b/spark/src/main/scala/org/apache/comet/rules/CometRule.scala @@ -29,6 +29,7 @@ import org.apache.spark.sql.execution.{ApplyColumnarRulesAndInsertTransitions, B import org.apache.spark.sql.execution.adaptive.{AdaptiveSparkPlanExec, InsertAdaptiveSparkPlan, QueryStageExec} import org.apache.spark.sql.execution.exchange.{BroadcastExchangeExec, Exchange} import org.apache.spark.sql.execution.reuse.ReuseExchangeAndSubquery +import org.apache.spark.sql.internal.SQLConf import org.apache.comet.{CometConf, ExtendedExplainInfo} import org.apache.comet.CometSparkSessionExtensions.isCometLoaded @@ -36,11 +37,29 @@ import org.apache.comet.shims.ShimCometStreaming object CometRule { - /** Comet's post-columnar rules, shared by `CometColumnar` and the plan-only preview. */ - def postColumnarRules(session: SparkSession, wholePlan: Boolean = false): Seq[Rule[SparkPlan]] = + /** + * Comet's post-columnar rules, shared by `CometColumnar` and the plan-only preview. + * + * @param preview + * true for the plan-only preview, which holds the whole plan and shows the plan Comet would + * execute. + */ + def postColumnarRules(session: SparkSession, preview: Boolean = false): Seq[Rule[SparkPlan]] = Seq( - RevertNativeForTransitionHeavyStages(session, wholePlan), - EliminateRedundantTransitions(session)) + RevertNativeForTransitionHeavyStages(session, wholePlan = preview), + EliminateRedundantTransitions(session), + CometCacheColumnarRule(preview)) + + /** + * Whether plan-only mode applies to `plan`, so that Comet only reports the plan it would + * execute and Spark executes `plan` unchanged. Mirrors the conversion rules' own guards; + * plan-only is scoped to exec being enabled. + */ + private[comet] def planOnlyApplies(conf: SQLConf, plan: SparkPlan): Boolean = + CometConf.COMET_EXPLAIN_PLAN_ONLY_ENABLED.get(conf) && + isCometLoaded(conf) && + !ShimCometStreaming.isStreamingPlan(plan) && + CometConf.COMET_EXEC_ENABLED.get(conf) /** * Canonical hashes of the subquery plans reported for the query this thread is preparing. Spark @@ -140,7 +159,7 @@ case class CometRule(session: SparkSession, queryStagePrep: Boolean = false) private val execRule = CometExecRule(session) override def apply(plan: SparkPlan): SparkPlan = { - if (planOnlyApplies(plan)) { + if (CometRule.planOnlyApplies(conf, plan)) { reportPlanOnlyCoverage(plan) plan } else { @@ -150,13 +169,6 @@ case class CometRule(session: SparkSession, queryStagePrep: Boolean = false) private def convert(plan: SparkPlan): SparkPlan = execRule.apply(scanRule.apply(plan)) - /** Mirrors the conversion rules' own guards; plan-only is scoped to exec being enabled. */ - private def planOnlyApplies(plan: SparkPlan): Boolean = - CometConf.COMET_EXPLAIN_PLAN_ONLY_ENABLED.get(conf) && - isCometLoaded(conf) && - !ShimCometStreaming.isStreamingPlan(plan) && - CometConf.COMET_EXEC_ENABLED.get(conf) - /** Logs the Comet plan for `plan` unless already reported. Never fails the query. */ private def reportPlanOnlyCoverage(plan: SparkPlan): Unit = { try { @@ -183,7 +195,7 @@ case class CometRule(session: SparkSession, queryStagePrep: Boolean = false) val withTransitions = ApplyColumnarRulesAndInsertTransitions(Seq.empty, outputsColumnar = false).apply(converted) val preview = CometRule - .postColumnarRules(session, wholePlan = true) + .postColumnarRules(session, preview = true) .foldLeft(withTransitions) { case (p, rule) => rule(p) } if (topLevel) ReuseExchangeAndSubquery.apply(preview) else preview } diff --git a/spark/src/main/scala/org/apache/spark/sql/comet/execution/arrow/ArrowCachedBatchSerializer.scala b/spark/src/main/scala/org/apache/spark/sql/comet/execution/arrow/ArrowCachedBatchSerializer.scala index 47691556848..5d159747978 100644 --- a/spark/src/main/scala/org/apache/spark/sql/comet/execution/arrow/ArrowCachedBatchSerializer.scala +++ b/spark/src/main/scala/org/apache/spark/sql/comet/execution/arrow/ArrowCachedBatchSerializer.scala @@ -26,7 +26,7 @@ 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, StartsWith, UnsafeProjection} +import org.apache.spark.sql.catalyst.expressions.{Attribute, Expression, GenericInternalRow, IsNotNull, IsNull, StartsWith} import org.apache.spark.sql.catalyst.util.TypeUtils import org.apache.spark.sql.columnar.{CachedBatch, SimpleMetricsCachedBatch, SimpleMetricsCachedBatchSerializer} import org.apache.spark.sql.comet.util.Utils @@ -696,11 +696,7 @@ class ArrowCachedBatchSerializer extends SimpleMetricsCachedBatchSerializer { convertCachedBatchToColumnarBatch(input, cacheAttributes, selectedAttributes, conf) .mapPartitions { batches => - val toUnsafe = UnsafeProjection.create(selectedAttributes, selectedAttributes) - - batches.flatMap { batch => - batch.rowIterator().asScala.map(row => toUnsafe(row).copy()) - } + new CachedBatchRowIterator(selectedAttributes).createObject(batches) } } } diff --git a/spark/src/main/scala/org/apache/spark/sql/comet/execution/arrow/CachedBatchRowIterator.scala b/spark/src/main/scala/org/apache/spark/sql/comet/execution/arrow/CachedBatchRowIterator.scala new file mode 100644 index 00000000000..82f73d25901 --- /dev/null +++ b/spark/src/main/scala/org/apache/spark/sql/comet/execution/arrow/CachedBatchRowIterator.scala @@ -0,0 +1,186 @@ +/* + * 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 org.apache.spark.sql.catalyst.InternalRow +import org.apache.spark.sql.catalyst.expressions.{Attribute, BoundReference, CodeGeneratorWithInterpretedFallback, InterpretedUnsafeProjection, LeafExpression, UnsafeProjection} +import org.apache.spark.sql.catalyst.expressions.codegen._ +import org.apache.spark.sql.catalyst.expressions.codegen.Block._ +import org.apache.spark.sql.internal.SQLConf +import org.apache.spark.sql.types.DataType +import org.apache.spark.sql.vectorized.{ColumnarBatch, ColumnVector} + +/** + * Reads vectors directly into Spark's reusable UnsafeRow buffer. The input iterator owns the + * batches and releases them on advancement or task completion. As with Spark's cache reader, + * callers must copy rows they retain across next(), but the returned row owns its variable-width + * values and remains valid when hasNext() releases the batch that supplied them. + * + * The generated reader hands each column read to GenerateUnsafeProjection as an expression, so it + * splits the field writes of a wide projection into methods of bounded size, as it does for any + * Spark projection. If a generated method still exceeds the huge-method limit, the reader backs + * off the way WholeStageCodegenExec does, here to Spark's projection of each batch row. + */ +private[arrow] class CachedBatchRowIterator(attributes: Seq[Attribute]) + extends CodeGeneratorWithInterpretedFallback[Iterator[ColumnarBatch], Iterator[InternalRow]] { + + private def fields: Seq[BoundReference] = attributes.zipWithIndex.map { case (attr, i) => + BoundReference(i, attr.dataType, attr.nullable) + } + + override protected def createCodeGeneratedObject( + batches: Iterator[ColumnarBatch]): Iterator[InternalRow] = { + val ctx = new CodegenContext + val vectorClass = classOf[ColumnVector].getName + val batchClass = classOf[ColumnarBatch].getName + val columns = ctx.addMutableState( + s"$vectorClass[]", + "columns", + v => s"$v = new $vectorClass[${attributes.length}];", + forceInline = true) + val rowId = ctx.addMutableState(CodeGenerator.JAVA_INT, "rowId", forceInline = true) + val reads = attributes.zipWithIndex.map { case (attr, i) => + VectorValue(s"$columns[$i]", rowId, attr.dataType, attr.nullable) + } + // With ctx.currentVars unset, GenerateUnsafeProjection splits the field writes into methods + // that take the input row as their argument. The reads above ignore it. + val projection = GenerateUnsafeProjection.createCode(ctx, reads) + val batchesRef = ctx.addReferenceObj("batches", batches, "scala.collection.Iterator") + val code = s""" + public Object generate(Object[] references) { + return new SpecificCachedBatchRowIterator(references); + } + + class SpecificCachedBatchRowIterator extends scala.collection.AbstractIterator { + private final Object[] references; + private final scala.collection.Iterator batches; + private int numRows = 0; + ${ctx.declareMutableStates()} + + public SpecificCachedBatchRowIterator(Object[] references) { + this.references = references; + this.batches = $batchesRef; + ${ctx.initMutableStates()} + } + + public boolean hasNext() { + while ($rowId >= numRows && batches.hasNext()) { + $batchClass batch = ($batchClass) batches.next(); + numRows = batch.numRows(); + $rowId = 0; + for (int ordinal = 0; ordinal < $columns.length; ordinal++) { + $columns[ordinal] = batch.column(ordinal); + } + } + return $rowId < numRows; + } + + public InternalRow next() { + if (!hasNext()) throw new java.util.NoSuchElementException(); + InternalRow ${ctx.INPUT_ROW} = null; + ${projection.code} + $rowId++; + return ${projection.value}; + } + + ${ctx.declareAddedFunctions()} + } + """ + val (compiled, stats) = + CodeGenerator.compile(new CodeAndComment(code, ctx.getPlaceHolderToComments())) + // Honor spark.sql.codegen.hugeMethodLimit as whole-stage codegen does, but never go above + // HotSpot's own limit: the config defaults to the largest method the JVM accepts, while this + // runs once per row and HotSpot never JIT-compiles a method longer than + // DEFAULT_JVM_HUGE_METHOD_LIMIT bytes. + val limit = + math.min(SQLConf.get.hugeMethodLimit, CodeGenerator.DEFAULT_JVM_HUGE_METHOD_LIMIT) + if (stats.maxMethodCodeSize > limit) { + logInfo( + s"Generated cache reader for ${attributes.length} columns has a " + + s"${stats.maxMethodCodeSize}-byte method, above the $limit-byte limit; " + + "projecting cached rows with UnsafeProjection instead") + new ProjectedRows(batches, UnsafeProjection.create(fields)) + } else { + compiled.generate(ctx.references.toArray).asInstanceOf[Iterator[InternalRow]] + } + } + + override protected def createInterpretedObject( + batches: Iterator[ColumnarBatch]): Iterator[InternalRow] = + new ProjectedRows(batches, InterpretedUnsafeProjection.createProjection(fields)) +} + +/** + * Projects each batch row through `projection`, which reuses one UnsafeRow and owns the values it + * writes, under the same contract as the generated reader. + */ +private[arrow] class ProjectedRows( + batches: Iterator[ColumnarBatch], + private[arrow] val projection: UnsafeProjection) + extends Iterator[InternalRow] { + private var batch: ColumnarBatch = _ + private var rowId = 0 + private var numRows = 0 + + override def hasNext: Boolean = { + while (rowId >= numRows && batches.hasNext) { + batch = batches.next() + numRows = batch.numRows() + rowId = 0 + } + rowId < numRows + } + + override def next(): InternalRow = { + if (!hasNext) throw new NoSuchElementException + val row = projection(batch.getRow(rowId)) + rowId += 1 + row + } +} + +/** + * The current row of one column of the batch a generated reader is reading. It exists only to be + * code generated, as an expression rather than through ctx.currentVars, which would stop + * GenerateUnsafeProjection from splitting the writer and leave every field in next(). + */ +private case class VectorValue( + column: String, + rowId: String, + dataType: DataType, + nullable: Boolean) + extends LeafExpression { + + override def eval(input: InternalRow): Any = + throw new UnsupportedOperationException(s"$nodeName is only code generated") + + override protected def doGenCode(ctx: CodegenContext, ev: ExprCode): ExprCode = { + val javaType = CodeGenerator.javaType(dataType) + val value = CodeGenerator.getValueFromVector(column, dataType, rowId) + if (nullable) { + ev.copy(code = code""" + boolean ${ev.isNull} = $column.isNullAt($rowId); + $javaType ${ev.value} = ${ev.isNull} ? ${CodeGenerator.defaultValue(dataType)} : ($value); + """) + } else { + ev.copy(code = code"$javaType ${ev.value} = $value;", isNull = FalseLiteral) + } + } +} diff --git a/spark/src/test/scala/org/apache/comet/exec/CometInMemoryCacheSuite.scala b/spark/src/test/scala/org/apache/comet/exec/CometInMemoryCacheSuite.scala index b12ef02c5f7..ca77452cb30 100644 --- a/spark/src/test/scala/org/apache/comet/exec/CometInMemoryCacheSuite.scala +++ b/spark/src/test/scala/org/apache/comet/exec/CometInMemoryCacheSuite.scala @@ -31,13 +31,13 @@ import org.apache.arrow.vector.compression.{CompressionCodec, CompressionUtil, N import org.apache.arrow.vector.types.pojo.ArrowType import org.apache.spark.CometDriverPlugin import org.apache.spark.SparkConf -import org.apache.spark.sql.{CometTestBase, QueryTest, Row} +import org.apache.spark.sql.{CometTestBase, DataFrame, QueryTest, Row} import org.apache.spark.sql.catalyst.expressions.{And, Attribute, AttributeReference, Expression, GreaterThanOrEqual, LessThan, Literal} import org.apache.spark.sql.columnar.{CachedBatch, SimpleMetricsCachedBatch} import org.apache.spark.sql.comet.{CometBroadcastHashJoinExec, CometInMemoryTableScanExec, CometSortExec, CometSortMergeJoinExec} import org.apache.spark.sql.comet.execution.arrow.{ArrowCachedBatchSerializer, CometCachedBatchHelper} import org.apache.spark.sql.comet.util.Utils -import org.apache.spark.sql.execution.SortExec +import org.apache.spark.sql.execution.{ColumnarToRowExec, FilterExec, RowToColumnarExec, SortExec} import org.apache.spark.sql.execution.adaptive.{AdaptiveSparkPlanExec, AQEShuffleReadExec, QueryStageExec, ShuffleQueryStageExec} import org.apache.spark.sql.execution.columnar.{CometInMemoryRelationHelper, InMemoryRelation, InMemoryTableScanExec} import org.apache.spark.sql.execution.exchange.{Exchange, ReusedExchangeExec, ShuffleExchangeLike} @@ -50,6 +50,7 @@ import org.apache.spark.storage.StorageLevel import org.apache.comet.{CometArrowAllocator, CometConf} import org.apache.comet.CometSparkSessionExtensions.{isSpark35Plus, isSpark40Plus} +import org.apache.comet.rules.CometCacheColumnarRule import org.apache.comet.vector.{CometPlainVector, CometVector} class CometInMemoryCacheSuite extends CometTestBase { @@ -351,6 +352,187 @@ class CometInMemoryCacheSuite extends CometTestBase { } } + test("Spark row consumers of Comet cache preserve values across batches") { + for { + adaptive <- Seq(false, true) + mode <- Seq("CODEGEN_ONLY", "NO_CODEGEN") + vectorized <- Seq(false, true) + } { + // Comet on with native execution off, so Spark operators consume the cache scan and the + // generated ones among them read its vectors through the fused transition. + withSQLConf( + CometConf.COMET_ENABLED.key -> "true", + CometConf.COMET_EXEC_IN_MEMORY_CACHE_ENABLED.key -> "true", + CometConf.COMET_EXEC_ENABLED.key -> "false", + CometConf.COMET_SHUFFLE_ENABLED.key -> "false", + SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> adaptive.toString, + SQLConf.CACHE_VECTORIZED_READER_ENABLED.key -> vectorized.toString, + SQLConf.COLUMN_BATCH_SIZE.key -> "7", + SQLConf.CODEGEN_FACTORY_MODE.key -> mode, + SQLConf.WHOLESTAGE_CODEGEN_ENABLED.key -> (mode == "CODEGEN_ONLY").toString, + SQLConf.AUTO_BROADCASTJOIN_THRESHOLD.key -> "-1", + SQLConf.SHUFFLE_PARTITIONS.key -> "2") { + val scalars = Seq( + "boolean", + "tinyint", + "smallint", + "int", + "bigint", + "float", + "double", + "decimal(10,2)", + "decimal(38,2)", + "date", + "timestamp", + "timestamp_ntz").zipWithIndex.map { case (dt, i) => + val value = dt match { + case "date" | "timestamp" | "timestamp_ntz" => + s"cast(date_add(DATE '2000-01-01', cast(id AS INT)) AS $dt)" + case _ => s"cast(id AS $dt)" + } + s"if(id % 3 = 0, null, $value) AS c$i" + } + val source = spark + .range(0, 41, 1, 2) + .selectExpr((Seq("id AS key") ++ scalars ++ Seq( + "if(id % 3 = 0, null, repeat(concat('字', id), cast(id + 1 AS INT))) AS s", + "if(id % 3 = 0, null, cast(concat('binary', id) AS BINARY)) AS b", + "if(id % 3 = 0, null, array(cast(id AS STRING), null)) AS a", + "if(id % 3 = 0, null, named_struct('x', id, 'a', array(cast(id AS STRING)))) AS st", + "if(id % 3 = 0, null, map('k', array(cast(id AS STRING), null))) AS m", + "null AS n")): _*) + + // Each query, and whether a generated Spark operator consumes the cache scan directly. + // The other consumers (the query root, exchanges and limits) read the row iterator. + def queries(df: DataFrame): Seq[(DataFrame, Boolean)] = Seq( + // The generated filter reads every column, so this covers the whole type matrix. + df.filter($"key" >= 0) -> true, + df.select("*") -> false, + df.selectExpr("s AS renamed", "key", "b", "a", "st", "m") -> true, + df.orderBy($"s".desc, $"key") -> false, + df.join(spark.range(41).toDF("join_key"), $"key" === $"join_key") + .select(df("*")) -> false, + df.selectExpr("count(*)") -> true, + df.limit(1) -> false) + + val expected = queries(source).map(_._1.collect().toSeq) + source.cache() + try { + assert(source.count() == 41) + val relation = + spark.sharedState.cacheManager.lookupCachedData(source).get.cachedRepresentation + val buffers = relation.cacheBuilder.cachedColumnBuffers.collect() + assert(buffers.length > 2) + assert(buffers.forall(_.getClass.getSimpleName == "CometCachedBatch")) + queries(source).zip(expected).foreach { case ((df, generatedConsumer), answer) => + val plan = df.queryExecution.executedPlan + checkAnswer(df, answer) + // Inspected after execution, when an adaptive plan is final. + val scans = collect(plan) { case scan: InMemoryTableScanExec => scan } + assert(scans.nonEmpty && scans.forall(_.supportsColumnar == vectorized), plan) + val transitions = collect(plan) { + case c: ColumnarToRowExec if collect(c.child) { case s: InMemoryTableScanExec => + s + }.nonEmpty => + c + } + val fused = generatedConsumer && vectorized && mode == "CODEGEN_ONLY" + assert(transitions.size == (if (fused) 1 else 0), plan) + } + } finally source.unpersist(blocking = true) + } + } + } + + test("Spark generated cache consumers respect runtime enable and codegen settings") { + val planOnly = Seq( + CometConf.COMET_EXPLAIN_PLAN_ONLY_ENABLED.key -> "true", + // Plan-only mode applies only while native execution is enabled. + CometConf.COMET_EXEC_ENABLED.key -> "true") + for { + adaptive <- Seq(false, true) + disabledSettings <- Seq( + Seq(CometConf.COMET_ENABLED.key -> "false"), + Seq(CometConf.COMET_EXEC_IN_MEMORY_CACHE_ENABLED.key -> "false"), + Seq(SQLConf.CODEGEN_FACTORY_MODE.key -> "NO_CODEGEN"), + Seq(SQLConf.WHOLESTAGE_CODEGEN_ENABLED.key -> "false"), + planOnly) + } { + withSQLConf( + CometConf.COMET_ENABLED.key -> "true", + CometConf.COMET_EXEC_IN_MEMORY_CACHE_ENABLED.key -> "true", + CometConf.COMET_EXEC_ENABLED.key -> "false", + CometConf.COMET_SHUFFLE_ENABLED.key -> "false", + SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> adaptive.toString, + SQLConf.CACHE_VECTORIZED_READER_ENABLED.key -> "true", + SQLConf.WHOLESTAGE_CODEGEN_ENABLED.key -> "true", + SQLConf.CODEGEN_FACTORY_MODE.key -> "CODEGEN_ONLY", + SQLConf.COLUMN_BATCH_SIZE.key -> "7", + SQLConf.SHUFFLE_PARTITIONS.key -> "2") { + val source = spark + .range(0, 41, 1, 2) + .selectExpr("id AS key", "if(id % 3 = 0, null, concat('字', id)) AS s") + def query = source + .filter("key >= 7") + .selectExpr("sum(key)", "sum(length(s))", "count(*)") + val expected = query.collect().toSeq + source.cache() + try { + val builder = spark.sharedState.cacheManager + .lookupCachedData(source) + .get + .cachedRepresentation + .cacheBuilder + // Materialize with fusion enabled, then disable and re-enable it on the same cache. + Seq(true, false, true).zipWithIndex.foreach { case (enabled, index) => + val settings = if (enabled) Seq.empty else disabledSettings + withSQLConf(settings: _*) { + val cold = index == 0 + val df = query + val plan = df.queryExecution.executedPlan + // Planning must not materialize the cache or replace AQE's cache-stage metadata. + assert(builder.isCachedColumnBuffersLoaded != cold, plan.toString) + // checkToRDD = false keeps checkAnswer from loading the cache with a query of its + // own, so the cold run's table-cache stage materializes and AQE re-plans above it. + QueryTest.checkAnswer(df, expected, checkToRDD = false) + assert(builder.isCachedColumnBuffersLoaded) + val transitions = collect(plan) { + case c: ColumnarToRowExec if collect(c.child) { case s: InMemoryTableScanExec => + s + }.nonEmpty => + c + } + assert(transitions.size == (if (enabled) 1 else 0), plan.toString) + assert(collect(plan) { case s: CometInMemoryTableScanExec => s }.isEmpty) + if (adaptive && isSpark35Plus) { + assert(collect(plan) { + case s: QueryStageExec + if s.getClass.getSimpleName == "TableCacheQueryStageExec" => + s + }.size == 1) + } + val scan = collect(plan) { case s: InMemoryTableScanExec => s }.head + assert(scan.supportsColumnar) + // A cache scan can also be the root of a columnar request or already have a + // transition. Applying the rule again must preserve those input/output contracts. + Seq(scan, ColumnarToRowExec(scan), RowToColumnarExec(scan)).foreach { boundary => + assert(CometCacheColumnarRule()(boundary).fastEquals(boundary)) + } + // The plan-only preview shows the plan Comet would execute, so it still fuses a + // generated consumer that the executed plan leaves alone in plan-only mode. + val consumer = FilterExec(Literal.TrueLiteral, scan) + val fusedConsumer = FilterExec(Literal.TrueLiteral, ColumnarToRowExec(scan)) + assert(CometCacheColumnarRule()(consumer).fastEquals(fusedConsumer) == enabled) + assert( + CometCacheColumnarRule(preview = true)(consumer).fastEquals(fusedConsumer) == + (enabled || disabledSettings == planOnly)) + } + } + } finally source.unpersist(blocking = true) + } + } + } + // Column expression and the reason the serializer has to decline it. private val unsupportedForArrowCache = Seq( // Interval types have no Arrow vector in Utils.getFieldVector. Without the schema check in diff --git a/spark/src/test/scala/org/apache/spark/sql/benchmark/CometInMemoryCacheBenchmark.scala b/spark/src/test/scala/org/apache/spark/sql/benchmark/CometInMemoryCacheBenchmark.scala index f33a53af299..9317b5b6bf6 100644 --- a/spark/src/test/scala/org/apache/spark/sql/benchmark/CometInMemoryCacheBenchmark.scala +++ b/spark/src/test/scala/org/apache/spark/sql/benchmark/CometInMemoryCacheBenchmark.scala @@ -27,6 +27,7 @@ import org.apache.spark.sql.SparkSession import org.apache.spark.sql.catalyst.expressions.AttributeReference import org.apache.spark.sql.comet.CometInMemoryTableScanExec import org.apache.spark.sql.comet.execution.arrow.ArrowCachedBatchSerializer +import org.apache.spark.sql.execution.ColumnarToRowExec import org.apache.spark.sql.execution.columnar.{CometInMemoryRelationHelper, DefaultCachedBatchSerializer, InMemoryRelation, InMemoryTableScanExec} import org.apache.spark.sql.execution.vectorized.OnHeapColumnVector import org.apache.spark.sql.internal.{SQLConf, StaticSQLConf} @@ -234,6 +235,7 @@ object CometInMemoryCacheBenchmark extends CometBenchmarkBase { runCodecBenchmark(flatRelation) runSparkOperatorBenchmark(flatRelation) } + runWideSparkOperatorBenchmark() } /** @@ -339,16 +341,28 @@ object CometInMemoryCacheBenchmark extends CometBenchmarkBase { /** * Reads that feed Spark operators rather than Comet ones, against Spark's own cache format. * - * Comet is off in every case, so this measures Spark consuming the cached data: the shape where - * Comet's format has something to lose, and the reason the feature is off by default. Both - * formats are cached from the same relation, one copy at a time as in runCodecBenchmark, and - * each case checks which serializer cached the relation it reads. + * Native execution is off in every case, so this measures Spark consuming the cached data: the + * shape where Comet's format has something to lose, and the reason the feature is off by + * default. Comet's format is read two ways. With Comet off, every Spark operator reads rows + * from the cache scan's row reader. With Comet on, a generated Spark operator instead reads the + * cached vectors through a ColumnarToRowExec fused into its generated code, which is how these + * reads run when Comet is enabled without native execution. With native execution on, Spark + * operators above the cache read CometInMemoryTableScan's batches through CometColumnarToRow + * instead, which this does not measure. + * + * Both formats are cached from the same relation, one copy at a time as in runCodecBenchmark, + * and each case checks which serializer cached the relation it reads and which reader it uses. */ private def runSparkOperatorBenchmark(relation: CachedRelation): Unit = { val view = s"${relation.table}_spark_operators" - val formats = Seq( - "Spark's cache format" -> classOf[DefaultCachedBatchSerializer].getName, - "Comet's cache format" -> classOf[ArrowCachedBatchSerializer].getName) + val reads = Seq( + SparkOperatorRead("Spark's cache format", sparkSerializer, sparkOperatorConf), + SparkOperatorRead("Comet's cache format, row reader", cometSerializer, sparkOperatorConf), + SparkOperatorRead( + "Comet's cache format, fused reader", + cometSerializer, + fusedReaderConf, + fused = true)) spark.catalog.clearCache() withTempTable(view) { @@ -356,19 +370,7 @@ object CometInMemoryCacheBenchmark extends CometBenchmarkBase { .sql(s"SELECT ${relation.columns.mkString(", ")} FROM ${relation.source}") .createOrReplaceTempView(view) - var cachedBy: String = null - def cacheBy(serializer: String): Unit = if (cachedBy != serializer) { - spark.catalog.uncacheTable(view) - cachedBy = null - withCacheSerializer(serializer) { - withSQLConf(sparkOperatorConf: _*) { - spark.catalog.cacheTable(view) - spark.table(view).count() - } - } - cachedBy = serializer - } - + val cache = new OneCachedCopy(view) Seq( ("row count only (0 of 6 columns)", s"SELECT count(*) FROM $view", 0), ("narrow projection (1 of 6 columns)", s"SELECT count(k) FROM $view", 1), @@ -381,23 +383,7 @@ object CometInMemoryCacheBenchmark extends CometBenchmarkBase { s"in-memory cache read by Spark operators, $label", relation.rows, output = output) - formats.foreach { case (name, serializer) => - var verified = false - // Re-caching in this case's format is setup, so it is outside the timer, and it only - // happens on the case's first call, which is a warmup iteration. - benchmark.addTimerCase(name) { timer => - cacheBy(serializer) - withSQLConf(sparkOperatorConf: _*) { - if (!verified) { - verifySparkOperatorRead(query, scanned, serializer) - verified = true - } - timer.startTiming() - spark.sql(query).noop() - timer.stopTiming() - } - } - } + reads.foreach(addSparkOperatorCase(benchmark, cache, query, scanned, _)) benchmark.run() } @@ -405,6 +391,97 @@ object CometInMemoryCacheBenchmark extends CometBenchmarkBase { } } + /** + * Spark operators reading every column of relations wider than the six-column one, at widths + * where a row reader that writes every column in one generated method would stop being JIT + * compiled (past about a hundred columns) or fail to compile at all (about 1500). + * + * Each case writes the relation to the noop sink, which takes rows straight from the cache + * scan, so both formats are read by their serializers' row readers. Past + * spark.sql.codegen.maxFields (100 by default) that is how any Spark operator reads a cached + * relation, because the scan stops offering columnar output and nothing can fuse with it. Every + * width holds about the same number of values, so the cases differ in width rather than in data + * volume. + */ + private def runWideSparkOperatorBenchmark(): Unit = { + val reads = Seq( + SparkOperatorRead("Spark's cache format", sparkSerializer, sparkOperatorConf), + SparkOperatorRead("Comet's cache format, row reader", cometSerializer, sparkOperatorConf)) + + Seq(100, 200, 1500).foreach { width => + val rows = 20 * 1000 * 1000 / width + val view = s"comet_cache_bench_wide_$width" + spark.catalog.clearCache() + withTempTable(view) { + spark + .range(0, rows, 1, 4) + .selectExpr((0 until width).map(i => + s"if(id % 8 = ${i % 8}, null, id + $i) AS c$i"): _*) + .createOrReplaceTempView(view) + + val cache = new OneCachedCopy(view) + val benchmark = new Benchmark( + s"in-memory cache read by Spark operators, all $width columns", + rows, + output = output) + reads.foreach(addSparkOperatorCase(benchmark, cache, s"SELECT * FROM $view", width, _)) + benchmark.run() + + spark.catalog.uncacheTable(view) + } + } + } + + /** How Spark operators read a cached relation: its format and the settings of the read. */ + private case class SparkOperatorRead( + name: String, + serializer: String, + conf: Seq[(String, String)], + fused: Boolean = false) + + /** + * Holds one cached copy of `view` at a time. Two copies could not coexist anyway: the cache + * manager keys on the plan rather than the name, so a second one would find the first. + */ + private class OneCachedCopy(view: String) { + private var cachedBy: String = _ + + def cacheBy(serializer: String): Unit = if (cachedBy != serializer) { + spark.catalog.uncacheTable(view) + cachedBy = null + withCacheSerializer(serializer) { + withSQLConf(sparkOperatorConf: _*) { + spark.catalog.cacheTable(view) + spark.table(view).count() + } + } + cachedBy = serializer + } + } + + private def addSparkOperatorCase( + benchmark: Benchmark, + cache: OneCachedCopy, + query: String, + scanned: Int, + read: SparkOperatorRead): Unit = { + var verified = false + // Re-caching in this case's format is setup, so it is outside the timer, and it only happens + // on the case's first call, which is a warmup iteration. + benchmark.addTimerCase(read.name) { timer => + cache.cacheBy(read.serializer) + withSQLConf(read.conf: _*) { + if (!verified) { + verifySparkOperatorRead(query, scanned, read) + verified = true + } + timer.startTiming() + spark.sql(query).noop() + timer.stopTiming() + } + } + } + // spark.sql.cache.serializer is static, and InMemoryRelation memoizes the serializer it names // for the life of the JVM. It looks the name up in the active session's conf when a relation is // cached, though, so setting it there directly and clearing the memoized instance around one @@ -422,10 +499,14 @@ object CometInMemoryCacheBenchmark extends CometBenchmarkBase { } } - // Pins what a Spark-operator case claims: no Comet operator anywhere, and one cache scan that - // reads the columns its label counts from a relation the named serializer cached. The last is - // what catches both formats silently reading one copy. - private def verifySparkOperatorRead(query: String, scanned: Int, serializer: String): Unit = { + // Pins what a Spark-operator case claims: no Comet operator anywhere, one cache scan that + // reads the columns its label counts from a relation the named serializer cached, and the + // reader the case names. The serializer check is what catches both formats silently reading one + // copy, and the reader check a fused case silently measuring the row reader. + private def verifySparkOperatorRead( + query: String, + scanned: Int, + read: SparkOperatorRead): Unit = { val executed = spark.sql(query).queryExecution.executedPlan val plan = executed.toString() assert(executed.find(_.nodeName.startsWith("Comet")).isEmpty, s"Expected no Comet:\n$plan") @@ -435,7 +516,15 @@ object CometInMemoryCacheBenchmark extends CometBenchmarkBase { scans.head.attributes.length == scanned, s"Expected the scan to read $scanned columns:\n$plan") val actual = scans.head.relation.cacheBuilder.serializer.getClass.getName - assert(actual == serializer, s"Expected a relation cached by $serializer, not $actual") + assert( + actual == read.serializer, + s"Expected a relation cached by ${read.serializer}, not $actual") + val transitions = executed.collect { + case c: ColumnarToRowExec if c.exists(_.isInstanceOf[InMemoryTableScanExec]) => c + } + assert( + transitions.length == (if (read.fused) 1 else 0), + s"Expected the ${if (read.fused) "fused" else "row"} reader:\n$plan") } /** What the cached relation behind `view` occupies, summed over its batches as written. */ @@ -602,4 +691,18 @@ object CometInMemoryCacheBenchmark extends CometBenchmarkBase { CometConf.COMET_ENABLED.key -> "false", CometConf.COMET_EXEC_ENABLED.key -> "false", "spark.sql.inMemoryColumnarStorage.batchSize" -> "10000") + + // Comet on but native execution and shuffle off, so the operators are still Spark's and the + // generated ones read Comet's cache through the fused reader. On-heap mode has to be enabled + // explicitly, as in cacheConf, or Comet stays unloaded and the reader is never fused. + private val fusedReaderConf: Seq[(String, String)] = Seq( + CometConf.COMET_ENABLED.key -> "true", + CometConf.COMET_EXEC_ENABLED.key -> "false", + CometConf.COMET_SHUFFLE_ENABLED.key -> "false", + CometConf.COMET_EXEC_IN_MEMORY_CACHE_ENABLED.key -> "true", + "spark.comet.exec.onHeap.enabled" -> "true", + "spark.sql.inMemoryColumnarStorage.batchSize" -> "10000") + + private val sparkSerializer = classOf[DefaultCachedBatchSerializer].getName + private val cometSerializer = classOf[ArrowCachedBatchSerializer].getName } diff --git a/spark/src/test/scala/org/apache/spark/sql/comet/execution/arrow/CachedBatchRowIteratorSuite.scala b/spark/src/test/scala/org/apache/spark/sql/comet/execution/arrow/CachedBatchRowIteratorSuite.scala new file mode 100644 index 00000000000..6d2a5385e81 --- /dev/null +++ b/spark/src/test/scala/org/apache/spark/sql/comet/execution/arrow/CachedBatchRowIteratorSuite.scala @@ -0,0 +1,212 @@ +/* + * 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 java.nio.charset.StandardCharsets.UTF_8 + +import org.scalatest.funsuite.AnyFunSuite + +import org.apache.arrow.memory.RootAllocator +import org.apache.arrow.vector.VarCharVector +import org.apache.spark.sql.catalyst.InternalRow +import org.apache.spark.sql.catalyst.expressions.{AttributeReference, InterpretedUnsafeProjection, UnsafeRow} +import org.apache.spark.sql.execution.vectorized.OnHeapColumnVector +import org.apache.spark.sql.internal.SQLConf +import org.apache.spark.sql.types.{IntegerType, LongType, StringType} +import org.apache.spark.sql.vectorized.{ColumnarBatch, ColumnVector} + +import org.apache.comet.vector.CometPlainVector + +class CachedBatchRowIteratorSuite extends AnyFunSuite { + Seq("CODEGEN_ONLY", "NO_CODEGEN").foreach { mode => + def withMode(settings: (String, String)*)(f: => Unit): Unit = { + val conf = new SQLConf + conf.setConfString(SQLConf.CODEGEN_FACTORY_MODE.key, mode) + settings.foreach { case (key, value) => conf.setConfString(key, value) } + SQLConf.withExistingConf(conf)(f) + } + + test(s"$mode: rows own Arrow values across batch release and reuse the output buffer") { + withMode() { + val allocator = new RootAllocator(Long.MaxValue) + val vectors = Seq(Seq("first", null), Seq("字" * 1000, "last")).map { values => + val vector = new VarCharVector("s", allocator) + values.zipWithIndex.foreach { case (value, i) => + if (value == null) vector.setNull(i) else vector.setSafe(i, value.getBytes(UTF_8)) + } + vector.setValueCount(values.size) + vector + } + try { + // Match the cache decoder: hasNext releases a consumed batch before the next is read. + val batches = vectors.iterator.flatMap { vector => + new Iterator[ColumnarBatch] { + private var emitted = false + override def hasNext: Boolean = { + if (emitted) vector.close() + !emitted + } + override def next(): ColumnarBatch = { + emitted = true + new ColumnarBatch(Array(new CometPlainVector(vector, false)), 2) + } + } + } + val attributes = Seq(AttributeReference("s", StringType, nullable = true)()) + val rows = new CachedBatchRowIterator(attributes).createObject(batches) + assert(rows.hasNext && rows.hasNext) + val first = rows.next().asInstanceOf[UnsafeRow] + val saved = first.copy() + assert(first.getUTF8String(0).toString == "first") + assert(rows.next() eq first) + assert(first.isNullAt(0)) + assert(rows.hasNext && rows.hasNext) + assert(first.isNullAt(0)) + assert(rows.next().getUTF8String(0).toString == "字" * 1000) + val last = rows.next() + assert(!rows.hasNext && !rows.hasNext) + assert(allocator.getAllocatedMemory == 0) + assert(last.getUTF8String(0).toString == "last") + assert(saved.getUTF8String(0).toString == "first") + intercept[NoSuchElementException](rows.next()) + } finally { + vectors.foreach(_.close()) + allocator.close() + } + } + } + + test(s"$mode: empty input, empty batches, and zero-column rows") { + withMode() { + val factory = new CachedBatchRowIterator(Seq.empty) + val empty = factory.createObject(Iterator.empty) + assert(!empty.hasNext) + intercept[NoSuchElementException](empty.next()) + val batches = Seq(0, 2, 0, 3, 0).map { n => + new ColumnarBatch(Array.empty[ColumnVector], n) + } + val rows = factory.createObject(batches.iterator) + assert(rows.map { row => + assert(row.isInstanceOf[UnsafeRow] && row.numFields == 0) + 1 + }.sum == 5) + intercept[NoSuchElementException](rows.next()) + } + } + + // Before the generated reader let GenerateUnsafeProjection split its writer, next() passed + // HotSpot's 8000-byte JIT limit near 100 columns and Janino's 64 KB method limit near 1500. + // The generated reader is only kept when none of its methods is above the former. + val reader = if (mode == "CODEGEN_ONLY") "the generated reader" else "the interpreted reader" + Seq(100, 200, 1500).foreach { width => + test(s"$mode: $width-column projections use $reader") { + withMode() { + val (attributes, batches) = wideInput(width) + try { + val rows = new CachedBatchRowIterator(attributes).createObject(batches.iterator) + rows match { + case projected: ProjectedRows => + assert(mode == "NO_CODEGEN", "The generated reader has a method too large to JIT") + assert(projected.projection.isInstanceOf[InterpretedUnsafeProjection]) + case _ => assert(mode == "CODEGEN_ONLY") + } + checkWideRows(attributes, rows) + } finally batches.foreach(_.close()) + } + } + } + + test( + s"$mode: a generated method above the huge-method limit falls back to UnsafeProjection") { + // Below any generated method, as if the reader had grown past HotSpot's limit. + withMode(SQLConf.WHOLESTAGE_HUGE_METHOD_LIMIT.key -> "1") { + val (attributes, batches) = wideInput(200) + try { + val rows = new CachedBatchRowIterator(attributes).createObject(batches.iterator) + val projection = rows.asInstanceOf[ProjectedRows].projection + assert(projection.isInstanceOf[InterpretedUnsafeProjection] == (mode == "NO_CODEGEN")) + checkWideRows(attributes, rows) + } finally batches.foreach(_.close()) + } + } + } + + private val wideTypes = Seq(IntegerType, LongType, StringType) + + /** + * Nullable and required int, bigint and string columns, two rows in each of two batches, so + * that every column is bound again at the batch boundary. Odd rows are null where a column + * allows it. + */ + private def wideInput(width: Int): (Seq[AttributeReference], Seq[ColumnarBatch]) = { + val attributes = (0 until width).map { i => + AttributeReference(s"c$i", wideTypes(i % 3), nullable = i % 2 == 0)() + } + val batches = Seq(0, 2).map { firstRow => + val columns = attributes.zipWithIndex.map { case (attr, i) => + val column = new OnHeapColumnVector(2, attr.dataType) + Seq(0, 1).foreach { r => + val row = firstRow + r + if (attr.nullable && row % 2 == 1) { + column.putNull(r) + } else { + attr.dataType match { + case IntegerType => column.putInt(r, wideValue(i, row).asInstanceOf[Int]) + case LongType => column.putLong(r, wideValue(i, row).asInstanceOf[Long]) + case _ => column.putByteArray(r, wideValue(i, row).toString.getBytes(UTF_8)) + } + } + } + column + } + new ColumnarBatch(columns.toArray[ColumnVector], 2) + } + (attributes, batches) + } + + private def wideValue(column: Int, row: Int): Any = wideTypes(column % 3) match { + case IntegerType => column * 10 + row + case LongType => column * 10000000000L + row + case _ => s"$column:$row" + } + + private def checkWideRows( + attributes: Seq[AttributeReference], + rows: Iterator[InternalRow]): Unit = { + (0 until 4).foreach { row => + assert(rows.hasNext) + val actual = rows.next() + assert(actual.numFields == attributes.size) + attributes.zipWithIndex.foreach { case (attr, i) => + if (attr.nullable && row % 2 == 1) { + assert(actual.isNullAt(i), s"c$i in row $row") + } else { + val value = attr.dataType match { + case IntegerType => actual.getInt(i) + case LongType => actual.getLong(i) + case _ => actual.getUTF8String(i).toString + } + assert(value == wideValue(i, row), s"c$i in row $row") + } + } + } + assert(!rows.hasNext) + } +}