From 2b8706351d514c4189c660027d90d60ba2f859ad Mon Sep 17 00:00:00 2001 From: Andy Grove Date: Sat, 3 Oct 2026 08:42:02 -0600 Subject: [PATCH 1/6] feat: run the operators above a typed Dataset operation natively Typed Dataset operations (map, flatMap, mapPartitions, mapGroups, cogroup) pass JVM objects between their operators, so they stay on Spark, and today the operators above them stay on Spark too until the next shuffle. Every typed operation ends in SerializeFromObjectExec, whose output is ordinary rows. With the new spark.comet.convert.typedDataset.enabled, CometExecRule puts a CometSparkToColumnarExec above it, so a partial aggregate, a broadcast join or a native shuffle above the operation runs natively. Spark inserts no columnar transitions below a RowToColumnarTransition, so the rule adds them to the subtree under the conversion with Spark's own ApplyColumnarRulesAndInsertTransitions. Without them the typed operation would read its Comet child through Spark's interpreted columnar-to-row path. CometExecRule no longer tags a ColumnarToRowTransition as an unsupported operator, which it did to these transitions on its second pass under AQE. Off by default: it is 1.4-2.2x faster when an aggregate over many groups sits above the typed operation, and slower when the work above is cheap. Adds CometTypedDatasetSuite and CometTypedDatasetBenchmark. --- .github/workflows/pr_build_linux.yml | 1 + .github/workflows/pr_build_macos.yml | 1 + .../adding_a_new_operator.md | 17 +- docs/source/user-guide/latest/datasources.md | 4 + docs/source/user-guide/latest/operators.md | 2 +- .../scala/org/apache/comet/CometConf.scala | 12 + .../apache/comet/rules/CometExecRule.scala | 41 +++ .../comet/exec/CometTypedDatasetSuite.scala | 262 ++++++++++++++++++ .../CometTypedDatasetBenchmark.scala | 192 +++++++++++++ 9 files changed, 525 insertions(+), 7 deletions(-) create mode 100644 spark/src/test/scala/org/apache/comet/exec/CometTypedDatasetSuite.scala create mode 100644 spark/src/test/scala/org/apache/spark/sql/benchmark/CometTypedDatasetBenchmark.scala diff --git a/.github/workflows/pr_build_linux.yml b/.github/workflows/pr_build_linux.yml index 57c695a2719..3a6fb58489f 100644 --- a/.github/workflows/pr_build_linux.yml +++ b/.github/workflows/pr_build_linux.yml @@ -519,6 +519,7 @@ jobs: org.apache.comet.exec.CometGenerateExecSuite org.apache.comet.exec.CometWindowExecSuite org.apache.comet.exec.CometJoinSuite + org.apache.comet.exec.CometTypedDatasetSuite org.apache.spark.sql.comet.CometMapInBatchSuite org.apache.spark.sql.execution.python.CometArrowPythonRunnerSuite org.apache.comet.CometNativeSuite diff --git a/.github/workflows/pr_build_macos.yml b/.github/workflows/pr_build_macos.yml index 51b6d493cb2..e902e95578a 100644 --- a/.github/workflows/pr_build_macos.yml +++ b/.github/workflows/pr_build_macos.yml @@ -223,6 +223,7 @@ jobs: org.apache.comet.exec.CometGenerateExecSuite org.apache.comet.exec.CometWindowExecSuite org.apache.comet.exec.CometJoinSuite + org.apache.comet.exec.CometTypedDatasetSuite org.apache.spark.sql.comet.CometMapInBatchSuite org.apache.spark.sql.execution.python.CometArrowPythonRunnerSuite org.apache.comet.CometNativeSuite diff --git a/docs/source/contributor-guide/adding_a_new_operator.md b/docs/source/contributor-guide/adding_a_new_operator.md index 20d526b150e..ba51aa25c92 100644 --- a/docs/source/contributor-guide/adding_a_new_operator.md +++ b/docs/source/contributor-guide/adding_a_new_operator.md @@ -139,12 +139,17 @@ and leaves the stage itself in place. `AppendColumnsExec`, `AppendColumnsWithObjectExec`, `MapGroupsExec`, and `CoGroupExec`. They convert rows to JVM objects, run an arbitrary user function on those objects, or convert them back, and most of them pass the objects to the next operator as an `ObjectType` column, which has no -Arrow representation. None of this can run natively. Per-row operators can still stay inside a -Comet plan: `MapElementsExec`, the operator behind `Dataset.map`, generates its call to the user -function as a Catalyst `Invoke` expression, so the deserializer, the call, and the serializer could -run together as one projection in the JVM codegen dispatcher. `mapPartitions`, `mapGroups`, and -`cogroup` pass the user function an iterator or a whole group, so there is no per-row expression to -build. A typed `filter` is planned as an ordinary `FilterExec`, not as one of these operators. +Arrow representation. None of this can run natively, so these operators stay on Spark. Every typed +operation ends in `SerializeFromObjectExec`, though, whose output is ordinary rows. With +`spark.comet.convert.typedDataset.enabled`, `CometExecRule` puts a `CometSparkToColumnarExec` above +it, so the operators above the typed operation can run natively. Spark inserts no columnar +transitions below a `RowToColumnarTransition`, so the rule inserts them for the typed operation's +own operators itself. Fusing the deserializer, the `Invoke` that calls the user function, and the +serializer of `Dataset.map` into one projection in the JVM codegen dispatcher was tried in +[#5714](https://github.com/apache/datafusion-comet/pull/5714) and dropped. The dispatcher only +calls into Spark's own classes, and the conversion gets nearly the same speedup for `map` while +also covering the operations that pass the user function an iterator or a whole group. A typed +`filter` is planned as an ordinary `FilterExec`, not as one of these operators. **Driver-side commands.** `ExecutedCommandExec` runs a `RunnableCommand`, such as DDL or `SET`, on the driver, so there is no data path for Comet to accelerate. diff --git a/docs/source/user-guide/latest/datasources.md b/docs/source/user-guide/latest/datasources.md index c6d9debaf04..d52178139df 100644 --- a/docs/source/user-guide/latest/datasources.md +++ b/docs/source/user-guide/latest/datasources.md @@ -63,6 +63,10 @@ This includes row-backed `ExistingRDD` inputs when `spark.comet.sparkToColumnar.supportedOperatorList` includes `RDDScan`. Spark still produces the RDD rows; conversion lets eligible downstream operators execute in Comet. +The same types apply to the output of typed `Dataset` operations, such as `map`, which Comet +converts when `spark.comet.convert.typedDataset.enabled=true`. A column of any other type keeps +the operators above the typed operation on Spark. + ## Data Catalogs ### Apache Iceberg diff --git a/docs/source/user-guide/latest/operators.md b/docs/source/user-guide/latest/operators.md index 676f295ccf9..f23d6939ca9 100644 --- a/docs/source/user-guide/latest/operators.md +++ b/docs/source/user-guide/latest/operators.md @@ -54,7 +54,7 @@ omitted from the tables below and may be reconsidered based on demand: - **Structured Streaming operators** (`StateStoreSaveExec`, `StateStoreRestoreExec`, `StreamingSymmetricHashJoinExec`, and similar): Comet targets batch execution. - **Cartesian / cross joins** (`CartesianProductExec`): rare and expensive, with little acceleration benefit. - **Pickled (non-Arrow) Python UDFs** (`BatchEvalPythonExec`): Comet accelerates Arrow-based Python UDFs only ([#4234](https://github.com/apache/datafusion-comet/pull/4234)). -- **Typed Dataset operators** (`DeserializeToObjectExec`, `SerializeFromObjectExec`, `MapElementsExec`, `MapPartitionsExec`, `MapGroupsExec`, `CoGroupExec`, `AppendColumnsExec`, and similar): produced by `map`, `mapPartitions`, `groupByKey`, `cogroup`, and other typed `Dataset` transformations. They exist to run user JVM functions on JVM objects, which Comet cannot do natively. +- **Typed Dataset operators** (`DeserializeToObjectExec`, `SerializeFromObjectExec`, `MapElementsExec`, `MapPartitionsExec`, `MapGroupsExec`, `CoGroupExec`, `AppendColumnsExec`, and similar): produced by `map`, `mapPartitions`, `groupByKey`, `cogroup`, and other typed `Dataset` transformations. They exist to run user JVM functions on JVM objects, which Comet cannot do natively. The operators above them can still run natively: set `spark.comet.convert.typedDataset.enabled=true` and Comet converts the output of a typed operation to Arrow. This is disabled by default because it can be slower than Spark when the operators above do little work, such as an aggregate over a few groups, which Spark compiles together with the typed operation into one loop. ## Scans diff --git a/spark/src/main/scala/org/apache/comet/CometConf.scala b/spark/src/main/scala/org/apache/comet/CometConf.scala index 797a2363250..8e0f9165239 100644 --- a/spark/src/main/scala/org/apache/comet/CometConf.scala +++ b/spark/src/main/scala/org/apache/comet/CometConf.scala @@ -203,6 +203,18 @@ object CometConf extends ShimCometConf { .booleanConf .createWithDefault(false) + val COMET_CONVERT_FROM_TYPED_DATASET_ENABLED: ConfigEntry[Boolean] = + conf("spark.comet.convert.typedDataset.enabled") + .category(CATEGORY_EXEC) + .doc("When enabled, the output of typed Dataset operations, such as `map`, `flatMap`, " + + "`mapPartitions` and `groupByKey(...).mapGroups`, will be converted to Arrow format so " + + "that the operators above them can run natively. The user function still runs in " + + "Spark. This pays off when the operators above do enough work, such as an " + + "aggregation over many groups, and can be slower when they are cheap, such as an " + + "aggregation over a few groups after a selective filter.") + .booleanConf + .createWithDefault(false) + val COMET_EXEC_ENABLED: ConfigEntry[Boolean] = conf(s"$COMET_EXEC_CONFIG_PREFIX.enabled") .category(CATEGORY_EXEC) .doc( diff --git a/spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala b/spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala index e5e250f75dc..776c2cfaaa5 100644 --- a/spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala +++ b/spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala @@ -464,6 +464,14 @@ case class CometExecRule(session: SparkSession) case op if shouldApplySparkToColumnar(conf, op) => convertToComet(op, CometSparkToColumnarExec).getOrElse(op) + // Typed Dataset operations (`map`, `flatMap`, `mapPartitions`, `mapGroups`, ...) pass JVM + // objects between their operators, so those stay on Spark. Each of them ends in + // `SerializeFromObjectExec`, though, whose output is ordinary rows, and converting those to + // Arrow lets the operators above the typed operation run natively. + case op: SerializeFromObjectExec + if CometConf.COMET_CONVERT_FROM_TYPED_DATASET_ENABLED.get(conf) => + convertTypedDatasetOutput(op) + // Spark 4.0+: replace only the per-task write, leaving DataWritingCommandExec - and // therefore Spark's commit protocol, stats trackers and SaveMode handling - in place. // `V1WritesUtils.getWriteFilesOpt` matches the `WriteFilesExecBase` trait there, which is @@ -583,6 +591,10 @@ case class CometExecRule(session: SparkSession) // Some execs should never be replaced. We include // these cases specially here so we do not add a misleading 'info' message. op + case _: ColumnarToRowTransition => + // A transition does no work of its own. This rule only meets one that + // `convertTypedDatasetOutput` inserted on an earlier pass over the same plan. + op case _: WriteFilesExec => // The write is converted at the enclosing DataWritingCommandExec above: on Spark 3.x // by replacing the whole command, on 4.0+ by converting this child from there. @@ -1166,6 +1178,35 @@ case class CometExecRule(session: SparkSession) private def hasEnabledHandler(op: SparkPlan): Boolean = allExecs.get(op.getClass).exists(_.enabledConfig.forall(_.get(op.conf))) + /** + * Converts the rows a typed Dataset operation produces to Arrow, so the operators above it can + * run natively. See [[CometConf.COMET_CONVERT_FROM_TYPED_DATASET_ENABLED]]. + * + * Spark inserts the columnar transitions after this rule, but it does not look below a + * `RowToColumnarTransition` such as `CometSparkToColumnarExec`. That is harmless above a leaf. + * Here the typed operation's own operators sit below the conversion, and without a transition + * they would read a Comet child through `CometExec.doExecute`, Spark's interpreted + * columnar-to-row path. So the subtree gets its transitions now, from Spark's own rule, and + * `EliminateRedundantTransitions` later replaces each one over a Comet child with Comet's own. + * Spark's rule leaves existing transitions alone, which matters because this rule runs over the + * same plan twice under AQE. + */ + private def convertTypedDatasetOutput(op: SerializeFromObjectExec): SparkPlan = { + val unsupported = op.output.filterNot(a => + CometSparkToColumnarExec.isTypeSupported(a.dataType, a.name, ListBuffer.empty)) + if (unsupported.nonEmpty) { + withFallbackReason( + op, + "Comet cannot convert the output of a typed Dataset operation to Arrow because it does " + + "not support the type of these columns: " + + unsupported.map(a => s"${a.name}: ${a.dataType.simpleString}").mkString(", ")) + } else { + val withTransitions = + ApplyColumnarRulesAndInsertTransitions(Seq.empty, outputsColumnar = false).apply(op) + convertToComet(withTransitions, CometSparkToColumnarExec).getOrElse(withTransitions) + } + } + private def shouldApplySparkToColumnar(conf: SQLConf, op: SparkPlan): Boolean = { // Only consider converting leaf nodes to columnar currently, so that all the following // operators can have a chance to be converted to columnar. Leaf operators that output diff --git a/spark/src/test/scala/org/apache/comet/exec/CometTypedDatasetSuite.scala b/spark/src/test/scala/org/apache/comet/exec/CometTypedDatasetSuite.scala new file mode 100644 index 00000000000..d0550a2f922 --- /dev/null +++ b/spark/src/test/scala/org/apache/comet/exec/CometTypedDatasetSuite.scala @@ -0,0 +1,262 @@ +/* + * 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.exec + +import java.util.concurrent.atomic.AtomicLong + +import org.apache.spark.sql.{CometTestBase, DataFrame, Dataset} +import org.apache.spark.sql.catalyst.expressions.aggregate.Partial +import org.apache.spark.sql.comet.{CometBroadcastHashJoinExec, CometHashAggregateExec, CometSparkToColumnarExec} +import org.apache.spark.sql.execution.{ColumnarToRowExec, ColumnarToRowTransition, InputAdapter, SerializeFromObjectExec, SparkPlan, WholeStageCodegenExec} +import org.apache.spark.sql.functions.{broadcast, col, size, sum} +import org.apache.spark.sql.internal.SQLConf + +import org.apache.comet.{CometConf, ExtendedExplainInfo} + +// Top-level, so the encoders need no outer pointer, which is the ordinary user shape. +case class TypedDsRec(a: Int, b: String) + +case class TypedDsWide(i: Int, s: String, d: java.math.BigDecimal, opt: Option[Long]) + +case class TypedDsNested(id: Int, inner: TypedDsRec, tags: Seq[String]) + +case class TypedDsInts(id: Int, xs: Seq[Int]) + +/** Counts calls to a user function. Comet tests run in local mode, so tasks see this object. */ +object TypedDsCounter { + val calls = new AtomicLong(0) +} + +/** Tests for [[CometConf.COMET_CONVERT_FROM_TYPED_DATASET_ENABLED]]. */ +class CometTypedDatasetSuite extends CometTestBase { + + import testImplicits._ + + /** Defines a test that runs with the typed Dataset output conversion enabled. */ + private def convertTest(name: String)(f: => Unit): Unit = + test(name) { + withSQLConf(CometConf.COMET_CONVERT_FROM_TYPED_DATASET_ENABLED.key -> "true")(f) + } + + /** `rows` rows of `TypedDsRec` in a Parquet table, read back as a typed Dataset. */ + private def withRecs(rows: Int = 200)(f: Dataset[TypedDsRec] => Unit): Unit = + withParquetTable((0 until rows).map(i => (i, (i % 13).toString)), "tbl") { + f(spark.sql("SELECT _1 AS a, _2 AS b FROM tbl").as[TypedDsRec]) + } + + private def conversions(plan: SparkPlan): Seq[CometSparkToColumnarExec] = + collectWithSubqueries(plan) { case c: CometSparkToColumnarExec => c } + + /** + * Spark inserts no columnar transitions below a `CometSparkToColumnarExec`, so the rule adds + * them for the typed operation's own operators. Without one, an operator reads its Comet child + * through `CometExec.doExecute`, Spark's interpreted columnar-to-row path, which gives the + * right answer slowly, so only the plan shows it. + */ + private def assertRowOperatorsReadThroughTransitions(plan: SparkPlan): Unit = { + def unwrap(p: SparkPlan): SparkPlan = p match { + case InputAdapter(child) => child + case other => other + } + val bare = collectWithSubqueries(plan) { + case p + if !p.supportsColumnar && !p.isInstanceOf[ColumnarToRowTransition] && + !p.isInstanceOf[WholeStageCodegenExec] && + p.children.exists(c => unwrap(c).supportsColumnar) => + p + } + assert(bare.isEmpty, s"row operators read a columnar child without a transition:\n$plan") + assert( + collectWithSubqueries(plan) { case c: ColumnarToRowExec => c }.isEmpty, + s"expected Comet's columnar-to-row transitions, not Spark's:\n$plan") + } + + /** + * Checks the answer, that everything above the typed operation runs natively, and that the + * operation's output is converted right above its `SerializeFromObjectExec`. + */ + private def checkConverted(df: => DataFrame): SparkPlan = { + val (_, plan) = + checkSparkAnswerAndOperator(df, includeClasses = Seq(classOf[CometSparkToColumnarExec])) + conversions(plan).foreach { c => + val converted = c.child match { + case w: WholeStageCodegenExec => w.child + case other => other + } + assert( + converted.isInstanceOf[SerializeFromObjectExec], + s"expected the conversion directly above SerializeFromObject:\n$plan") + } + assertRowOperatorsReadThroughTransitions(plan) + // The transitions are added on the first of the rule's passes under AQE. A later pass must + // not report them as operators Comet failed to convert. + val reasons = new ExtendedExplainInfo().getFallbackReasons(plan) + assert(!reasons.exists(_.contains("ColumnarToRow")), reasons) + plan + } + + private def nativePartialAggregates(plan: SparkPlan): Seq[CometHashAggregateExec] = + collectWithSubqueries(plan) { + case a: CometHashAggregateExec if a.modes.contains(Partial) => a + } + + convertTest("the aggregate and the shuffle above a map run natively") { + Seq("true", "false").foreach { aqe => + withSQLConf(SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> aqe) { + withRecs() { ds => + val plan = checkConverted(ds.map(r => TypedDsRec(r.a + 1, r.b)).groupBy("b").count()) + assert(nativePartialAggregates(plan).nonEmpty, s"AQE $aqe:\n$plan") + } + } + } + } + + convertTest("flatMap, mapPartitions, mapGroups and cogroup") { + withRecs() { ds => + checkConverted(ds.flatMap(r => Seq(r, TypedDsRec(-r.a, r.b))).groupBy("b").count()) + checkConverted(ds.mapPartitions(_.map(r => TypedDsRec(r.a + 1, r.b))).groupBy("b").count()) + checkConverted( + ds.groupByKey(_.b) + .mapGroups((k, rows) => TypedDsRec(rows.map(_.a).sum, k)) + .groupBy("a") + .count()) + checkConverted( + ds.groupByKey(_.b) + .cogroup(ds.groupByKey(_.b))((k, left, right) => + Iterator(TypedDsRec(left.size + right.size, k))) + .groupBy("a") + .count()) + } + } + + convertTest("a Dataset built from an RDD of objects") { + val ds = spark.sparkContext + .parallelize((0 until 200).map(i => TypedDsRec(i, (i % 13).toString)), 4) + .toDS() + checkConverted(ds.groupBy("b").agg(sum("a"))) + } + + convertTest("a broadcast join above the typed operation runs natively") { + withRecs() { ds => + withParquetTable((0 until 13).map(i => (i.toString, s"name$i")), "dim") { + val dim = spark.table("dim").toDF("b", "name") + val plan = checkConverted(ds.map(r => TypedDsRec(r.a * 2, r.b)).join(broadcast(dim), "b")) + assert( + collectWithSubqueries(plan) { case j: CometBroadcastHashJoinExec => j }.nonEmpty, + plan) + } + } + } + + convertTest("decimal, string, Option, nested struct and string array fields") { + val wide = (1 to 40).map(i => + (i, s"s$i", new java.math.BigDecimal(s"$i.25"), if (i % 3 == 0) None else Some(i * 10L))) + withParquetTable(wide, "wide") { + val ds = spark.sql("SELECT _1 AS i, _2 AS s, _3 AS d, _4 AS opt FROM wide").as[TypedDsWide] + checkConverted( + ds.map(r => TypedDsWide(r.i % 4, r.s, r.d, r.opt.map(_ + 1))) + .groupBy("i") + .agg(sum("d"), sum("opt"))) + } + val nested = (1 to 30).map(i => (i, (i, s"n$i"), Seq(s"t$i", null))) + withParquetTable(nested, "nested") { + val ds = spark + .sql( + "SELECT _1 AS id, named_struct('a', _2._1, 'b', _2._2) AS inner, _3 AS tags " + + "FROM nested") + .as[TypedDsNested] + checkConverted( + ds.map(r => TypedDsNested(r.id % 3, TypedDsRec(r.inner.a, r.inner.b), r.tags.reverse)) + .groupBy("id") + .count()) + } + } + + convertTest("nothing above the typed operation consumes Arrow, so nothing is converted") { + withRecs() { ds => + val (_, plan) = checkSparkAnswer(ds.map(r => TypedDsRec(r.a + 1, r.b)).toDF()) + assert(conversions(plan).isEmpty, plan) + assertRowOperatorsReadThroughTransitions(plan) + } + } + + convertTest("columns the conversion does not support keep the operators above on Spark") { + withParquetTable((0 until 20).map(i => (i, Seq(i, i + 1))), "ints") { + val ds = spark.sql("SELECT _1 AS id, _2 AS xs FROM ints").as[TypedDsInts] + val (_, plan) = checkSparkAnswerAndFallbackReason( + // The aggregate reads `xs`, or Spark would prune it from the serializer. + ds.map(r => TypedDsInts(r.id % 3, r.xs.reverse)).groupBy("id").agg(sum(size(col("xs")))), + "Comet cannot convert the output of a typed Dataset operation to Arrow because it " + + "does not support the type of these columns: xs: array") + assert(conversions(plan).isEmpty, plan) + } + } + + convertTest("the user function runs once per row") { + withRecs(500) { ds => + TypedDsCounter.calls.set(0) + val df = ds + .map { r => + TypedDsCounter.calls.incrementAndGet() + TypedDsRec(r.a, r.b) + } + .groupBy("b") + .count() + assert(df.collect().map(_.getLong(1)).sum == 500) + assert(conversions(df.queryExecution.executedPlan).nonEmpty, df.queryExecution.executedPlan) + assert(TypedDsCounter.calls.get() == 500) + } + } + + convertTest("an exception from the user function fails the query") { + // One partition, so the failing row is past the first batch, where an error from a JVM input + // reaches the native plan through the Arrow stream rather than on the JVM. + withTempPath { path => + spark + .range(20000) + .selectExpr("CAST(id AS INT) AS a", "CAST(id % 13 AS STRING) AS b") + .coalesce(1) + .write + .parquet(path.toString) + val df = spark.read + .parquet(path.toString) + .as[TypedDsRec] + .map { r => + if (r.a == 15000) throw new IllegalArgumentException("typed map failed at 15000") + r + } + .groupBy("b") + .count() + assert(conversions(df.queryExecution.executedPlan).nonEmpty, df.queryExecution.executedPlan) + val e = intercept[Exception](df.collect()) + assert( + causeChain(e).exists(t => Option(t.getMessage).exists(_.contains("failed at 15000"))), + e) + } + } + + test("off by default") { + withRecs() { ds => + val (_, plan) = checkSparkAnswer(ds.map(r => TypedDsRec(r.a + 1, r.b)).groupBy("b").count()) + assert(conversions(plan).isEmpty, plan) + assert(nativePartialAggregates(plan).isEmpty, plan) + } + } +} diff --git a/spark/src/test/scala/org/apache/spark/sql/benchmark/CometTypedDatasetBenchmark.scala b/spark/src/test/scala/org/apache/spark/sql/benchmark/CometTypedDatasetBenchmark.scala new file mode 100644 index 00000000000..eb52920eb07 --- /dev/null +++ b/spark/src/test/scala/org/apache/spark/sql/benchmark/CometTypedDatasetBenchmark.scala @@ -0,0 +1,192 @@ +/* + * 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.benchmark + +import org.apache.spark.benchmark.Benchmark +import org.apache.spark.sql.{DataFrame, Dataset, Encoder, Encoders, Row} +import org.apache.spark.sql.catalyst.expressions.aggregate.Partial +import org.apache.spark.sql.comet.{CometHashAggregateExec, CometPlan, CometSparkToColumnarExec} +import org.apache.spark.sql.execution.SparkPlan +import org.apache.spark.sql.functions.{col, count, length, lit, sum} +import org.apache.spark.sql.internal.SQLConf + +import org.apache.comet.CometConf + +// Top-level, so the encoders need no outer pointer, which is the ordinary user shape. +case class TypedDatasetBenchRec(a: Long, b: String) + +case class TypedDatasetBenchWide(a: Long, b: String, c: Long, d: String) + +/** + * Compares three ways to run a query over the output of a typed Dataset operation: + * + * - Spark: Comet disabled. + * - Comet: the default. The typed operation runs in Spark, and Comet takes over again at the + * shuffle above it, so the operators in between stay on Spark. + * - Comet, converted: `spark.comet.convert.typedDataset.enabled`, which converts the output of + * the typed operation to Arrow so the operators above it run natively. + * + * The cases sweep how much work sits above the typed operation, from an aggregate over 100 groups + * that Spark's whole-stage codegen fuses with the operation to one over a million groups. Every + * arm's result and plan are checked before it is timed, and the Comet arm runs again at the end + * of each case to show the noise. To run this benchmark: + * {{{ + * SPARK_GENERATE_BENCHMARK_FILES=1 make benchmark-org.apache.spark.sql.benchmark.CometTypedDatasetBenchmark + * }}} + * Results will be written to "spark/benchmarks/CometTypedDatasetBenchmark-**results.txt". + */ +object CometTypedDatasetBenchmark extends CometBenchmarkBase { + + private val numRows = 4 * 1024 * 1024 + private val loKeys = 100 + private val hiKeys = 1024 * 1024 + + // Declared rather than imported from `spark.implicits`, which would start the session while + // this object initializes. + private implicit val recEncoder: Encoder[TypedDatasetBenchRec] = + Encoders.product[TypedDatasetBenchRec] + private implicit val wideEncoder: Encoder[TypedDatasetBenchWide] = + Encoders.product[TypedDatasetBenchWide] + + private case class Arm(name: String, confs: Seq[(String, String)], check: SparkPlan => Unit) + + // No AQE and one shuffle partition, so every arm plans the same shape on every iteration. + private val planConfs = + Seq(SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "false", SQLConf.SHUFFLE_PARTITIONS.key -> "1") + + private def cometConfs(convert: Boolean) = planConfs ++ Seq( + CometConf.COMET_ENABLED.key -> "true", + CometConf.COMET_EXEC_ENABLED.key -> "true", + CometConf.COMET_CONVERT_FROM_TYPED_DATASET_ENABLED.key -> convert.toString) + + private def conversions(plan: SparkPlan): Seq[CometSparkToColumnarExec] = + collect(plan) { case c: CometSparkToColumnarExec => c } + + private val arms = Seq( + Arm( + "Spark", + planConfs :+ (CometConf.COMET_ENABLED.key -> "false"), + plan => assert(collect(plan) { case c: CometPlan => c }.isEmpty, plan)), + Arm("Comet", cometConfs(convert = false), plan => assert(conversions(plan).isEmpty, plan)), + Arm( + "Comet, converted", + cometConfs(convert = true), + plan => { + assert(conversions(plan).nonEmpty, plan) + assert( + collect(plan) { + case a: CometHashAggregateExec if a.modes.contains(Partial) => a + }.nonEmpty, + plan) + })) + + /** The table read back as a typed Dataset, grouping on `keys`. */ + private def typed(keys: String): Dataset[TypedDatasetBenchRec] = + spark + .table("parquetV1Table") + .select(col("id").as("a"), col(keys).as("b")) + .as[TypedDatasetBenchRec] + + /** Totals the per-group results in `s`, so each case returns one row. */ + private def total(grouped: DataFrame): DataFrame = grouped.agg(sum("s"), count(lit(1))) + + private def mapThenGroup(keys: String): DataFrame = + total( + typed(keys).map(r => TypedDatasetBenchRec(r.a + 1, r.b)).groupBy("b").agg(sum("a").as("s"))) + + private def mapThenFilterThenGroup(keys: String): DataFrame = + total( + typed(keys) + .map(r => TypedDatasetBenchRec(r.a * 2, r.b)) + .filter(col("a") % 3 === 0) + .groupBy("b") + .agg(sum("a").as("s"))) + + private val cases: Seq[(String, () => DataFrame)] = Seq( + "map -> group by 100 keys" -> (() => mapThenGroup("lo")), + "map -> group by 1M keys" -> (() => mapThenGroup("hi")), + "map -> filter -> group by 100 keys" -> (() => mapThenFilterThenGroup("lo")), + "map -> filter -> group by 1M keys" -> (() => mapThenFilterThenGroup("hi")), + "map -> group by long key, 1M keys" -> (() => + total( + typed("lo") + .map(r => TypedDatasetBenchRec(r.a % hiKeys, r.b)) + .groupBy("a") + .agg(sum(length(col("b"))).as("s")))), + "map, 4 columns -> group by 1M keys" -> (() => + total( + typed("hi") + .map(r => TypedDatasetBenchWide(r.a + 1, r.b, r.a * 2, r.b + "!")) + .groupBy("b") + .agg((sum("a") + sum("c") + sum(length(col("d")))).as("s")))), + "mapPartitions -> group by 1M keys" -> (() => + total( + typed("hi") + .mapPartitions(_.map(r => TypedDatasetBenchRec(r.a + 1, r.b))) + .groupBy("b") + .agg(sum("a").as("s"))))) + + private def runArm(arm: Arm, query: () => DataFrame): (Seq[Row], SparkPlan) = { + var result: (Seq[Row], SparkPlan) = null + withSQLConf(arm.confs: _*) { + val df = query() + result = (df.collect().toSeq, df.queryExecution.executedPlan) + } + result + } + + override def runCometBenchmark(mainArgs: Array[String]): Unit = { + withTempPath { dir => + withTempTable("parquetV1Table") { + prepareTable( + dir, + spark + .range(numRows) + .selectExpr( + "id", + s"CAST(id % $loKeys AS STRING) AS lo", + s"CAST(id % $hiKeys AS STRING) AS hi")) + + cases.foreach { case (name, query) => + runBenchmark(name) { + // Check every arm's plan and result before timing anything. + val results = arms.map(runArm(_, query)) + arms.zip(results).foreach { case (arm, (rows, plan)) => + arm.check(plan) + assert( + rows == results.head._1, + s"${arm.name} returned $rows, expected ${results.head._1}") + } + + val benchmark = new Benchmark(name, numRows, output = output) + (arms :+ arms(1).copy(name = "Comet (repeat)")).foreach { arm => + benchmark.addCase(arm.name) { _ => + withSQLConf(arm.confs: _*) { + query().collect() + } + } + } + benchmark.run() + } + } + } + } + } +} From 1ec0e66265af56a3b2cfc4e04e6a160efac3c4fc Mon Sep 17 00:00:00 2001 From: Andy Grove Date: Sat, 3 Oct 2026 09:40:07 -0600 Subject: [PATCH 2/6] fix: keep wide-decimal shuffles above a typed Dataset conversion on JVM shuffle A typed Dataset conversion moves the shuffle above it from Comet's columnar shuffle, which partitions with Spark's hash, to native shuffle. Native shuffle hashes decimals wider than 18 digits differently from Spark (#5994), so a join on such keys with an input that stays on columnar shuffle, for example one with an array column the conversion declines, put matching keys in different partitions and returned 9 rows instead of 100. Such a shuffle now stays on columnar shuffle unless it has one partition, the same rule #6005 proposes for every native shuffle. Also declare the benchmark's row count as a Long, which the strict Spark 3.5 compile requires. --- .../contributor-guide/native_shuffle.md | 6 +++ .../shuffle/CometShuffleExchangeExec.scala | 40 +++++++++++++++++-- .../comet/exec/CometTypedDatasetSuite.scala | 30 ++++++++++++++ .../CometTypedDatasetBenchmark.scala | 2 +- 4 files changed, 74 insertions(+), 4 deletions(-) diff --git a/docs/source/contributor-guide/native_shuffle.md b/docs/source/contributor-guide/native_shuffle.md index b882acd25a0..4ef01578412 100644 --- a/docs/source/contributor-guide/native_shuffle.md +++ b/docs/source/contributor-guide/native_shuffle.md @@ -71,6 +71,12 @@ Native shuffle (`CometExchange`) is selected when all of the following condition Spark's `mapsort` normalization makes physical entry order irrelevant. A collated string at any depth still disqualifies the key. The config defaults to `false` pending measurement of the nested hashing paths, so by default a complex hash key falls back to JVM shuffle. + - A hash key that is or contains a decimal wider than 18 digits stays on JVM shuffle when the + shuffle's stage starts at a typed `Dataset` conversion + (`spark.comet.convert.typedDataset.enabled`) and the shuffle has more than one partition. + Native shuffle hashes such decimals differently from Spark + ([#5994](https://github.com/apache/datafusion-comet/issues/5994)). Without the conversion + this shuffle would have used JVM shuffle, and a join partner may still use it. ## Architecture diff --git a/spark/src/main/scala/org/apache/spark/sql/comet/execution/shuffle/CometShuffleExchangeExec.scala b/spark/src/main/scala/org/apache/spark/sql/comet/execution/shuffle/CometShuffleExchangeExec.scala index 7a0d231a141..ffd475fefd5 100644 --- a/spark/src/main/scala/org/apache/spark/sql/comet/execution/shuffle/CometShuffleExchangeExec.scala +++ b/spark/src/main/scala/org/apache/spark/sql/comet/execution/shuffle/CometShuffleExchangeExec.scala @@ -34,11 +34,11 @@ import org.apache.spark.sql.catalyst.expressions.{Attribute, BoundReference, Exp import org.apache.spark.sql.catalyst.expressions.codegen.LazilyGeneratedOrdering import org.apache.spark.sql.catalyst.plans.logical.Statistics import org.apache.spark.sql.catalyst.plans.physical._ -import org.apache.spark.sql.comet.{CometFilterExec, CometMetricNode, CometNativeExec, CometNativeScanExec, CometPlan, CometProjectExec, CometSinkPlaceHolder, NativeExecContext} +import org.apache.spark.sql.comet.{CometFilterExec, CometMetricNode, CometNativeExec, CometNativeScanExec, CometPlan, CometProjectExec, CometScanWrapper, CometSinkPlaceHolder, CometSparkToColumnarExec, NativeExecContext} import org.apache.spark.sql.comet.execution.arrow.CometArrowStream import org.apache.spark.sql.execution._ -import org.apache.spark.sql.execution.adaptive.ShuffleQueryStageExec -import org.apache.spark.sql.execution.exchange.{ENSURE_REQUIREMENTS, ShuffleExchangeExec, ShuffleExchangeLike, ShuffleOrigin} +import org.apache.spark.sql.execution.adaptive.{QueryStageExec, ShuffleQueryStageExec} +import org.apache.spark.sql.execution.exchange.{ENSURE_REQUIREMENTS, Exchange, ShuffleExchangeExec, ShuffleExchangeLike, ShuffleOrigin} import org.apache.spark.sql.execution.metric.{SQLMetric, SQLMetrics, SQLShuffleReadMetricsReporter, SQLShuffleWriteMetricsReporter} import org.apache.spark.sql.internal.SQLConf import org.apache.spark.sql.types.{ArrayType, BinaryType, BooleanType, ByteType, DataType, DateType, DecimalType, DoubleType, FloatType, IntegerType, LongType, MapType, ShortType, StringType, StructField, StructType, TimestampNTZType, TimestampType} @@ -524,6 +524,29 @@ object CometShuffleExchangeExec None } + /** Whether `dt` is, or contains, a decimal that native shuffle does not hash as Spark does. */ + private def hasWideDecimal(dt: DataType): Boolean = dt match { + case d: DecimalType => d.precision > 18 + case StructType(fields) => fields.exists(f => hasWideDecimal(f.dataType)) + case ArrayType(elementType, _) => hasWideDecimal(elementType) + case MapType(keyType, valueType, _) => hasWideDecimal(keyType) || hasWideDecimal(valueType) + case _ => false + } + + /** + * Whether the stage feeding a shuffle starts at a typed Dataset conversion (see + * [[CometConf.COMET_CONVERT_FROM_TYPED_DATASET_ENABLED]]). `CometExecRule` decides the shuffle + * before it removes its placeholders, so the conversion is still inside its `CometScanWrapper`. + */ + private def readsTypedDatasetConversion(plan: SparkPlan): Boolean = plan match { + case _: Exchange | _: QueryStageExec => false + case CometScanWrapper(_, conversion: CometSparkToColumnarExec) => + conversion.child.isInstanceOf[SerializeFromObjectExec] + case conversion: CometSparkToColumnarExec => + conversion.child.isInstanceOf[SerializeFromObjectExec] + case other => other.children.exists(readsTypedDatasetConversion) + } + /** * Reasons the native shuffle path cannot handle this shuffle. Empty means native is supported. * Pure: does not tag the node. @@ -622,6 +645,17 @@ object CometShuffleExchangeExec reasons += s"unsupported hash partitioning data type for native shuffle: $dt" } } + // A typed Dataset conversion moves the shuffle above it from Comet's columnar shuffle, + // which partitions with Spark's hash, to native shuffle. Native shuffle hashes a decimal + // wider than 18 digits differently from Spark (#5994), so a join with an input that is + // still on the columnar shuffle would put matching keys in different partitions. Leave + // such a shuffle where it was. A single partition hashes nothing. + if (partitioning.numPartitions > 1 && + expressions.exists(e => hasWideDecimal(e.dataType)) && + readsTypedDatasetConversion(s.child)) { + reasons += "a shuffle above a typed Dataset conversion that hashes a decimal wider " + + "than 18 digits stays on Comet's columnar shuffle, which hashes it as Spark does" + } case SinglePartition => // we already checked that the input types are supported case RangePartitioning(orderings, _) => diff --git a/spark/src/test/scala/org/apache/comet/exec/CometTypedDatasetSuite.scala b/spark/src/test/scala/org/apache/comet/exec/CometTypedDatasetSuite.scala index d0550a2f922..08f9cbfa8e2 100644 --- a/spark/src/test/scala/org/apache/comet/exec/CometTypedDatasetSuite.scala +++ b/spark/src/test/scala/org/apache/comet/exec/CometTypedDatasetSuite.scala @@ -24,6 +24,7 @@ import java.util.concurrent.atomic.AtomicLong import org.apache.spark.sql.{CometTestBase, DataFrame, Dataset} import org.apache.spark.sql.catalyst.expressions.aggregate.Partial import org.apache.spark.sql.comet.{CometBroadcastHashJoinExec, CometHashAggregateExec, CometSparkToColumnarExec} +import org.apache.spark.sql.comet.execution.shuffle.{CometColumnarShuffle, CometShuffleExchangeExec} import org.apache.spark.sql.execution.{ColumnarToRowExec, ColumnarToRowTransition, InputAdapter, SerializeFromObjectExec, SparkPlan, WholeStageCodegenExec} import org.apache.spark.sql.functions.{broadcast, col, size, sum} import org.apache.spark.sql.internal.SQLConf @@ -39,6 +40,10 @@ case class TypedDsNested(id: Int, inner: TypedDsRec, tags: Seq[String]) case class TypedDsInts(id: Int, xs: Seq[Int]) +case class TypedDsDecimal(k: java.math.BigDecimal, v: Long) + +case class TypedDsDecimalInts(k: java.math.BigDecimal, xs: Seq[Int]) + /** Counts calls to a user function. Comet tests run in local mode, so tasks see this object. */ object TypedDsCounter { val calls = new AtomicLong(0) @@ -209,6 +214,31 @@ class CometTypedDatasetSuite extends CometTestBase { } } + convertTest("a join on wide decimal keys with an input that is not converted") { + // The left input converts and the right one, with its array column, does not. Native + // shuffle hashes a decimal wider than 18 digits differently from Spark's partitioner (#5994), + // so the shuffle above the conversion has to stay on Comet's columnar shuffle like the right + // input's, or matching keys land in different partitions and the join loses rows. + withSQLConf( + SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "false", + SQLConf.AUTO_BROADCASTJOIN_THRESHOLD.key -> "-1", + SQLConf.SHUFFLE_PARTITIONS.key -> "10") { + val left = spark + .range(0, 100, 1, 2) + .map(i => TypedDsDecimal(new java.math.BigDecimal(i.longValue), i.longValue)) + .alias("l") + val right = spark + .range(0, 100, 1, 2) + .map(i => TypedDsDecimalInts(new java.math.BigDecimal(i.longValue), Seq(i.intValue))) + .alias("r") + val (_, plan) = checkSparkAnswer( + left.join(right, col("l.k") === col("r.k")).select(col("l.v"), col("r.xs"))) + assert(conversions(plan).nonEmpty, plan) + val shuffles = collectWithSubqueries(plan) { case s: CometShuffleExchangeExec => s } + assert(shuffles.nonEmpty && shuffles.forall(_.shuffleType == CometColumnarShuffle), plan) + } + } + convertTest("the user function runs once per row") { withRecs(500) { ds => TypedDsCounter.calls.set(0) diff --git a/spark/src/test/scala/org/apache/spark/sql/benchmark/CometTypedDatasetBenchmark.scala b/spark/src/test/scala/org/apache/spark/sql/benchmark/CometTypedDatasetBenchmark.scala index eb52920eb07..3686fbc328f 100644 --- a/spark/src/test/scala/org/apache/spark/sql/benchmark/CometTypedDatasetBenchmark.scala +++ b/spark/src/test/scala/org/apache/spark/sql/benchmark/CometTypedDatasetBenchmark.scala @@ -54,7 +54,7 @@ case class TypedDatasetBenchWide(a: Long, b: String, c: Long, d: String) */ object CometTypedDatasetBenchmark extends CometBenchmarkBase { - private val numRows = 4 * 1024 * 1024 + private val numRows = 4L * 1024 * 1024 private val loKeys = 100 private val hiKeys = 1024 * 1024 From 61e08895b40d9e33d1ddf54170afb7bf991f7071 Mon Sep 17 00:00:00 2001 From: Andy Grove Date: Sat, 3 Oct 2026 10:09:47 -0600 Subject: [PATCH 3/6] refactor: simplify the typed Dataset conversion's shuffle guard and tests - Use Spark's existsRecursively(DecimalType.isByteArrayDecimalType) for the wide-decimal check, fold the duplicated conversion cases in readsTypedDatasetConversion, and check the config and earlier native shuffle reasons before walking the plan. Add a TODO to drop the guard once #5994 lands. - Tests: assert the join's exchanges with checkCometExchange, drop conditions that cannot change the transition assertion, write the Parquet table once for both AQE settings, and take checkConverted's query by value. - Benchmark: import spark.implicits instead of declaring encoders. --- .../shuffle/CometShuffleExchangeExec.scala | 28 +++++++------------ .../comet/exec/CometTypedDatasetSuite.scala | 27 ++++++++---------- .../CometTypedDatasetBenchmark.scala | 9 ++---- 3 files changed, 23 insertions(+), 41 deletions(-) diff --git a/spark/src/main/scala/org/apache/spark/sql/comet/execution/shuffle/CometShuffleExchangeExec.scala b/spark/src/main/scala/org/apache/spark/sql/comet/execution/shuffle/CometShuffleExchangeExec.scala index ffd475fefd5..92e03629dbe 100644 --- a/spark/src/main/scala/org/apache/spark/sql/comet/execution/shuffle/CometShuffleExchangeExec.scala +++ b/spark/src/main/scala/org/apache/spark/sql/comet/execution/shuffle/CometShuffleExchangeExec.scala @@ -37,7 +37,7 @@ import org.apache.spark.sql.catalyst.plans.physical._ import org.apache.spark.sql.comet.{CometFilterExec, CometMetricNode, CometNativeExec, CometNativeScanExec, CometPlan, CometProjectExec, CometScanWrapper, CometSinkPlaceHolder, CometSparkToColumnarExec, NativeExecContext} import org.apache.spark.sql.comet.execution.arrow.CometArrowStream import org.apache.spark.sql.execution._ -import org.apache.spark.sql.execution.adaptive.{QueryStageExec, ShuffleQueryStageExec} +import org.apache.spark.sql.execution.adaptive.ShuffleQueryStageExec import org.apache.spark.sql.execution.exchange.{ENSURE_REQUIREMENTS, Exchange, ShuffleExchangeExec, ShuffleExchangeLike, ShuffleOrigin} import org.apache.spark.sql.execution.metric.{SQLMetric, SQLMetrics, SQLShuffleReadMetricsReporter, SQLShuffleWriteMetricsReporter} import org.apache.spark.sql.internal.SQLConf @@ -524,24 +524,14 @@ object CometShuffleExchangeExec None } - /** Whether `dt` is, or contains, a decimal that native shuffle does not hash as Spark does. */ - private def hasWideDecimal(dt: DataType): Boolean = dt match { - case d: DecimalType => d.precision > 18 - case StructType(fields) => fields.exists(f => hasWideDecimal(f.dataType)) - case ArrayType(elementType, _) => hasWideDecimal(elementType) - case MapType(keyType, valueType, _) => hasWideDecimal(keyType) || hasWideDecimal(valueType) - case _ => false - } - /** * Whether the stage feeding a shuffle starts at a typed Dataset conversion (see * [[CometConf.COMET_CONVERT_FROM_TYPED_DATASET_ENABLED]]). `CometExecRule` decides the shuffle * before it removes its placeholders, so the conversion is still inside its `CometScanWrapper`. */ private def readsTypedDatasetConversion(plan: SparkPlan): Boolean = plan match { - case _: Exchange | _: QueryStageExec => false - case CometScanWrapper(_, conversion: CometSparkToColumnarExec) => - conversion.child.isInstanceOf[SerializeFromObjectExec] + case _: Exchange => false + case CometScanWrapper(_, wrapped) => readsTypedDatasetConversion(wrapped) case conversion: CometSparkToColumnarExec => conversion.child.isInstanceOf[SerializeFromObjectExec] case other => other.children.exists(readsTypedDatasetConversion) @@ -647,11 +637,13 @@ object CometShuffleExchangeExec } // A typed Dataset conversion moves the shuffle above it from Comet's columnar shuffle, // which partitions with Spark's hash, to native shuffle. Native shuffle hashes a decimal - // wider than 18 digits differently from Spark (#5994), so a join with an input that is - // still on the columnar shuffle would put matching keys in different partitions. Leave - // such a shuffle where it was. A single partition hashes nothing. - if (partitioning.numPartitions > 1 && - expressions.exists(e => hasWideDecimal(e.dataType)) && + // wider than 18 digits differently from Spark, so a join with an input that is still on + // the columnar shuffle would put matching keys in different partitions. Leave such a + // shuffle where it was. A single partition hashes nothing. + // TODO: remove once native hashing matches Spark for wide decimals (#5994). + if (CometConf.COMET_CONVERT_FROM_TYPED_DATASET_ENABLED.get(conf) && reasons.isEmpty && + partitioning.numPartitions > 1 && + expressions.exists(_.dataType.existsRecursively(DecimalType.isByteArrayDecimalType)) && readsTypedDatasetConversion(s.child)) { reasons += "a shuffle above a typed Dataset conversion that hashes a decimal wider " + "than 18 digits stays on Comet's columnar shuffle, which hashes it as Spark does" diff --git a/spark/src/test/scala/org/apache/comet/exec/CometTypedDatasetSuite.scala b/spark/src/test/scala/org/apache/comet/exec/CometTypedDatasetSuite.scala index 08f9cbfa8e2..8df48186fea 100644 --- a/spark/src/test/scala/org/apache/comet/exec/CometTypedDatasetSuite.scala +++ b/spark/src/test/scala/org/apache/comet/exec/CometTypedDatasetSuite.scala @@ -24,8 +24,7 @@ import java.util.concurrent.atomic.AtomicLong import org.apache.spark.sql.{CometTestBase, DataFrame, Dataset} import org.apache.spark.sql.catalyst.expressions.aggregate.Partial import org.apache.spark.sql.comet.{CometBroadcastHashJoinExec, CometHashAggregateExec, CometSparkToColumnarExec} -import org.apache.spark.sql.comet.execution.shuffle.{CometColumnarShuffle, CometShuffleExchangeExec} -import org.apache.spark.sql.execution.{ColumnarToRowExec, ColumnarToRowTransition, InputAdapter, SerializeFromObjectExec, SparkPlan, WholeStageCodegenExec} +import org.apache.spark.sql.execution.{ColumnarToRowExec, ColumnarToRowTransition, SerializeFromObjectExec, SparkPlan, WholeStageCodegenExec} import org.apache.spark.sql.functions.{broadcast, col, size, sum} import org.apache.spark.sql.internal.SQLConf @@ -76,15 +75,11 @@ class CometTypedDatasetSuite extends CometTestBase { * right answer slowly, so only the plan shows it. */ private def assertRowOperatorsReadThroughTransitions(plan: SparkPlan): Unit = { - def unwrap(p: SparkPlan): SparkPlan = p match { - case InputAdapter(child) => child - case other => other - } + // `InputAdapter` and `WholeStageCodegenExec` report their child's `supportsColumnar`. val bare = collectWithSubqueries(plan) { case p if !p.supportsColumnar && !p.isInstanceOf[ColumnarToRowTransition] && - !p.isInstanceOf[WholeStageCodegenExec] && - p.children.exists(c => unwrap(c).supportsColumnar) => + p.children.exists(_.supportsColumnar) => p } assert(bare.isEmpty, s"row operators read a columnar child without a transition:\n$plan") @@ -97,7 +92,7 @@ class CometTypedDatasetSuite extends CometTestBase { * Checks the answer, that everything above the typed operation runs natively, and that the * operation's output is converted right above its `SerializeFromObjectExec`. */ - private def checkConverted(df: => DataFrame): SparkPlan = { + private def checkConverted(df: DataFrame): SparkPlan = { val (_, plan) = checkSparkAnswerAndOperator(df, includeClasses = Seq(classOf[CometSparkToColumnarExec])) conversions(plan).foreach { c => @@ -123,9 +118,9 @@ class CometTypedDatasetSuite extends CometTestBase { } convertTest("the aggregate and the shuffle above a map run natively") { - Seq("true", "false").foreach { aqe => - withSQLConf(SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> aqe) { - withRecs() { ds => + withRecs() { ds => + Seq("true", "false").foreach { aqe => + withSQLConf(SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> aqe) { val plan = checkConverted(ds.map(r => TypedDsRec(r.a + 1, r.b)).groupBy("b").count()) assert(nativePartialAggregates(plan).nonEmpty, s"AQE $aqe:\n$plan") } @@ -231,11 +226,11 @@ class CometTypedDatasetSuite extends CometTestBase { .range(0, 100, 1, 2) .map(i => TypedDsDecimalInts(new java.math.BigDecimal(i.longValue), Seq(i.intValue))) .alias("r") - val (_, plan) = checkSparkAnswer( - left.join(right, col("l.k") === col("r.k")).select(col("l.v"), col("r.xs"))) + val df = left.join(right, col("l.k") === col("r.k")).select(col("l.v"), col("r.xs")) + val (_, plan) = checkSparkAnswer(df) assert(conversions(plan).nonEmpty, plan) - val shuffles = collectWithSubqueries(plan) { case s: CometShuffleExchangeExec => s } - assert(shuffles.nonEmpty && shuffles.forall(_.shuffleType == CometColumnarShuffle), plan) + // One columnar shuffle for each input of the join. + checkCometExchange(df, 2, native = false) } } diff --git a/spark/src/test/scala/org/apache/spark/sql/benchmark/CometTypedDatasetBenchmark.scala b/spark/src/test/scala/org/apache/spark/sql/benchmark/CometTypedDatasetBenchmark.scala index 3686fbc328f..bfd2b1a063f 100644 --- a/spark/src/test/scala/org/apache/spark/sql/benchmark/CometTypedDatasetBenchmark.scala +++ b/spark/src/test/scala/org/apache/spark/sql/benchmark/CometTypedDatasetBenchmark.scala @@ -20,7 +20,7 @@ package org.apache.spark.sql.benchmark import org.apache.spark.benchmark.Benchmark -import org.apache.spark.sql.{DataFrame, Dataset, Encoder, Encoders, Row} +import org.apache.spark.sql.{DataFrame, Dataset, Row} import org.apache.spark.sql.catalyst.expressions.aggregate.Partial import org.apache.spark.sql.comet.{CometHashAggregateExec, CometPlan, CometSparkToColumnarExec} import org.apache.spark.sql.execution.SparkPlan @@ -58,12 +58,7 @@ object CometTypedDatasetBenchmark extends CometBenchmarkBase { private val loKeys = 100 private val hiKeys = 1024 * 1024 - // Declared rather than imported from `spark.implicits`, which would start the session while - // this object initializes. - private implicit val recEncoder: Encoder[TypedDatasetBenchRec] = - Encoders.product[TypedDatasetBenchRec] - private implicit val wideEncoder: Encoder[TypedDatasetBenchWide] = - Encoders.product[TypedDatasetBenchWide] + import spark.implicits._ private case class Arm(name: String, confs: Seq[(String, String)], check: SparkPlan => Unit) From 2737816792153ea2bf2161f48b699dd253a55087 Mon Sep 17 00:00:00 2001 From: Andy Grove Date: Sun, 4 Oct 2026 11:28:38 -0600 Subject: [PATCH 4/6] fix: preserve typed Dataset limit short-circuiting --- .../apache/comet/rules/CometExecRule.scala | 36 ++++++++++++++++++- .../comet/exec/CometTypedDatasetSuite.scala | 21 +++++++++++ 2 files changed, 56 insertions(+), 1 deletion(-) diff --git a/spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala b/spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala index 5eb97b2b123..65680a934d3 100644 --- a/spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala +++ b/spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala @@ -162,6 +162,14 @@ object CometExecRule { */ val SKIP_COMET_BROADCAST_TAG: org.apache.spark.sql.catalyst.trees.TreeNodeTag[Unit] = org.apache.spark.sql.catalyst.trees.TreeNodeTag[Unit]("comet.skipCometBroadcast") + + /** + * Tag set on a `SerializeFromObjectExec` below a physical limit. Converting its output to Arrow + * fills batches before the limit consumes them, which can evaluate typed Dataset user code for + * rows that Spark's row-based limit would never request. + */ + private val SKIP_TYPED_DATASET_CONVERSION_UNDER_LIMIT: TreeNodeTag[Unit] = + TreeNodeTag[Unit]("comet.skipTypedDatasetConversionUnderLimit") } /** @@ -369,6 +377,23 @@ case class CometExecRule(session: SparkSession) */ // spotless:on private def transform(plan: SparkPlan): SparkPlan = { + def tagTypedDatasetOutputsUnderLimit(limitChild: SparkPlan): Unit = + limitChild.foreach { + case op: SerializeFromObjectExec => + op.setTagValue(CometExecRule.SKIP_TYPED_DATASET_CONVERSION_UNDER_LIMIT, ()) + case _ => + } + + // Conversion is bottom-up, so mark typed serializers whose limit ancestors have not been + // visited yet. TreeNode tags survive the child copies made during transformUp, while an + // identity set would not. + plan.foreach { + case limit: CollectLimitExec => tagTypedDatasetOutputsUnderLimit(limit.child) + case limit: LocalLimitExec => tagTypedDatasetOutputsUnderLimit(limit.child) + case limit: GlobalLimitExec => tagTypedDatasetOutputsUnderLimit(limit.child) + case _ => + } + def convertNode(op: SparkPlan): SparkPlan = op match { // Scan marker produced by an optional, out-of-tree scan contrib (e.g. contrib/delta). // Matched by trait (no compile-time dependency on the contrib) and present only when that @@ -479,7 +504,16 @@ case class CometExecRule(session: SparkSession) // Arrow lets the operators above the typed operation run natively. case op: SerializeFromObjectExec if CometConf.COMET_CONVERT_FROM_TYPED_DATASET_ENABLED.get(conf) => - convertTypedDatasetOutput(op) + if (op + .getTagValue(CometExecRule.SKIP_TYPED_DATASET_CONVERSION_UNDER_LIMIT) + .isDefined) { + withFallbackReason( + op, + "Comet does not convert the output of a typed Dataset operation below a limit " + + "because Arrow batching could evaluate rows beyond Spark's row-level limit") + } else { + convertTypedDatasetOutput(op) + } // Spark 4.0+: replace only the per-task write, leaving DataWritingCommandExec - and // therefore Spark's commit protocol, stats trackers and SaveMode handling - in place. diff --git a/spark/src/test/scala/org/apache/comet/exec/CometTypedDatasetSuite.scala b/spark/src/test/scala/org/apache/comet/exec/CometTypedDatasetSuite.scala index 8df48186fea..d3673d17282 100644 --- a/spark/src/test/scala/org/apache/comet/exec/CometTypedDatasetSuite.scala +++ b/spark/src/test/scala/org/apache/comet/exec/CometTypedDatasetSuite.scala @@ -277,6 +277,27 @@ class CometTypedDatasetSuite extends CometTestBase { } } + convertTest("a limit does not evaluate typed Dataset rows beyond the result") { + Seq("true", "false").foreach { aqe => + withSQLConf(SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> aqe) { + val df = spark + .range(0, 100, 1, 1) + .map { i => + if (i == 30L) { + throw new IllegalArgumentException("unexpected evaluation of row 30") + } + i + 1L + } + .toDF() + .limit(1) + val (_, plan) = checkSparkAnswerAndFallbackReason( + df, + "Comet does not convert the output of a typed Dataset operation below a limit") + assert(conversions(plan).isEmpty, s"AQE $aqe:\n$plan") + } + } + } + test("off by default") { withRecs() { ds => val (_, plan) = checkSparkAnswer(ds.map(r => TypedDsRec(r.a + 1, r.b)).groupBy("b").count()) From c202f515888da0a6a81f526e4abe152ed4c724c7 Mon Sep 17 00:00:00 2001 From: Andy Grove Date: Sun, 4 Oct 2026 12:35:48 -0600 Subject: [PATCH 5/6] fix: leave typed Dataset output unconverted below any reader that stops early The limit guard covered only physical limits, so a mapPartitions function such as `_.take(1)`, or code reading Dataset.rdd, still read typed rows through an Arrow batch and ran the user function on rows that Spark never reaches. Walk down from each reader that can stop early (a limit, a top-k over sorted input, a MapPartitionsExec, and a DeserializeToObjectExec at the plan root) and stop at an operator that reads all of its input first, so a limit above an aggregate keeps the conversion. --- .../adding_a_new_operator.md | 14 ++- .../apache/comet/rules/CometExecRule.scala | 91 ++++++++++++------- .../comet/exec/CometTypedDatasetSuite.scala | 57 +++++++++--- 3 files changed, 114 insertions(+), 48 deletions(-) diff --git a/docs/source/contributor-guide/adding_a_new_operator.md b/docs/source/contributor-guide/adding_a_new_operator.md index ba51aa25c92..56a983ac762 100644 --- a/docs/source/contributor-guide/adding_a_new_operator.md +++ b/docs/source/contributor-guide/adding_a_new_operator.md @@ -144,12 +144,16 @@ operation ends in `SerializeFromObjectExec`, though, whose output is ordinary ro `spark.comet.convert.typedDataset.enabled`, `CometExecRule` puts a `CometSparkToColumnarExec` above it, so the operators above the typed operation can run natively. Spark inserts no columnar transitions below a `RowToColumnarTransition`, so the rule inserts them for the typed operation's -own operators itself. Fusing the deserializer, the `Invoke` that calls the user function, and the +own operators itself. Spark computes a typed operation's rows one at a time, as they are read, while +the conversion fills a whole Arrow batch first. So the rule leaves the output unconverted where a +limit, a `mapPartitions` function, or code reading `Dataset.rdd` could stop reading it early, unless +an operator that reads all of its input first, such as an exchange, a sort, or a hash aggregate, +sits in between. Fusing the deserializer, the `Invoke` that calls the user function, and the serializer of `Dataset.map` into one projection in the JVM codegen dispatcher was tried in -[#5714](https://github.com/apache/datafusion-comet/pull/5714) and dropped. The dispatcher only -calls into Spark's own classes, and the conversion gets nearly the same speedup for `map` while -also covering the operations that pass the user function an iterator or a whole group. A typed -`filter` is planned as an ordinary `FilterExec`, not as one of these operators. +[#5714](https://github.com/apache/datafusion-comet/pull/5714) and dropped. The dispatcher only calls +into Spark's own classes, and the conversion gets nearly the same speedup for `map` while also +covering the operations that pass the user function an iterator or a whole group. A typed `filter` +is planned as an ordinary `FilterExec`, not as one of these operators. **Driver-side commands.** `ExecutedCommandExec` runs a `RunnableCommand`, such as DDL or `SET`, on the driver, so there is no data path for Comet to accelerate. diff --git a/spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala b/spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala index 65680a934d3..1f8896539bc 100644 --- a/spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala +++ b/spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala @@ -23,7 +23,7 @@ import scala.collection.mutable.ListBuffer import scala.jdk.CollectionConverters._ import org.apache.spark.sql.SparkSession -import org.apache.spark.sql.catalyst.expressions.{Divide, DoubleLiteral, EqualNullSafe, EqualTo, Expression, FloatLiteral, GreaterThan, GreaterThanOrEqual, KnownFloatingPointNormalized, LessThan, LessThanOrEqual, NamedExpression, Remainder} +import org.apache.spark.sql.catalyst.expressions.{Divide, DoubleLiteral, EqualNullSafe, EqualTo, Expression, FloatLiteral, GreaterThan, GreaterThanOrEqual, KnownFloatingPointNormalized, LessThan, LessThanOrEqual, NamedExpression, Remainder, SortOrder} import org.apache.spark.sql.catalyst.expressions.aggregate.{AggregateMode, Final, Partial, PartialMerge} import org.apache.spark.sql.catalyst.optimizer.NormalizeNaNAndZero import org.apache.spark.sql.catalyst.rules.Rule @@ -47,7 +47,7 @@ import org.apache.spark.sql.execution.datasources.v2.{BatchScanExec, V2CommandEx import org.apache.spark.sql.execution.datasources.v2.csv.CSVScan import org.apache.spark.sql.execution.datasources.v2.json.JsonScan import org.apache.spark.sql.execution.datasources.v2.parquet.ParquetScan -import org.apache.spark.sql.execution.exchange.{BroadcastExchangeExec, BroadcastExchangeLike, ReusedExchangeExec, ShuffleExchangeExec, ShuffleExchangeLike} +import org.apache.spark.sql.execution.exchange.{BroadcastExchangeExec, BroadcastExchangeLike, Exchange, ReusedExchangeExec, ShuffleExchangeExec, ShuffleExchangeLike} import org.apache.spark.sql.execution.joins.{BroadcastHashJoinExec, BroadcastNestedLoopJoinExec, ShuffledHashJoinExec, SortMergeJoinExec} import org.apache.spark.sql.execution.window.WindowExec import org.apache.spark.sql.internal.SQLConf @@ -164,12 +164,11 @@ object CometExecRule { org.apache.spark.sql.catalyst.trees.TreeNodeTag[Unit]("comet.skipCometBroadcast") /** - * Tag set on a `SerializeFromObjectExec` below a physical limit. Converting its output to Arrow - * fills batches before the limit consumes them, which can evaluate typed Dataset user code for - * rows that Spark's row-based limit would never request. + * Tag set on a `SerializeFromObjectExec` whose output an operator above it can stop reading + * early, naming that operator. See `tagPartiallyReadTypedDatasetOutputs`. */ - private val SKIP_TYPED_DATASET_CONVERSION_UNDER_LIMIT: TreeNodeTag[Unit] = - TreeNodeTag[Unit]("comet.skipTypedDatasetConversionUnderLimit") + private val TYPED_DATASET_PARTIAL_READER: TreeNodeTag[String] = + TreeNodeTag[String]("comet.typedDatasetPartialReader") } /** @@ -377,21 +376,8 @@ case class CometExecRule(session: SparkSession) */ // spotless:on private def transform(plan: SparkPlan): SparkPlan = { - def tagTypedDatasetOutputsUnderLimit(limitChild: SparkPlan): Unit = - limitChild.foreach { - case op: SerializeFromObjectExec => - op.setTagValue(CometExecRule.SKIP_TYPED_DATASET_CONVERSION_UNDER_LIMIT, ()) - case _ => - } - - // Conversion is bottom-up, so mark typed serializers whose limit ancestors have not been - // visited yet. TreeNode tags survive the child copies made during transformUp, while an - // identity set would not. - plan.foreach { - case limit: CollectLimitExec => tagTypedDatasetOutputsUnderLimit(limit.child) - case limit: LocalLimitExec => tagTypedDatasetOutputsUnderLimit(limit.child) - case limit: GlobalLimitExec => tagTypedDatasetOutputsUnderLimit(limit.child) - case _ => + if (CometConf.COMET_CONVERT_FROM_TYPED_DATASET_ENABLED.get(conf)) { + tagPartiallyReadTypedDatasetOutputs(plan) } def convertNode(op: SparkPlan): SparkPlan = op match { @@ -504,15 +490,15 @@ case class CometExecRule(session: SparkSession) // Arrow lets the operators above the typed operation run natively. case op: SerializeFromObjectExec if CometConf.COMET_CONVERT_FROM_TYPED_DATASET_ENABLED.get(conf) => - if (op - .getTagValue(CometExecRule.SKIP_TYPED_DATASET_CONVERSION_UNDER_LIMIT) - .isDefined) { - withFallbackReason( - op, - "Comet does not convert the output of a typed Dataset operation below a limit " + - "because Arrow batching could evaluate rows beyond Spark's row-level limit") - } else { - convertTypedDatasetOutput(op) + op.getTagValue(CometExecRule.TYPED_DATASET_PARTIAL_READER) match { + case Some(reader) => + withFallbackReason( + op, + "Comet does not convert the output of a typed Dataset operation when " + + s"$reader can stop reading it early, because filling an Arrow batch would " + + "run the user function on rows that Spark never reaches") + case None => + convertTypedDatasetOutput(op) } // Spark 4.0+: replace only the per-task write, leaving DataWritingCommandExec - and @@ -1302,6 +1288,49 @@ case class CometExecRule(session: SparkSession) private def hasEnabledHandler(op: SparkPlan): Boolean = allExecs.get(op.getClass).exists(_.enabledConfig.forall(_.get(op.conf))) + /** + * Tags each `SerializeFromObjectExec` whose output an operator above it can stop reading early + * with that operator, so `transform` does not convert it. Spark computes the rows of a typed + * Dataset operation one at a time, as the operator above reads them, while the conversion fills + * a whole Arrow batch first. Below a limit, a `mapPartitions` function such as `_.take(1)`, or + * code reading `Dataset.rdd`, the conversion would run the user function on rows Spark never + * reaches, and a function that throws on one of them would fail a query that succeeds in Spark. + * An operator that reads all of its input before it returns a row ends the search, since Spark + * computes every row below it anyway: an exchange, a sort, a hash aggregate, or a top-k over + * input that is not already sorted. + * + * Conversion is bottom-up, so this runs first. TreeNode tags survive the child copies made + * during transformUp, while an identity set would not. + */ + private def tagPartiallyReadTypedDatasetOutputs(plan: SparkPlan): Unit = { + def visit(op: SparkPlan, partialReader: Option[String]): Unit = { + val childReader = op match { + case serialize: SerializeFromObjectExec => + partialReader.foreach( + serialize.setTagValue(CometExecRule.TYPED_DATASET_PARTIAL_READER, _)) + partialReader + case _: CollectLimitExec | _: LocalLimitExec | _: GlobalLimitExec => Some("a limit") + // A top-k reads only its first rows when its input is already sorted. + case topK: TakeOrderedAndProjectExec + if SortOrder.orderingSatisfies(topK.child.outputOrdering, topK.sortOrder) => + Some("a limit") + case _: MapPartitionsExec => Some("a mapPartitions function") + case _: Exchange | _: SortExec | _: HashAggregateExec | _: ObjectHashAggregateExec | + _: TakeOrderedAndProjectExec => + None + case _ => partialReader + } + op.children.foreach(visit(_, childReader)) + } + // `Dataset.rdd` plans a `DeserializeToObjectExec` at the root, and the RDD's own code + // decides how much of it to read, as `take(1)` does. + val rootReader = plan match { + case _: DeserializeToObjectExec => Some("code reading Dataset.rdd") + case _ => None + } + visit(plan, rootReader) + } + /** * Converts the rows a typed Dataset operation produces to Arrow, so the operators above it can * run natively. See [[CometConf.COMET_CONVERT_FROM_TYPED_DATASET_ENABLED]]. diff --git a/spark/src/test/scala/org/apache/comet/exec/CometTypedDatasetSuite.scala b/spark/src/test/scala/org/apache/comet/exec/CometTypedDatasetSuite.scala index d3673d17282..d61de020bd6 100644 --- a/spark/src/test/scala/org/apache/comet/exec/CometTypedDatasetSuite.scala +++ b/spark/src/test/scala/org/apache/comet/exec/CometTypedDatasetSuite.scala @@ -277,27 +277,60 @@ class CometTypedDatasetSuite extends CometTestBase { } } + /** + * One partition whose user function throws on row 30. Spark computes typed rows as they are + * read, so a reader that stops before row 30 never reaches it, but an Arrow batch would. + */ + private def failsOnRow30: Dataset[Long] = + spark + .range(0, 100, 1, 1) + .map { i => + if (i == 30L) { + throw new IllegalArgumentException("unexpected evaluation of row 30") + } + i + 1L + } + convertTest("a limit does not evaluate typed Dataset rows beyond the result") { Seq("true", "false").foreach { aqe => withSQLConf(SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> aqe) { - val df = spark - .range(0, 100, 1, 1) - .map { i => - if (i == 30L) { - throw new IllegalArgumentException("unexpected evaluation of row 30") - } - i + 1L - } - .toDF() - .limit(1) val (_, plan) = checkSparkAnswerAndFallbackReason( - df, - "Comet does not convert the output of a typed Dataset operation below a limit") + failsOnRow30.toDF().limit(1), + "Comet does not convert the output of a typed Dataset operation when a limit can " + + "stop reading it early") assert(conversions(plan).isEmpty, s"AQE $aqe:\n$plan") } } } + convertTest("a mapPartitions function does not evaluate typed Dataset rows it never reads") { + Seq("true", "false").foreach { aqe => + withSQLConf(SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> aqe) { + // The filter between the two typed operations is what the conversion would run natively. + val (_, plan) = checkSparkAnswerAndFallbackReason( + failsOnRow30.filter(col("value") > 0L).mapPartitions(_.take(1)).toDF(), + "Comet does not convert the output of a typed Dataset operation when a mapPartitions " + + "function can stop reading it early") + assert(conversions(plan).isEmpty, s"AQE $aqe:\n$plan") + } + } + } + + convertTest("code reading Dataset.rdd does not evaluate typed Dataset rows it never reads") { + assert(failsOnRow30.filter(col("value") > 0L).rdd.take(1).toSeq == Seq(1L)) + } + + convertTest("a limit above an aggregate keeps the conversion the aggregate reads all of") { + withRecs() { ds => + Seq("true", "false").foreach { aqe => + withSQLConf(SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> aqe) { + // There are 13 groups, so the limit keeps all of them in whatever order they come. + checkConverted(ds.map(r => TypedDsRec(r.a + 1, r.b)).groupBy("b").count().limit(20)) + } + } + } + } + test("off by default") { withRecs() { ds => val (_, plan) = checkSparkAnswer(ds.map(r => TypedDsRec(r.a + 1, r.b)).groupBy("b").count()) From 1e238811fe3c04d61c940c9d06fcc5ac708925a1 Mon Sep 17 00:00:00 2001 From: Andy Grove Date: Mon, 5 Oct 2026 07:55:06 -0600 Subject: [PATCH 6/6] fix: recognize Dataset.rdd plans whose deserializer Spark removed The guard for code reading Dataset.rdd looked for the DeserializeToObjectExec that Dataset.rdd puts at the root of the plan. When the Dataset ends in a typed operation such as map, Spark's EliminateSerialization drops that deserializer together with the operation's serializer, so the root is the operation itself, or a typed filter over it. The output of an earlier typed operation was then still converted, and map(f).filter(...).map(g).rdd.take(1) ran f on rows that Spark never reaches. The guard now takes any root that produces objects, under a filter or a project, as code reading Dataset.rdd. A Dataset's own plan ends in rows, so no other plan has such a root. --- .../org/apache/comet/rules/CometExecRule.scala | 17 +++++++++++------ .../comet/exec/CometTypedDatasetSuite.scala | 11 ++++++++++- 2 files changed, 21 insertions(+), 7 deletions(-) diff --git a/spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala b/spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala index fb3ff0e848c..33c1930ddc7 100644 --- a/spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala +++ b/spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala @@ -1340,13 +1340,18 @@ case class CometExecRule(session: SparkSession) } op.children.foreach(visit(_, childReader)) } - // `Dataset.rdd` plans a `DeserializeToObjectExec` at the root, and the RDD's own code - // decides how much of it to read, as `take(1)` does. - val rootReader = plan match { - case _: DeserializeToObjectExec => Some("code reading Dataset.rdd") - case _ => None + // `Dataset.rdd` reads the objects the plan produces, and the RDD's own code decides how many + // of them to read, as `take(1)` does. The plan's root is the `DeserializeToObjectExec` that + // `Dataset.rdd` adds. When the Dataset ends in a typed operation such as `map`, Spark's + // `EliminateSerialization` drops that deserializer together with the operation's serializer, + // so the root is the operation itself, which produces objects too, or a typed filter or a + // project over it. A Dataset's own plan ends in rows, so no other plan has such a root. + def producesObjects(op: SparkPlan): Boolean = op match { + case _: ObjectProducerExec => true + case _: FilterExec | _: ProjectExec => producesObjects(op.children.head) + case _ => false } - visit(plan, rootReader) + visit(plan, if (producesObjects(plan)) Some("code reading Dataset.rdd") else None) } /** diff --git a/spark/src/test/scala/org/apache/comet/exec/CometTypedDatasetSuite.scala b/spark/src/test/scala/org/apache/comet/exec/CometTypedDatasetSuite.scala index d61de020bd6..9b809e24e8d 100644 --- a/spark/src/test/scala/org/apache/comet/exec/CometTypedDatasetSuite.scala +++ b/spark/src/test/scala/org/apache/comet/exec/CometTypedDatasetSuite.scala @@ -317,7 +317,16 @@ class CometTypedDatasetSuite extends CometTestBase { } convertTest("code reading Dataset.rdd does not evaluate typed Dataset rows it never reads") { - assert(failsOnRow30.filter(col("value") > 0L).rdd.take(1).toSeq == Seq(1L)) + Seq("true", "false").foreach { aqe => + withSQLConf(SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> aqe) { + val filtered = failsOnRow30.filter(col("value") > 0L) + assert(filtered.rdd.take(1).toSeq == Seq(1L)) + // When the Dataset ends in a typed operation, Spark drops the deserializer that `rdd` + // adds, so the root of the plan is that operation, or a typed filter over it. + assert(filtered.map(_ + 1L).rdd.take(1).toSeq == Seq(2L)) + assert(filtered.map(_ + 1L).filter((v: Long) => v > 0L).rdd.take(1).toSeq == Seq(2L)) + } + } } convertTest("a limit above an aggregate keeps the conversion the aggregate reads all of") {