From 3563e191a70db0e9b11b82ad0346d2c7a7334674 Mon Sep 17 00:00:00 2001 From: Andy Grove Date: Sat, 5 Sep 2026 10:24:26 -0600 Subject: [PATCH 1/6] feat: fuse the typed Dataset map sandwich into a Comet projection Collapses the SerializeFromObject / MapElements / DeserializeToObject island that `ds.map(f)` produces into a projection over the deserializer's child, so a typed map no longer forces a Spark fallback in the middle of an otherwise native plan. The projection converts through the ordinary CometProjectExec path and the fused tree routes through the JVM codegen dispatcher, so there is no proto or native change. Neither object operator can run natively on its own: one outputs and the other consumes a raw JVM object reference, which has no Arrow representation. Fusing works because CometBatchKernelCodegen.canHandle only type-checks the root and the bound references, so the object may exist strictly inside the tree. Spark already builds the fused expression -- the rule mirrors MapElementsExec.doConsume and the two doConsume methods whole-stage codegen chains around it. With more than one output column the rule emits two stacked projections: an inner one producing a single CreateNamedStruct, and an outer one extracting the fields with GetStructField. That is a correctness requirement, not tidiness: N separate projection expressions would each carry their own copy of the closure Invoke and each become its own kernel, calling user code N times per row where Spark calls it once. The struct is tagged FORCE_DISPATCH so it compiles into one kernel, and subexpression elimination inside that kernel collapses the shared Invoke back to a single call. Off by default. The rewrite pays off when the typed operation sits between native operators and can cost a little when it is at the top of the plan, so flipping the default needs benchmarks first. Part of #5572. Closes #5710. --- .github/workflows/pr_build_linux.yml | 1 + .github/workflows/pr_build_macos.yml | 1 + docs/source/user-guide/latest/operators.md | 17 + .../scala/org/apache/comet/CometConf.scala | 14 + .../apache/comet/rules/CometExecRule.scala | 17 +- .../comet/rules/RewriteTypedDatasetMap.scala | 243 +++++++++++++ .../apache/comet/serde/CometScalaUDF.scala | 19 + .../apache/comet/serde/QueryPlanSerde.scala | 5 + .../apache/comet/CometTypedDatasetSuite.scala | 342 ++++++++++++++++++ 9 files changed, 657 insertions(+), 2 deletions(-) create mode 100644 spark/src/main/scala/org/apache/comet/rules/RewriteTypedDatasetMap.scala create mode 100644 spark/src/test/scala/org/apache/comet/CometTypedDatasetSuite.scala diff --git a/.github/workflows/pr_build_linux.yml b/.github/workflows/pr_build_linux.yml index 313fc705492..3e7cf5fd542 100644 --- a/.github/workflows/pr_build_linux.yml +++ b/.github/workflows/pr_build_linux.yml @@ -471,6 +471,7 @@ jobs: org.apache.comet.CometUuidExpressionSuite org.apache.comet.serde.CometScalarFunctionSuite org.apache.comet.CometFallbackInvarianceSuite + org.apache.comet.CometTypedDatasetSuite fail-fast: false name: ${{ matrix.profile.name }} [${{ matrix.suite.name }}] runs-on: ubuntu-24.04 diff --git a/.github/workflows/pr_build_macos.yml b/.github/workflows/pr_build_macos.yml index a7df2b50ab8..f82d20a1df2 100644 --- a/.github/workflows/pr_build_macos.yml +++ b/.github/workflows/pr_build_macos.yml @@ -244,6 +244,7 @@ jobs: org.apache.comet.CometUuidExpressionSuite org.apache.comet.serde.CometScalarFunctionSuite org.apache.comet.CometFallbackInvarianceSuite + org.apache.comet.CometTypedDatasetSuite fail-fast: false name: ${{ matrix.os }}/${{ matrix.profile.name }} [${{ matrix.suite.name }}] diff --git a/docs/source/user-guide/latest/operators.md b/docs/source/user-guide/latest/operators.md index 3a18d86606d..65566ffff88 100644 --- a/docs/source/user-guide/latest/operators.md +++ b/docs/source/user-guide/latest/operators.md @@ -120,6 +120,23 @@ omitted from the tables below and may be reconsidered based on demand: | ------------------------ | ------ | ----------------------------------------------------------------- | | `DataWritingCommandExec` | ⚠️ | Experimental native Parquet writes, disabled by default (opt-in). | +## Typed Dataset operations + +Typed `Dataset` operations bracket a JVM object in a `DeserializeToObject` / `SerializeFromObject` +pair. Neither operator can run natively on its own, because its input or output is a raw JVM object +reference that has no Arrow representation. What Comet can do is fuse a whole sandwich back into a +projection when everything between the pair is per-row, so the typed operation no longer forces a +Spark fallback island in the middle of an otherwise native plan. The user closure still runs on the +JVM, once per row, inside Comet's codegen dispatcher, so results match Spark exactly. + +| Operator | Status | Notes | +| ---------------------------------------------------------------------------- | ------ | ------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------- | +| `SerializeFromObjectExec` + `MapElementsExec` + `DeserializeToObjectExec` | ⚠️ | The `ds.map(f)` sandwich is fused into a projection when `spark.comet.exec.typedDatasetMap.enabled=true` (off by default). Worth enabling when the typed operation sits between native operators ([#5710](https://github.com/apache/datafusion-comet/issues/5710)). | +| `MapPartitionsExec`, `FlatMapGroupsExec`, `CoGroupExec`, `AppendColumnsExec` | 🔜 | These consume iterators or groups rather than rows, so there is no per-row expression to fuse. `AppendColumnsExec` is per-row but widens the schema and is not handled yet. | + +Typed filters (`ds.filter(func)`) do not produce this sandwich at all: Catalyst lowers them to an +ordinary `FilterExec` whose condition is an `Invoke` of the closure, which Comet already dispatches. + ## Python and UDF | Operator | Status | Notes | diff --git a/spark/src/main/scala/org/apache/comet/CometConf.scala b/spark/src/main/scala/org/apache/comet/CometConf.scala index f4aeacc8478..83915ea73e3 100644 --- a/spark/src/main/scala/org/apache/comet/CometConf.scala +++ b/spark/src/main/scala/org/apache/comet/CometConf.scala @@ -396,6 +396,20 @@ object CometConf extends ShimCometConf { .booleanConf .createWithDefault(true) + val COMET_EXEC_TYPED_DATASET_MAP_ENABLED: ConfigEntry[Boolean] = + conf("spark.comet.exec.typedDatasetMap.enabled") + .category(CATEGORY_EXEC) + .doc("Experimental. Whether to fuse the `SerializeFromObject` / `MapElements` / " + + "`DeserializeToObject` operator sandwich that a typed `Dataset.map` produces into a " + + "Comet projection, so the typed operation no longer forces a Spark fallback island in " + + "the middle of an otherwise native plan. The user closure still runs on the JVM, once " + + "per row, inside Comet's Arrow-direct codegen dispatcher, so results match Spark. Off " + + "by default: the rewrite only pays off when the typed operation sits between native " + + "operators, and can cost a little when it is at the top of the plan. Requires " + + s"${COMET_SCALA_UDF_CODEGEN_ENABLED.key}=true.") + .booleanConf + .createWithDefault(false) + val COMET_SHUFFLE_NATIVE_HASH_PARTITIONING_ENABLED: ConfigEntry[Boolean] = conf("spark.comet.shuffle.native.partitioning.hash.enabled") .withAlternative("spark.comet.native.shuffle.partitioning.hash.enabled") 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 1d2e48e52b4..f45aa1cf6ca 100644 --- a/spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala +++ b/spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala @@ -707,13 +707,26 @@ case class CometExecRule(session: SparkSession) normalizedPlan } + // Collapse the SerializeFromObject / MapElements / DeserializeToObject sandwich that a typed + // `Dataset.map` produces into a projection, before the bottom-up conversion sees the + // individual operators. `transform()` visits children first, so by the time it reached + // `DeserializeToObjectExec` the island would already have fallen back. + val planWithTypedMapFused = + if (CometConf.COMET_EXEC_TYPED_DATASET_MAP_ENABLED.get(conf)) { + planWithJoinRewritten.transformUp { case p => + RewriteTypedDatasetMap.rewrite(p) + } + } else { + planWithJoinRewritten + } + // Tag Partial aggregates that must not be converted to Comet because a // corresponding Final or PartialMerge cannot be converted and the intermediate buffer // formats are incompatible. This runs before transform() so the tags are checked // during the bottom-up conversion. Tags persist through AQE stage creation. - tagUnsafePartialAggregates(planWithJoinRewritten) + tagUnsafePartialAggregates(planWithTypedMapFused) - var newPlan = transform(planWithJoinRewritten) + var newPlan = transform(planWithTypedMapFused) // if the plan cannot be run fully natively then explain why (when appropriate // config is enabled) diff --git a/spark/src/main/scala/org/apache/comet/rules/RewriteTypedDatasetMap.scala b/spark/src/main/scala/org/apache/comet/rules/RewriteTypedDatasetMap.scala new file mode 100644 index 00000000000..14095c66aeb --- /dev/null +++ b/spark/src/main/scala/org/apache/comet/rules/RewriteTypedDatasetMap.scala @@ -0,0 +1,243 @@ +/* + * 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.api.java.function.MapFunction +import org.apache.spark.internal.Logging +import org.apache.spark.sql.catalyst.expressions.{Alias, AttributeReference, AttributeSeq, AttributeSet, BindReferences, BoundReference, CreateNamedStruct, Expression, GetStructField, Literal, NamedExpression} +import org.apache.spark.sql.catalyst.expressions.objects.Invoke +import org.apache.spark.sql.catalyst.plans.logical.FunctionUtils +import org.apache.spark.sql.execution.{DeserializeToObjectExec, MapElementsExec, ProjectExec, SerializeFromObjectExec, SparkPlan} +import org.apache.spark.sql.internal.SQLConf +import org.apache.spark.sql.types.ObjectType + +import org.apache.comet.CometConf +import org.apache.comet.CometSparkSessionExtensions.withFallbackReason +import org.apache.comet.codegen.CometBatchKernelCodegen +import org.apache.comet.serde.CometScalaUDF + +/** + * Collapses the three-operator island that a typed `Dataset.map` produces + * + * {{{ + * SerializeFromObject [serializer...] + * +- MapElements , obj: T + * +- DeserializeToObject , obj: T + * +- child + * }}} + * + * into a projection over `child` whose expressions are the fused + * serializer(closure(deserializer(child))) trees. The projection then converts through the + * ordinary `CometProjectExec` path, and the fused tree routes through the JVM codegen dispatcher, + * so the whole thing stays inside the Comet pipeline. No proto or native change is involved. + * + * '''Why the operators cannot be handled individually.''' `DeserializeToObjectExec` outputs a + * single `ObjectType` attribute and `SerializeFromObjectExec` consumes one. `ObjectType` is + * outside `QueryPlanSerde.supportedDataType` and outside + * `CometBatchKernelCodegen.isSupportedDataType`, because a JVM object reference cannot live in an + * Arrow vector. Fusing works because `CometBatchKernelCodegen.canHandle` only type-checks the + * root and the bound references, so the object may exist strictly ''inside'' the tree. See + * https://github.com/apache/datafusion-comet/issues/5710. + * + * '''Why the expression is safe to rebuild.''' Spark already constructs it. This rule mirrors + * `MapElementsExec.doConsume` (the `Invoke` on the closure literal), + * `DeserializeToObjectExec.doConsume` and `SerializeFromObjectExec.doConsume`, which whole-stage + * codegen chains to produce exactly the same fused tree. + * + * '''Shape of the output.''' With one output column the projection is a single fused expression. + * With N > 1 the rule emits two stacked projections: an inner one producing a single + * `CreateNamedStruct` column, and an outer one extracting the N fields with `GetStructField`. The + * struct matters for correctness, not tidiness: N separate projection expressions would each + * carry their own copy of the closure `Invoke` and each become its own dispatch kernel, calling + * the user closure N times per row where Spark calls it once. The inner struct is tagged + * `CometScalaUDF.FORCE_DISPATCH` so it compiles into one kernel, and subexpression elimination + * inside that kernel collapses the shared `Invoke` back to a single call per row. + * + * Only `MapElementsExec` is fusable. `MapPartitionsExec`, `FlatMapGroupsExec` and `CoGroupExec` + * consume iterators or groups rather than rows, so no per-row expression exists for them. A chain + * of adjacent `MapElementsExec` nodes is fused as a whole, because `ds.map(f).map(g)` leaves both + * under a single Serialize/Deserialize pair. + */ +object RewriteTypedDatasetMap extends Logging { + + /** Name of the single intermediate struct column when there is more than one output column. */ + private val FUSED_COLUMN = "comet_fused_object" + + def rewrite(plan: SparkPlan): SparkPlan = plan match { + case serialize: SerializeFromObjectExec => + // `ds.map(f).map(g)` leaves two adjacent `MapElements` under one Serialize/Deserialize pair, + // so walk the whole chain rather than expecting exactly one. + val chain = mapElementsChain(serialize.child) + chain.lastOption.map(_.child) match { + case Some(deserialize: DeserializeToObjectExec) => + fuse(serialize, chain, deserialize).getOrElse(plan) + case _ => plan + } + case _ => plan + } + + /** Consecutive `MapElementsExec` nodes, outermost first. */ + private def mapElementsChain(plan: SparkPlan): Seq[MapElementsExec] = plan match { + case m: MapElementsExec => m +: mapElementsChain(m.child) + case _ => Nil + } + + private def fuse( + serialize: SerializeFromObjectExec, + chain: Seq[MapElementsExec], + deserialize: DeserializeToObjectExec): Option[SparkPlan] = { + val child = deserialize.child + val serializer = serialize.serializer + val objType = chain.head.outputObjectType + + // Every serializer element is an `Alias` for encoder-generated serializers. The outer + // projection rebuilds them with `withNewChildren`, which needs exactly one child, and that + // is also what preserves `exprId` / qualifier / metadata across the rewrite. + if (!serializer.forall(_.isInstanceOf[Alias])) { + return declineQuietly( + serialize, + "serializer contains a non-Alias element: " + + serializer + .filterNot(_.isInstanceOf[Alias]) + .map(_.getClass.getSimpleName) + .mkString(", ")) + } + + // The serializer reads the object through `BoundReference(0, objType)`. Anything else means a + // shape this rule has not been reasoned about, so leave it alone rather than guess. + val badRefs = serializer.flatMap(_.collect { + case b: BoundReference if b.ordinal != 0 || b.dataType != objType => b + }) + if (badRefs.nonEmpty) { + return declineQuietly( + serialize, + s"serializer reads unexpected bound references: ${badRefs.mkString(", ")}") + } + + // Compose the chain innermost-first, so `ds.map(f).map(g)` becomes g(f(deserializer)). Each + // step mirrors `MapElementsExec.doConsume`, including how it picks the specialized + // `Function1.apply$mc..$sp` name from the operator's own input and output object types. + val callFunc = chain.reverse.foldLeft(deserialize.deserializer) { (arg, m) => + val (funcClass, funcName) = m.func match { + case _: MapFunction[_, _] => classOf[MapFunction[_, _]] -> "call" + case _ => + FunctionUtils.getFunctionOneName(m.outputObjectType, m.child.output.head.dataType) + } + Invoke( + Literal.create(m.func, ObjectType(funcClass)), + funcName, + m.outputObjectType, + arg :: Nil, + propagateNull = false) + } + + // Substitute the closure call for the object the serializer reads. `transform` on an `Alias` + // preserves `exprId`, qualifier and metadata via `otherCopyArgs`, so the rewritten projection + // keeps `SerializeFromObjectExec`'s exact output attributes and parents stay valid. + val fused = serializer.map { ne => + ne.transform { case _: BoundReference => callFunc }.asInstanceOf[NamedExpression] + } + + // A projection may only reference its child's output. The deserializer reads `child.output` + // and the object reference is gone, so this should hold; check rather than assume. + val childOutput = AttributeSet(child.output) + val dangling = fused.flatMap(f => (f.references -- childOutput).toSeq) + if (dangling.nonEmpty) { + return declineQuietly( + serialize, + s"fused expression references attributes outside the child: ${dangling.mkString(", ")}") + } + + if (fused.length == 1) { + // Single output column: the fused expression is already one dispatch kernel, so there is + // nothing for the struct wrapper to deduplicate. + val single = fused.head + val value = single.children.head + dispatchable(value).map { _ => + value.setTagValue(CometScalaUDF.FORCE_DISPATCH, ()) + ProjectExec(fused, child) + } + } else { + // Without subexpression elimination the single kernel would still evaluate the shared + // `Invoke` once per struct field, so the closure would run N times per row. Spark runs it + // once; decline rather than change how many times user code executes. + if (!SQLConf.get.subexpressionEliminationEnabled) { + return declineQuietly( + serialize, + s"${serializer.length} output columns require subexpression elimination to keep the " + + "closure to one call per row, but " + + s"${SQLConf.SUBEXPRESSION_ELIMINATION_ENABLED.key}=false") + } + + val structExpr = + CreateNamedStruct(fused.flatMap(ne => Seq(Literal(ne.name), ne.children.head))) + dispatchable(structExpr).map { _ => + structExpr.setTagValue(CometScalaUDF.FORCE_DISPATCH, ()) + val structAlias = Alias(structExpr, FUSED_COLUMN)() + val inner = ProjectExec(Seq(structAlias), child) + val structAttr = structAlias.toAttribute + // `CreateNamedStruct` is never null and copies each field's nullability, so + // `GetStructField(structAttr, i).nullable` equals the original expression's nullability. + // The rewritten output attributes therefore match `SerializeFromObjectExec.output` exactly. + val outer = fused.zipWithIndex.map { case (ne, i) => + ne.withNewChildren(Seq(GetStructField(structAttr, i, Some(ne.name)))) + .asInstanceOf[NamedExpression] + } + ProjectExec(outer, inner) + } + } + } + + /** + * Plan-time gate. The rewrite is only worth making if the fused tree can actually reach the + * dispatcher; otherwise the resulting projection would fall back to Spark anyway, and we would + * have replaced a whole-stage-fused island with an unfused one. Binds the same way + * `CometScalaUDF.emitJvmCodegenDispatch` does so `canHandle` sees the tree it will really get. + */ + private def dispatchable(fusedValue: Expression): Option[Unit] = { + if (!CometConf.COMET_SCALA_UDF_CODEGEN_ENABLED.get()) { + logDebug( + "RewriteTypedDatasetMap: not rewriting because " + + s"${CometConf.COMET_SCALA_UDF_CODEGEN_ENABLED.key}=false") + return None + } + // Same binding as `emitJvmCodegenDispatch`, so `canHandle` sees the tree it will really get. + val attrs = fusedValue.collect { case a: AttributeReference => a }.distinct + val bound = BindReferences.bindReference(fusedValue, AttributeSeq(attrs)) + CometBatchKernelCodegen.canHandle(bound) match { + case Some(reason) => + logDebug(s"RewriteTypedDatasetMap: not rewriting because $reason") + None + case None => Some(()) + } + } + + /** + * Leave the sandwich alone and record why on the operator, so `EXPLAIN` shows the reason the + * typed operation stayed on Spark rather than the bare "not supported" the un-rewritten + * operators would otherwise produce. + */ + private def declineQuietly( + serialize: SerializeFromObjectExec, + reason: String): Option[SparkPlan] = { + withFallbackReason(serialize, s"Cannot fuse typed Dataset map: $reason") + None + } +} diff --git a/spark/src/main/scala/org/apache/comet/serde/CometScalaUDF.scala b/spark/src/main/scala/org/apache/comet/serde/CometScalaUDF.scala index 26fe0c591ab..e5ed77cab19 100644 --- a/spark/src/main/scala/org/apache/comet/serde/CometScalaUDF.scala +++ b/spark/src/main/scala/org/apache/comet/serde/CometScalaUDF.scala @@ -21,6 +21,7 @@ package org.apache.comet.serde import org.apache.spark.SparkEnv import org.apache.spark.sql.catalyst.expressions.{Attribute, AttributeReference, AttributeSeq, BindReferences, Expression, Literal, RuntimeReplaceable, ScalaUDF} +import org.apache.spark.sql.catalyst.trees.TreeNodeTag import org.apache.spark.sql.types.BinaryType import org.apache.comet.CometConf @@ -53,6 +54,24 @@ import org.apache.comet.udf.codegen.CometScalaUDFCodegen */ object CometScalaUDF extends CometExpressionSerde[ScalaUDF] { + /** + * Marks a subtree as "dispatch this whole thing as one kernel", overriding the normal per-node + * serde lookup in `QueryPlanSerde.exprToProtoInternal`. + * + * Needed when a rewrite builds a tree whose root has a perfectly good native serde but whose + * children must not be converted independently. `RewriteTypedDatasetMap` is the motivating + * case: it fuses a typed `Dataset.map` into a `CreateNamedStruct` over N serializer expressions + * that all share one `Invoke` of the user closure. Letting `CometCreateNamedStruct` convert + * each field separately would emit N dispatch protos and call the closure N times per row. + * Tagging the root produces one kernel instead, and subexpression elimination inside it + * collapses the shared `Invoke` to a single call per row. + * + * A tag rather than a Comet-specific `Expression` subclass on purpose: the rewritten plan stays + * built entirely from stock Spark expressions, so it still executes correctly if the enclosing + * operator ends up falling back to Spark for an unrelated reason. + */ + val FORCE_DISPATCH: TreeNodeTag[Unit] = TreeNodeTag[Unit]("comet.forceCodegenDispatch") + override def convert(expr: ScalaUDF, inputs: Seq[Attribute], binding: Boolean): Option[Expr] = emitJvmCodegenDispatch(expr, inputs, binding) diff --git a/spark/src/main/scala/org/apache/comet/serde/QueryPlanSerde.scala b/spark/src/main/scala/org/apache/comet/serde/QueryPlanSerde.scala index 479682ee640..9fb9024529a 100644 --- a/spark/src/main/scala/org/apache/comet/serde/QueryPlanSerde.scala +++ b/spark/src/main/scala/org/apache/comet/serde/QueryPlanSerde.scala @@ -1001,6 +1001,11 @@ object QueryPlanSerde extends Logging with CometExprShim with CometTypeShim { sparkVersionSpecificExprToProtoInternal(expr, inputs, binding) .orElse(expr match { + case _ if expr.getTagValue(CometScalaUDF.FORCE_DISPATCH).isDefined => + // A rewrite rule asked for this whole subtree to compile into one kernel rather than + // letting each child pick its own serde. See `CometScalaUDF.FORCE_DISPATCH`. + CometScalaUDF.emitJvmCodegenDispatch(expr, inputs, binding) + case UnaryExpression(child) if expr.prettyName == "promote_precision" => // `UnaryExpression` includes `PromotePrecision` for Spark 3.3 // `PromotePrecision` is just a wrapper, don't need to serialize it. diff --git a/spark/src/test/scala/org/apache/comet/CometTypedDatasetSuite.scala b/spark/src/test/scala/org/apache/comet/CometTypedDatasetSuite.scala new file mode 100644 index 00000000000..69716c9d407 --- /dev/null +++ b/spark/src/test/scala/org/apache/comet/CometTypedDatasetSuite.scala @@ -0,0 +1,342 @@ +/* + * 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 + +import java.util.concurrent.atomic.AtomicLong + +import org.apache.spark.api.java.function.MapFunction +import org.apache.spark.sql.CometTestBase +import org.apache.spark.sql.Encoders +import org.apache.spark.sql.comet.CometProjectExec +import org.apache.spark.sql.execution.{DeserializeToObjectExec, MapElementsExec, MapPartitionsExec, SerializeFromObjectExec, SparkPlan} +import org.apache.spark.sql.execution.adaptive.AdaptiveSparkPlanHelper +import org.apache.spark.sql.internal.SQLConf + +/** Top-level so `NewInstance` needs no outer pointer, which is the ordinary user shape. */ +case class TypedRec(a: Int, b: String) + +case class TypedWide(i: Int, s: String, d: java.math.BigDecimal, opt: Option[Long]) + +case class TypedNested(id: Int, inner: TypedRec, tags: Seq[String]) + +case class TypedDec(id: Int, d: java.math.BigDecimal) + +/** JVM-static counter. Comet tests run in local mode, so driver and executor share this. */ +object TypedMapCounter { + val calls = new AtomicLong(0) + + def reset(): Unit = calls.set(0) +} + +/** + * Tests for [[org.apache.comet.rules.RewriteTypedDatasetMap]], which fuses the + * `SerializeFromObject` / `MapElements` / `DeserializeToObject` sandwich a typed `Dataset.map` + * produces into a Comet projection. See https://github.com/apache/datafusion-comet/issues/5710. + */ +class CometTypedDatasetSuite extends CometTestBase with AdaptiveSparkPlanHelper { + + import testImplicits._ + + private def withFusion(pairs: (String, String)*)(f: => Unit): Unit = + withSQLConf( + (Seq( + CometConf.COMET_EXEC_TYPED_DATASET_MAP_ENABLED.key -> "true", + CometConf.COMET_SCALA_UDF_CODEGEN_ENABLED.key -> "true") ++ pairs): _*)(f) + + private def objectOperators(plan: SparkPlan): Seq[SparkPlan] = + collectWithSubqueries(plan) { + case p: SerializeFromObjectExec => p + case p: DeserializeToObjectExec => p + case p: MapElementsExec => p + case p: MapPartitionsExec => p + } + + private def assertNoObjectOperators(plan: SparkPlan): Unit = + assert( + objectOperators(plan).isEmpty, + s"expected the typed sandwich to be fused away, but plan still has " + + s"${objectOperators(plan).map(_.nodeName).mkString(", ")}:\n$plan") + + test("ds.map produces a fully native plan - single output column") { + withFusion() { + withParquetTable((0 until 100).map(i => (i, i.toString)), "tbl") { + val ds = spark.sql("select _1 as a, _2 as b from tbl").as[TypedRec] + val (_, cometPlan) = checkSparkAnswerAndOperator(ds.map(_.a + 1).toDF()) + assertNoObjectOperators(cometPlan) + } + } + } + + test("ds.map produces a fully native plan - multiple output columns") { + withFusion() { + withParquetTable((0 until 100).map(i => (i, i.toString)), "tbl") { + val ds = spark.sql("select _1 as a, _2 as b from tbl").as[TypedRec] + val (_, cometPlan) = + checkSparkAnswerAndOperator(ds.map(r => TypedRec(r.a + 1, r.b + "!")).toDF()) + assertNoObjectOperators(cometPlan) + // Inner projection builds the struct, outer one unpacks it. + assert( + collectWithSubqueries(cometPlan) { case p: CometProjectExec => p }.size >= 2, + s"expected stacked projections for the struct fuse:\n$cometPlan") + } + } + } + + test("output schema is unchanged by the rewrite") { + withParquetTable((0 until 20).map(i => (i, i.toString)), "tbl") { + // `withSQLConf` returns Unit on Spark 3.x, so capture rather than return from the block. + def schemaOf(fused: Boolean): String = { + var schema: String = null + withSQLConf(CometConf.COMET_EXEC_TYPED_DATASET_MAP_ENABLED.key -> fused.toString) { + schema = spark + .sql("select _1 as a, _2 as b from tbl") + .as[TypedRec] + .map(r => TypedRec(r.a + 1, r.b)) + .toDF() + .schema + .treeString + } + schema + } + assert(schemaOf(fused = true) === schemaOf(fused = false)) + } + } + + test("closure runs exactly once per row with multiple output columns") { + withFusion() { + withParquetTable((0 until 50).map(i => (i, i.toString)), "tbl") { + val ds = spark.sql("select _1 as a, _2 as b from tbl").as[TypedRec] + TypedMapCounter.reset() + val fused = ds + .map { r => + TypedMapCounter.calls.incrementAndGet() + TypedRec(r.a * 2, r.b) + } + .toDF() + assertNoObjectOperators(fused.queryExecution.executedPlan) + assert(fused.collect().length === 50) + // The whole point of the struct wrapper: N output columns must not mean N closure calls. + assert( + TypedMapCounter.calls.get() === 50, + s"expected 50 closure calls for 50 rows, got ${TypedMapCounter.calls.get()}") + } + } + } + + test("fusion unblocks the aggregate and shuffle above it") { + withFusion() { + withParquetTable((0 until 100).map(i => (i, (i % 7).toString)), "tbl") { + val ds = spark.sql("select _1 as a, _2 as b from tbl").as[TypedRec] + val (_, cometPlan) = + checkSparkAnswerAndOperator(ds.map(r => TypedRec(r.a + 1, r.b)).groupBy("b").count()) + assertNoObjectOperators(cometPlan) + } + } + } + + test("wide record with decimal, string and Option fields") { + withFusion() { + val rows = (1 to 40).map(i => + ( + i, + s"s$i", + new java.math.BigDecimal(s"$i.25"), + if (i % 3 == 0) null else Long.box(i * 10L))) + withTempPath { path => + rows.toDF("i", "s", "d", "opt").write.parquet(path.toString) + withParquetTable(path.toString, "tbl") { + val ds = spark.table("tbl").as[TypedWide] + val (_, cometPlan) = checkSparkAnswerAndOperator( + ds.map(r => TypedWide(r.i + 1, r.s, r.d, r.opt.map(_ + 1))).toDF()) + assertNoObjectOperators(cometPlan) + } + } + } + } + + test("nested struct and array fields round-trip") { + withFusion() { + withTempPath { path => + (1 to 30) + .map(i => (i, (i, s"n$i"), Seq(s"t$i", s"u$i"))) + .toDF("id", "inner", "tags") + .selectExpr("id", "named_struct('a', inner._1, 'b', inner._2) as inner", "tags") + .write + .parquet(path.toString) + withParquetTable(path.toString, "tbl") { + val ds = spark.table("tbl").as[TypedNested] + val (_, cometPlan) = checkSparkAnswerAndOperator( + ds.map(r => TypedNested(r.id + 1, TypedRec(r.inner.a, r.inner.b), r.tags.reverse)) + .toDF()) + assertNoObjectOperators(cometPlan) + } + } + } + } + + test("null field values in the input and the output") { + withFusion() { + withTempPath { path => + (1 to 30) + .map(i => (i, if (i % 4 == 0) null else s"s$i")) + .toDF("a", "b") + .write + .parquet(path.toString) + withParquetTable(path.toString, "tbl") { + val ds = spark.table("tbl").as[TypedRec] + val (_, cometPlan) = checkSparkAnswerAndOperator( + ds.map(r => TypedRec(r.a, if (r.a % 5 == 0) null else r.b)).toDF()) + assertNoObjectOperators(cometPlan) + } + } + } + } + + test("AssertNotNull inside the fused kernel still raises like Spark") { + // Returning null for a non-nullable top-level product is an error in Spark. The serializer's + // `assertnotnull` has to survive the fuse, or Comet would silently emit a null row instead. + withFusion() { + withParquetTable((0 until 20).map(i => (i, i.toString)), "tbl") { + val ds = spark.sql("select _1 as a, _2 as b from tbl").as[TypedRec] + val df = ds.map(r => if (r.a == 7) null else TypedRec(r.a, r.b)).toDF() + assertNoObjectOperators(df.queryExecution.executedPlan) + val err = intercept[Exception](df.collect()) + // Spark 3.x raises NullPointerException("Null value appeared ..."), Spark 4.x raises + // SparkRuntimeException("[NOT_NULL_ASSERT_VIOLATION] NULL value appeared ..."). Match the + // part they share. + assert( + Option(err.getMessage).exists( + _.toLowerCase.contains("value appeared in non-nullable field")), + s"expected a not-null assertion failure, got: ${err.getMessage}") + } + } + } + + test("decimal overflow in the serializer is caught before the Arrow write") { + // The concern on #5710 was that an encoder-declared decimal(38,18) could receive a wider + // value that Spark nulls at row materialization but the kernel's Arrow DecimalVector write + // would not -- the shape of the Iceberg truncate(w, decimal) bug from #5575. It does not + // apply: the encoder's serializer already wraps the value in `CheckOverflow`, so the fused + // tree raises (ANSI) or nulls (non-ANSI) exactly where Spark does, ahead of the write. + withFusion() { + withParquetTable((1 to 5).map(i => (i, new java.math.BigDecimal(s"$i.5"))), "tbl") { + val ds = spark.sql("select _1 as id, _2 as d from tbl").as[TypedDec] + def overflowed = + ds.map(r => TypedDec(r.id, r.d.multiply(new java.math.BigDecimal("1" + "0" * 30)))) + .toDF() + + assertNoObjectOperators(overflowed.queryExecution.executedPlan) + + withSQLConf(SQLConf.ANSI_ENABLED.key -> "true") { + val err = intercept[Exception](overflowed.collect()) + assert( + err.getMessage.contains("cannot be represented as Decimal(38, 18)"), + s"expected an ANSI overflow error, got: ${err.getMessage}") + } + withSQLConf(SQLConf.ANSI_ENABLED.key -> "false") { + // Non-ANSI nulls the overflowing value; Comet must produce the same nulls as Spark. + val (_, cometPlan) = checkSparkAnswerAndOperator(overflowed) + assertNoObjectOperators(cometPlan) + } + } + } + } + + test("Java MapFunction takes the same path as a Scala closure") { + withFusion() { + withParquetTable((0 until 40).map(i => (i, i.toString)), "tbl") { + val ds = spark.sql("select _1 as a, _2 as b from tbl").as[TypedRec] + val fn = new MapFunction[TypedRec, TypedRec] { + override def call(r: TypedRec): TypedRec = TypedRec(r.a + 100, r.b) + } + val (_, cometPlan) = + checkSparkAnswerAndOperator(ds.map(fn, Encoders.product[TypedRec]).toDF()) + assertNoObjectOperators(cometPlan) + } + } + } + + test("chained maps fuse into one native pipeline") { + withFusion() { + withParquetTable((0 until 40).map(i => (i, i.toString)), "tbl") { + val ds = spark.sql("select _1 as a, _2 as b from tbl").as[TypedRec] + val (_, cometPlan) = checkSparkAnswerAndOperator( + ds.map(r => TypedRec(r.a + 1, r.b)).map(r => TypedRec(r.a * 2, r.b + "x")).toDF()) + assertNoObjectOperators(cometPlan) + } + } + } + + test("off by default") { + withParquetTable((0 until 20).map(i => (i, i.toString)), "tbl") { + val ds = spark.sql("select _1 as a, _2 as b from tbl").as[TypedRec] + val df = ds.map(r => TypedRec(r.a + 1, r.b)).toDF() + df.collect() + assert( + objectOperators(df.queryExecution.executedPlan).nonEmpty, + "rewrite must not apply unless it is explicitly enabled") + } + } + + test("declines when the codegen dispatcher is disabled") { + withSQLConf( + CometConf.COMET_EXEC_TYPED_DATASET_MAP_ENABLED.key -> "true", + CometConf.COMET_SCALA_UDF_CODEGEN_ENABLED.key -> "false") { + withParquetTable((0 until 20).map(i => (i, i.toString)), "tbl") { + val ds = spark.sql("select _1 as a, _2 as b from tbl").as[TypedRec] + val df = ds.map(r => TypedRec(r.a + 1, r.b)).toDF() + checkSparkAnswer(df) + assert( + objectOperators(df.queryExecution.executedPlan).nonEmpty, + "without the dispatcher there is nothing to fuse into") + } + } + } + + test("declines multi-column fusion when subexpression elimination is off") { + // Without CSE the single kernel would evaluate the shared closure Invoke once per struct + // field, so the rule must decline rather than change how many times user code runs. + withFusion(SQLConf.SUBEXPRESSION_ELIMINATION_ENABLED.key -> "false") { + withParquetTable((0 until 20).map(i => (i, i.toString)), "tbl") { + val ds = spark.sql("select _1 as a, _2 as b from tbl").as[TypedRec] + val df = ds.map(r => TypedRec(r.a + 1, r.b)).toDF() + checkSparkAnswer(df) + assert( + objectOperators(df.queryExecution.executedPlan).nonEmpty, + "multi-column fusion needs CSE to keep the closure to one call per row") + } + } + } + + test("mapPartitions is not fused") { + withFusion() { + withParquetTable((0 until 20).map(i => (i, i.toString)), "tbl") { + val ds = spark.sql("select _1 as a, _2 as b from tbl").as[TypedRec] + val df = ds.mapPartitions(it => it.map(r => TypedRec(r.a + 1, r.b))).toDF() + checkSparkAnswer(df) + assert( + collectWithSubqueries(df.queryExecution.executedPlan) { case p: MapPartitionsExec => + p + }.nonEmpty, + "mapPartitions is iterator-shaped and must be left alone") + } + } + } +} From caa1797f763f754093e130f77732b18e94cda385 Mon Sep 17 00:00:00 2001 From: Andy Grove Date: Sat, 5 Sep 2026 11:38:27 -0600 Subject: [PATCH 2/6] refactor: address review of the typed Dataset map fuse Fixes a real limitation found while reviewing: `canHandle`'s `spark.sql.codegen.maxFields` gate counted `BoundReference` *occurrences*, but the fused struct puts the shared deserializer in every field on purpose, so a 12-column typed map counted 12 + 12*12 = 156 fields against a default cap of 100 and silently declined to fuse. WSCG gates on the operator's schema, where each column contributes once, and the kernel likewise emits one typed field and one getter per ordinal -- so count distinct ordinals. Verified: a 12-column record declined before, fuses now, and is pinned by a new test. Reuse and duplication: - Extract `CometScalaUDF.canDispatch` / `bindForDispatch`. The rule was re-deriving the bind-then-canHandle gate and had already drifted: it omitted the `RuntimeReplaceable` unwrap, so the plan-time prediction could disagree with what the serde does. - Move `FORCE_DISPATCH` to `QueryPlanSerde`, which reads it. Matches the convention every other behavior tag follows (SKIP_COMET_SCAN_TAG, SKIP_COMET_BROADCAST_TAG, COMET_UNSAFE_PARTIAL all live on their reader), and documents that the tag bypasses the per-expression policy layer. - Check the tag ahead of the version shim, not inside the `orElse`, so the override is genuinely unconditional -- the shim matches on Invoke/StaticInvoke, exactly what this rule synthesizes. Simplification: fold the dispatch gate and tag-set into one `forceDispatch` so the two facts cannot travel separately; hoist the CSE precondition above the output-count branch; `AttributeSet(fused) -- child.outputSet` for the dangling check; rename `declineQuietly` to `decline` (it is the opposite of quiet -- it writes an EXPLAIN-visible reason) and route the dispatcher-disabled case through it so that reason reaches EXPLAIN too. Altitude: correct the pre-pass comment, which claimed a technical necessity that does not exist -- `convertNode` only tags `DeserializeToObjectExec`, it never replaces it, so the sandwich is intact when the walk reaches the parent. Records that doing it in `convertNode` would additionally expose the profitability signal this rewrite lacks. Move the pre-pass above `normalizePlan` so the synthesized projections get the same NaN / -0.0 normalization as every other `ProjectExec`. Tests: replace a tautological schema assertion (Dataset.schema comes from the analyzed plan, which a physical rule cannot change) with one on the executed plan's output attributes; pin both decline messages with `checkSparkAnswerAndFallbackReason`; collapse the four-case operator collector to `ObjectConsumerExec | ObjectProducerExec`; stop building the failure clue on the happy path; add `withTypedRecs` / `withParquetRoundTrip` fixtures. --- .../codegen/CometBatchKernelCodegen.scala | 17 +- .../apache/comet/rules/CometExecRule.scala | 41 ++-- .../comet/rules/RewriteTypedDatasetMap.scala | 169 +++++++------- .../apache/comet/serde/CometScalaUDF.scala | 64 +++--- .../apache/comet/serde/QueryPlanSerde.scala | 41 +++- .../apache/comet/CometTypedDatasetSuite.scala | 212 ++++++++++-------- 6 files changed, 321 insertions(+), 223 deletions(-) diff --git a/spark/src/main/scala/org/apache/comet/codegen/CometBatchKernelCodegen.scala b/spark/src/main/scala/org/apache/comet/codegen/CometBatchKernelCodegen.scala index 3e2456234ec..c0aa042254c 100644 --- a/spark/src/main/scala/org/apache/comet/codegen/CometBatchKernelCodegen.scala +++ b/spark/src/main/scala/org/apache/comet/codegen/CometBatchKernelCodegen.scala @@ -124,9 +124,22 @@ object CometBatchKernelCodegen extends Logging with CometExprTraitShim with Come // typed input field count, the typed-getter switch, and the constant pool. Refuse here so the // operator falls back to Spark cleanly rather than tripping a Janino compile failure // mid-execution (Comet has no recovery for that). + // + // Count each input ordinal once. WSCG gates on the operator's *schema* + // (`plan.schema.map(_.dataType).map(numOfNestedFields).sum`), so a column read more than once + // contributes once; the kernel likewise emits one typed field and one getter per ordinal, not + // per occurrence. Counting occurrences instead would scale with how often the tree happens to + // repeat a column, which is what `RewriteTypedDatasetMap` does deliberately: it puts the same + // shared subtree in every struct field so subexpression elimination can collapse it. A + // 12-column typed map reads 12 ordinals in each of 12 fields, which counted per occurrence + // came to 156 and refused a kernel that is really 24 fields wide. val maxFields = SQLConf.get.wholeStageMaxNumFields - val totalFields = numOfNestedFields(boundExpr.dataType) + - boundExpr.collect { case b: BoundReference => numOfNestedFields(b.dataType) }.sum + val inputFields = boundExpr + .collect { case b: BoundReference => b.ordinal -> b.dataType } + .distinct + .map { case (_, dt) => numOfNestedFields(dt) } + .sum + val totalFields = numOfNestedFields(boundExpr.dataType) + inputFields if (totalFields > maxFields) { return Some( s"codegen dispatch: too many nested fields ($totalFields > " + 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 f45aa1cf6ca..353fffeeceb 100644 --- a/spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala +++ b/spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala @@ -697,7 +697,29 @@ case class CometExecRule(session: SparkSession) plan } } else { - val normalizedPlan = normalizePlan(plan) + // Collapse the SerializeFromObject / MapElements / DeserializeToObject sandwich that a typed + // `Dataset.map` produces into a projection. Runs before `normalizePlan` so the projections it + // synthesizes get the same NaN / -0.0 normalization every other `ProjectExec` in the plan + // gets; encoder serializer trees contain no comparison operators today, so this is + // future-proofing rather than a live fix. + // + // A pre-pass rather than a `convertNode` case for now. Note that `convertNode` *could* do it: + // the bottom-up walk only tags `DeserializeToObjectExec` with a fallback reason, it never + // replaces it, so the sandwich is still structurally intact when the walk reaches + // `SerializeFromObjectExec`. Doing it there would additionally let the rule see whether the + // deserializer's child converted to a `CometNativeExec` and decline when it did not, which is + // the profitability signal this rewrite currently lacks. See + // https://github.com/apache/datafusion-comet/issues/5710. + val planWithTypedMapFused = + if (CometConf.COMET_EXEC_TYPED_DATASET_MAP_ENABLED.get(conf)) { + plan.transformUp { case p => + RewriteTypedDatasetMap.rewrite(p) + } + } else { + plan + } + + val normalizedPlan = normalizePlan(planWithTypedMapFused) val planWithJoinRewritten = if (CometConf.COMET_FORCE_SHJ.get()) { normalizedPlan.transformUp { case p => @@ -707,26 +729,13 @@ case class CometExecRule(session: SparkSession) normalizedPlan } - // Collapse the SerializeFromObject / MapElements / DeserializeToObject sandwich that a typed - // `Dataset.map` produces into a projection, before the bottom-up conversion sees the - // individual operators. `transform()` visits children first, so by the time it reached - // `DeserializeToObjectExec` the island would already have fallen back. - val planWithTypedMapFused = - if (CometConf.COMET_EXEC_TYPED_DATASET_MAP_ENABLED.get(conf)) { - planWithJoinRewritten.transformUp { case p => - RewriteTypedDatasetMap.rewrite(p) - } - } else { - planWithJoinRewritten - } - // Tag Partial aggregates that must not be converted to Comet because a // corresponding Final or PartialMerge cannot be converted and the intermediate buffer // formats are incompatible. This runs before transform() so the tags are checked // during the bottom-up conversion. Tags persist through AQE stage creation. - tagUnsafePartialAggregates(planWithTypedMapFused) + tagUnsafePartialAggregates(planWithJoinRewritten) - var newPlan = transform(planWithTypedMapFused) + var newPlan = transform(planWithJoinRewritten) // if the plan cannot be run fully natively then explain why (when appropriate // config is enabled) diff --git a/spark/src/main/scala/org/apache/comet/rules/RewriteTypedDatasetMap.scala b/spark/src/main/scala/org/apache/comet/rules/RewriteTypedDatasetMap.scala index 14095c66aeb..8f8488d931e 100644 --- a/spark/src/main/scala/org/apache/comet/rules/RewriteTypedDatasetMap.scala +++ b/spark/src/main/scala/org/apache/comet/rules/RewriteTypedDatasetMap.scala @@ -21,7 +21,7 @@ package org.apache.comet.rules import org.apache.spark.api.java.function.MapFunction import org.apache.spark.internal.Logging -import org.apache.spark.sql.catalyst.expressions.{Alias, AttributeReference, AttributeSeq, AttributeSet, BindReferences, BoundReference, CreateNamedStruct, Expression, GetStructField, Literal, NamedExpression} +import org.apache.spark.sql.catalyst.expressions.{Alias, AttributeSet, BoundReference, CreateNamedStruct, Expression, GetStructField, Literal, NamedExpression} import org.apache.spark.sql.catalyst.expressions.objects.Invoke import org.apache.spark.sql.catalyst.plans.logical.FunctionUtils import org.apache.spark.sql.execution.{DeserializeToObjectExec, MapElementsExec, ProjectExec, SerializeFromObjectExec, SparkPlan} @@ -30,8 +30,7 @@ import org.apache.spark.sql.types.ObjectType import org.apache.comet.CometConf import org.apache.comet.CometSparkSessionExtensions.withFallbackReason -import org.apache.comet.codegen.CometBatchKernelCodegen -import org.apache.comet.serde.CometScalaUDF +import org.apache.comet.serde.{CometScalaUDF, QueryPlanSerde} /** * Collapses the three-operator island that a typed `Dataset.map` produces @@ -85,11 +84,11 @@ object RewriteTypedDatasetMap extends Logging { // `ds.map(f).map(g)` leaves two adjacent `MapElements` under one Serialize/Deserialize pair, // so walk the whole chain rather than expecting exactly one. val chain = mapElementsChain(serialize.child) - chain.lastOption.map(_.child) match { - case Some(deserialize: DeserializeToObjectExec) => - fuse(serialize, chain, deserialize).getOrElse(plan) - case _ => plan - } + chain.lastOption + .map(_.child) + .collect { case deserialize: DeserializeToObjectExec => deserialize } + .flatMap(fuse(serialize, chain, _)) + .getOrElse(plan) case _ => plan } @@ -107,26 +106,36 @@ object RewriteTypedDatasetMap extends Logging { val serializer = serialize.serializer val objType = chain.head.outputObjectType - // Every serializer element is an `Alias` for encoder-generated serializers. The outer - // projection rebuilds them with `withNewChildren`, which needs exactly one child, and that - // is also what preserves `exprId` / qualifier / metadata across the rewrite. - if (!serializer.forall(_.isInstanceOf[Alias])) { - return declineQuietly( + // Cheapest precondition first: a config that cannot change mid-plan. + if (!CometConf.COMET_SCALA_UDF_CODEGEN_ENABLED.get()) { + return decline( + serialize, + s"${CometConf.COMET_SCALA_UDF_CODEGEN_ENABLED.key}=false, so there is no dispatcher to " + + "fuse into") + } + + // Every serializer element is an `Alias` for encoder-generated serializers -- Spark builds them + // via `ExpressionEncoder.namedExpressions`, which aliases every field of the flattened + // `CreateNamedStruct`. The outer projection relies on that: it rebuilds each element with + // `withNewChildren`, which needs exactly one child, and that is what preserves `exprId` / + // qualifier / metadata across the rewrite. + val nonAliases = serializer.filterNot(_.isInstanceOf[Alias]) + if (nonAliases.nonEmpty) { + return decline( serialize, "serializer contains a non-Alias element: " + - serializer - .filterNot(_.isInstanceOf[Alias]) - .map(_.getClass.getSimpleName) - .mkString(", ")) + nonAliases.map(_.getClass.getSimpleName).mkString(", ")) } - // The serializer reads the object through `BoundReference(0, objType)`. Anything else means a - // shape this rule has not been reasoned about, so leave it alone rather than guess. + // The serializer reads the object through `BoundReference(0, objType)` -- Spark asserts as much + // in `ExpressionEncoder` ("all serializer expressions must use the same BoundReference") and + // `ScalaReflection.serializerFor` builds it at ordinal 0. Anything else is a shape this rule + // has not been reasoned about, so leave it alone rather than guess. val badRefs = serializer.flatMap(_.collect { case b: BoundReference if b.ordinal != 0 || b.dataType != objType => b }) if (badRefs.nonEmpty) { - return declineQuietly( + return decline( serialize, s"serializer reads unexpected bound references: ${badRefs.mkString(", ")}") } @@ -148,95 +157,87 @@ object RewriteTypedDatasetMap extends Logging { propagateNull = false) } - // Substitute the closure call for the object the serializer reads. `transform` on an `Alias` - // preserves `exprId`, qualifier and metadata via `otherCopyArgs`, so the rewritten projection - // keeps `SerializeFromObjectExec`'s exact output attributes and parents stay valid. + // Substitute the closure call for the object the serializer reads. The guard above proved every + // `BoundReference` here is that object, and the pattern restates it so the substitution cannot + // outlive its precondition. `transform` on an `Alias` preserves `exprId`, qualifier and + // metadata via `otherCopyArgs`, so the rewritten projection keeps + // `SerializeFromObjectExec`'s exact output attributes and parents stay valid. val fused = serializer.map { ne => - ne.transform { case _: BoundReference => callFunc }.asInstanceOf[NamedExpression] + ne.transform { + case b: BoundReference if b.ordinal == 0 && b.dataType == objType => callFunc + }.asInstanceOf[NamedExpression] } - // A projection may only reference its child's output. The deserializer reads `child.output` - // and the object reference is gone, so this should hold; check rather than assume. - val childOutput = AttributeSet(child.output) - val dangling = fused.flatMap(f => (f.references -- childOutput).toSeq) + // A projection may only reference its child's output. Spark asserts the serializer itself has + // no free references (`ExpressionEncoder`), and the substituted tree bottoms out in the + // deserializer, which reads `child.output` -- so this should hold; check rather than assume. + val dangling = AttributeSet(fused) -- child.outputSet if (dangling.nonEmpty) { - return declineQuietly( + return decline( serialize, s"fused expression references attributes outside the child: ${dangling.mkString(", ")}") } - if (fused.length == 1) { - // Single output column: the fused expression is already one dispatch kernel, so there is - // nothing for the struct wrapper to deduplicate. - val single = fused.head - val value = single.children.head - dispatchable(value).map { _ => - value.setTagValue(CometScalaUDF.FORCE_DISPATCH, ()) - ProjectExec(fused, child) - } - } else { - // Without subexpression elimination the single kernel would still evaluate the shared - // `Invoke` once per struct field, so the closure would run N times per row. Spark runs it - // once; decline rather than change how many times user code executes. - if (!SQLConf.get.subexpressionEliminationEnabled) { - return declineQuietly( - serialize, - s"${serializer.length} output columns require subexpression elimination to keep the " + - "closure to one call per row, but " + - s"${SQLConf.SUBEXPRESSION_ELIMINATION_ENABLED.key}=false") - } + // Without subexpression elimination the single kernel would evaluate the shared `Invoke` once + // per struct field, so the closure would run N times per row. Spark runs it once; decline + // rather than change how many times user code executes. One output column has nothing to + // deduplicate, so it does not need CSE. + if (fused.length > 1 && !SQLConf.get.subexpressionEliminationEnabled) { + return decline( + serialize, + s"${fused.length} output columns require subexpression elimination to keep the closure " + + s"to one call per row, but ${SQLConf.SUBEXPRESSION_ELIMINATION_ENABLED.key}=false") + } - val structExpr = - CreateNamedStruct(fused.flatMap(ne => Seq(Literal(ne.name), ne.children.head))) - dispatchable(structExpr).map { _ => - structExpr.setTagValue(CometScalaUDF.FORCE_DISPATCH, ()) - val structAlias = Alias(structExpr, FUSED_COLUMN)() - val inner = ProjectExec(Seq(structAlias), child) - val structAttr = structAlias.toAttribute - // `CreateNamedStruct` is never null and copies each field's nullability, so - // `GetStructField(structAttr, i).nullable` equals the original expression's nullability. - // The rewritten output attributes therefore match `SerializeFromObjectExec.output` exactly. - val outer = fused.zipWithIndex.map { case (ne, i) => - ne.withNewChildren(Seq(GetStructField(structAttr, i, Some(ne.name)))) - .asInstanceOf[NamedExpression] + fused match { + // One output column is already one kernel; the struct wrapper would have nothing to dedupe. + case Seq(only) => + forceDispatch(only.children.head).map(_ => ProjectExec(fused, child)) + + case _ => + forceDispatch( + CreateNamedStruct(fused.flatMap(ne => Seq(Literal(ne.name), ne.children.head)))).map { + structExpr => + val structAlias = Alias(structExpr, FUSED_COLUMN)() + val inner = ProjectExec(Seq(structAlias), child) + val structAttr = structAlias.toAttribute + // `CreateNamedStruct` is never null and copies each field's nullability, so + // `GetStructField(structAttr, i).nullable` equals the original expression's + // nullability. The rewritten output attributes therefore match + // `SerializeFromObjectExec.output` exactly. + val outer = fused.zipWithIndex.map { case (ne, i) => + ne.withNewChildren(Seq(GetStructField(structAttr, i, Some(ne.name)))) + .asInstanceOf[NamedExpression] + } + ProjectExec(outer, inner) } - ProjectExec(outer, inner) - } } } /** - * Plan-time gate. The rewrite is only worth making if the fused tree can actually reach the - * dispatcher; otherwise the resulting projection would fall back to Spark anyway, and we would - * have replaced a whole-stage-fused island with an unfused one. Binds the same way - * `CometScalaUDF.emitJvmCodegenDispatch` does so `canHandle` sees the tree it will really get. + * Tag `expr` to compile into a single kernel, or `None` when the dispatcher would refuse it. + * + * The gate matters: if the fused tree cannot reach the dispatcher, the projection this rule + * builds would fall back to Spark anyway, and we would have traded a whole-stage-fused island + * for an unfused one. Delegating to [[CometScalaUDF.canDispatch]] rather than re-deriving the + * binding is what keeps the prediction identical to what the serde will really do. */ - private def dispatchable(fusedValue: Expression): Option[Unit] = { - if (!CometConf.COMET_SCALA_UDF_CODEGEN_ENABLED.get()) { - logDebug( - "RewriteTypedDatasetMap: not rewriting because " + - s"${CometConf.COMET_SCALA_UDF_CODEGEN_ENABLED.key}=false") - return None - } - // Same binding as `emitJvmCodegenDispatch`, so `canHandle` sees the tree it will really get. - val attrs = fusedValue.collect { case a: AttributeReference => a }.distinct - val bound = BindReferences.bindReference(fusedValue, AttributeSeq(attrs)) - CometBatchKernelCodegen.canHandle(bound) match { + private def forceDispatch[T <: Expression](expr: T): Option[T] = + CometScalaUDF.canDispatch(expr) match { case Some(reason) => logDebug(s"RewriteTypedDatasetMap: not rewriting because $reason") None - case None => Some(()) + case None => + expr.setTagValue(QueryPlanSerde.FORCE_DISPATCH, ()) + Some(expr) } - } /** * Leave the sandwich alone and record why on the operator, so `EXPLAIN` shows the reason the * typed operation stayed on Spark rather than the bare "not supported" the un-rewritten * operators would otherwise produce. */ - private def declineQuietly( - serialize: SerializeFromObjectExec, - reason: String): Option[SparkPlan] = { + private def decline(serialize: SerializeFromObjectExec, reason: String): Option[SparkPlan] = { withFallbackReason(serialize, s"Cannot fuse typed Dataset map: $reason") None } diff --git a/spark/src/main/scala/org/apache/comet/serde/CometScalaUDF.scala b/spark/src/main/scala/org/apache/comet/serde/CometScalaUDF.scala index e5ed77cab19..0393b4e9b7f 100644 --- a/spark/src/main/scala/org/apache/comet/serde/CometScalaUDF.scala +++ b/spark/src/main/scala/org/apache/comet/serde/CometScalaUDF.scala @@ -21,7 +21,6 @@ package org.apache.comet.serde import org.apache.spark.SparkEnv import org.apache.spark.sql.catalyst.expressions.{Attribute, AttributeReference, AttributeSeq, BindReferences, Expression, Literal, RuntimeReplaceable, ScalaUDF} -import org.apache.spark.sql.catalyst.trees.TreeNodeTag import org.apache.spark.sql.types.BinaryType import org.apache.comet.CometConf @@ -54,26 +53,47 @@ import org.apache.comet.udf.codegen.CometScalaUDFCodegen */ object CometScalaUDF extends CometExpressionSerde[ScalaUDF] { + override def convert(expr: ScalaUDF, inputs: Seq[Attribute], binding: Boolean): Option[Expr] = + emitJvmCodegenDispatch(expr, inputs, binding) + /** - * Marks a subtree as "dispatch this whole thing as one kernel", overriding the normal per-node - * serde lookup in `QueryPlanSerde.exprToProtoInternal`. + * Bind `expr` the way [[emitJvmCodegenDispatch]] will, returning the attributes it reads in + * ordinal order alongside the bound tree. * - * Needed when a rewrite builds a tree whose root has a perfectly good native serde but whose - * children must not be converted independently. `RewriteTypedDatasetMap` is the motivating - * case: it fuses a typed `Dataset.map` into a `CreateNamedStruct` over N serializer expressions - * that all share one `Invoke` of the user closure. Letting `CometCreateNamedStruct` convert - * each field separately would emit N dispatch protos and call the closure N times per row. - * Tagging the root produces one kernel instead, and subexpression elimination inside it - * collapses the shared `Invoke` to a single call per row. - * - * A tag rather than a Comet-specific `Expression` subclass on purpose: the rewritten plan stays - * built entirely from stock Spark expressions, so it still executes correctly if the enclosing - * operator ends up falling back to Spark for an unrelated reason. + * Callers that only need to know whether a tree is dispatchable should use [[canDispatch]] + * rather than repeating this. */ - val FORCE_DISPATCH: TreeNodeTag[Unit] = TreeNodeTag[Unit]("comet.forceCodegenDispatch") + private def bindForDispatch(expr: Expression): (Seq[AttributeReference], Expression) = { + // `RuntimeReplaceable` expressions (e.g. Spark 4's `StructsToJson`) have a `doGenCode` that + // always throws "Cannot generate code for expression". Catalyst's `ReplaceExpressions` rule + // normally rewrites them to their `replacement` form before codegen runs. Comet's serde + // sometimes works with the pre-rewrite form (via shim reconstruction) for matching purposes, + // so unwrap to the replacement here before binding so the kernel compiles. + val target = expr match { + case rr: RuntimeReplaceable => rr.replacement + case other => other + } + // Bind against only the AttributeReferences the tree actually reads, so ordinals align with + // the data args we ship. + val attrs = target.collect { case a: AttributeReference => a }.distinct + (attrs, BindReferences.bindReference(target, AttributeSeq(attrs))) + } - override def convert(expr: ScalaUDF, inputs: Seq[Attribute], binding: Boolean): Option[Expr] = - emitJvmCodegenDispatch(expr, inputs, binding) + /** + * Plan-time gate: `None` if the dispatcher would accept `expr`, `Some(reason)` otherwise. + * + * Exposed so a rule that decides whether to rewrite a plan into a dispatchable shape predicts + * the same answer [[emitJvmCodegenDispatch]] will give it later. Do not re-derive the binding + * at the call site -- the two would drift, and a stale prediction commits a plan neither side + * wants. + */ + def canDispatch(expr: Expression): Option[String] = { + if (!CometConf.COMET_SCALA_UDF_CODEGEN_ENABLED.get()) { + return Some( + s"${CometConf.COMET_SCALA_UDF_CODEGEN_ENABLED.key}=false; expression has no native path") + } + CometBatchKernelCodegen.canHandle(bindForDispatch(expr)._2) + } /** * Bind `expr`, closure-serialize it, and emit a `JvmScalarUdf` proto routed through @@ -104,15 +124,7 @@ object CometScalaUDF extends CometExpressionSerde[ScalaUDF] { // normally rewrites them to their `replacement` form before codegen runs. Comet's serde // sometimes works with the pre-rewrite form (via shim reconstruction) for matching purposes, // so unwrap to the replacement here before binding so the kernel compiles. - val target = expr match { - case rr: RuntimeReplaceable => rr.replacement - case other => other - } - - // Bind against only the AttributeReferences the tree actually reads, so ordinals align with - // the data args we ship. - val attrs = target.collect { case a: AttributeReference => a }.distinct - val boundExpr = BindReferences.bindReference(target, AttributeSeq(attrs)) + val (attrs, boundExpr) = bindForDispatch(expr) // Gate at plan time. Surface the reason via withFallbackReason rather than crashing Janino // at execute. diff --git a/spark/src/main/scala/org/apache/comet/serde/QueryPlanSerde.scala b/spark/src/main/scala/org/apache/comet/serde/QueryPlanSerde.scala index 9fb9024529a..8f83814b3b1 100644 --- a/spark/src/main/scala/org/apache/comet/serde/QueryPlanSerde.scala +++ b/spark/src/main/scala/org/apache/comet/serde/QueryPlanSerde.scala @@ -30,6 +30,7 @@ import org.apache.spark.sql.catalyst.expressions._ import org.apache.spark.sql.catalyst.expressions.aggregate._ import org.apache.spark.sql.catalyst.expressions.objects.StaticInvoke import org.apache.spark.sql.catalyst.expressions.xml.{XPathBoolean, XPathDouble, XPathFloat, XPathInt, XPathList, XPathLong, XPathShort, XPathString} +import org.apache.spark.sql.catalyst.trees.TreeNodeTag import org.apache.spark.sql.comet.DecimalPrecision import org.apache.spark.sql.execution.{ScalarSubquery, SparkPlan} import org.apache.spark.sql.execution.datasources.parquet.ParquetUtils @@ -52,6 +53,31 @@ import org.apache.comet.shims.{CometExprShim, CometTypeShim} */ object QueryPlanSerde extends Logging with CometExprShim with CometTypeShim { + /** + * Marks a subtree as "dispatch this whole thing as one kernel", overriding the normal per-node + * serde lookup in [[exprToProtoInternal]]. + * + * Needed when a rewrite builds a tree whose root has a perfectly good native serde but whose + * children must not be converted independently. `RewriteTypedDatasetMap` is the motivating + * case: it fuses a typed `Dataset.map` into a `CreateNamedStruct` over N serializer expressions + * that all share one `Invoke` of the user closure. Letting `CometCreateNamedStruct` convert + * each field separately would emit N dispatch protos and call the closure N times per row. + * Tagging the root produces one kernel instead, and subexpression elimination inside it + * collapses the shared `Invoke` to a single call per row. + * + * A tag rather than a Comet-specific `Expression` subclass on purpose: the rewritten plan stays + * built entirely from stock Spark expressions, so it still executes correctly if the enclosing + * operator ends up falling back to Spark for an unrelated reason. (Comet has no `Expression` + * subclasses at all, and its whole expression story is serdes over stock Spark nodes.) + * + * Caveat for future callers: the tag routes straight to `CometScalaUDF.emitJvmCodegenDispatch`, + * bypassing the per-expression policy layer in `exprToProtoInternal`'s `convert` helper -- so + * `CometConf.isExprEnabled`, `getSupportLevel` and `allowIncompatible` are *not* consulted for + * a tagged node. Only tag trees you have already decided are dispatchable, via + * [[CometScalaUDF.canDispatch]]. + */ + val FORCE_DISPATCH: TreeNodeTag[Unit] = TreeNodeTag[Unit]("comet.forceCodegenDispatch") + private[comet] val arrayExpressions: Map[Class[_ <: Expression], CometExpressionSerde[_]] = Map( // ArrayAppend is a concrete expression only on Spark 3.x. Spark 4.0+ marks it // RuntimeReplaceable and rewrites it to ArrayInsert before serde, so this entry is @@ -998,14 +1024,19 @@ object QueryPlanSerde extends Logging with CometExprShim with CometTypeShim { } } + // A rewrite rule asked for this whole subtree to compile into one kernel rather than letting + // each child pick its own serde. Checked ahead of the version shim so the override is + // unconditional: the shim pattern-matches on `Invoke`/`StaticInvoke`, which is exactly the + // shape `RewriteTypedDatasetMap` synthesizes. See `FORCE_DISPATCH`. + if (expr.getTagValue(FORCE_DISPATCH).isDefined) { + return CometScalaUDF + .emitJvmCodegenDispatch(expr, inputs, binding) + .map(attachExprIdAndContext(expr, _)) + } + sparkVersionSpecificExprToProtoInternal(expr, inputs, binding) .orElse(expr match { - case _ if expr.getTagValue(CometScalaUDF.FORCE_DISPATCH).isDefined => - // A rewrite rule asked for this whole subtree to compile into one kernel rather than - // letting each child pick its own serde. See `CometScalaUDF.FORCE_DISPATCH`. - CometScalaUDF.emitJvmCodegenDispatch(expr, inputs, binding) - case UnaryExpression(child) if expr.prettyName == "promote_precision" => // `UnaryExpression` includes `PromotePrecision` for Spark 3.3 // `PromotePrecision` is just a wrapper, don't need to serialize it. diff --git a/spark/src/test/scala/org/apache/comet/CometTypedDatasetSuite.scala b/spark/src/test/scala/org/apache/comet/CometTypedDatasetSuite.scala index 69716c9d407..82076235b57 100644 --- a/spark/src/test/scala/org/apache/comet/CometTypedDatasetSuite.scala +++ b/spark/src/test/scala/org/apache/comet/CometTypedDatasetSuite.scala @@ -22,11 +22,9 @@ package org.apache.comet import java.util.concurrent.atomic.AtomicLong import org.apache.spark.api.java.function.MapFunction -import org.apache.spark.sql.CometTestBase -import org.apache.spark.sql.Encoders +import org.apache.spark.sql.{CometTestBase, DataFrame, Dataset, Encoders} import org.apache.spark.sql.comet.CometProjectExec -import org.apache.spark.sql.execution.{DeserializeToObjectExec, MapElementsExec, MapPartitionsExec, SerializeFromObjectExec, SparkPlan} -import org.apache.spark.sql.execution.adaptive.AdaptiveSparkPlanHelper +import org.apache.spark.sql.execution.{MapPartitionsExec, ObjectConsumerExec, ObjectProducerExec, SparkPlan} import org.apache.spark.sql.internal.SQLConf /** Top-level so `NewInstance` needs no outer pointer, which is the ordinary user shape. */ @@ -38,6 +36,21 @@ case class TypedNested(id: Int, inner: TypedRec, tags: Seq[String]) case class TypedDec(id: Int, d: java.math.BigDecimal) +/** Twelve columns, to pin the `spark.sql.codegen.maxFields` accounting. */ +case class TypedWide12( + c1: Int, + c2: Int, + c3: Int, + c4: Int, + c5: Int, + c6: Int, + c7: Int, + c8: Int, + c9: Int, + c10: Int, + c11: Int, + c12: Int) + /** JVM-static counter. Comet tests run in local mode, so driver and executor share this. */ object TypedMapCounter { val calls = new AtomicLong(0) @@ -50,7 +63,7 @@ object TypedMapCounter { * `SerializeFromObject` / `MapElements` / `DeserializeToObject` sandwich a typed `Dataset.map` * produces into a Comet projection. See https://github.com/apache/datafusion-comet/issues/5710. */ -class CometTypedDatasetSuite extends CometTestBase with AdaptiveSparkPlanHelper { +class CometTypedDatasetSuite extends CometTestBase { import testImplicits._ @@ -60,24 +73,44 @@ class CometTypedDatasetSuite extends CometTestBase with AdaptiveSparkPlanHelper CometConf.COMET_EXEC_TYPED_DATASET_MAP_ENABLED.key -> "true", CometConf.COMET_SCALA_UDF_CODEGEN_ENABLED.key -> "true") ++ pairs): _*)(f) - private def objectOperators(plan: SparkPlan): Seq[SparkPlan] = - collectWithSubqueries(plan) { - case p: SerializeFromObjectExec => p - case p: DeserializeToObjectExec => p - case p: MapElementsExec => p - case p: MapPartitionsExec => p + /** The common fixture: `rows` rows of `(a: Int, b: String)` via Parquet, read back typed. */ + private def withTypedRecs(rows: Int = 100)(f: Dataset[TypedRec] => Unit): Unit = + withParquetTable((0 until rows).map(i => (i, i.toString)), "tbl") { + f(spark.sql("select _1 as a, _2 as b from tbl").as[TypedRec]) } - private def assertNoObjectOperators(plan: SparkPlan): Unit = - assert( - objectOperators(plan).isEmpty, - s"expected the typed sandwich to be fused away, but plan still has " + - s"${objectOperators(plan).map(_.nodeName).mkString(", ")}:\n$plan") + /** + * Round-trips `df` through Parquet and registers it as `tbl`. Needed rather than + * `withParquetTable(df, name)` because that registers the DataFrame directly, leaving a + * `LocalTableScanExec` that the native-plan assertions do not tolerate. + */ + private def withParquetRoundTrip(df: => DataFrame)(f: => Unit): Unit = + withTempPath { path => + df.write.parquet(path.toString) + withParquetTable(path.toString, "tbl")(f) + } + + /** + * Every operator in the typed-Dataset family, via the two traits Spark uses to mark them. Wider + * than the `ds.map` sandwich on purpose: it also covers `FlatMapGroupsExec` / + * `AppendColumnsExec`, which this rule does not fuse, so a future change that starts fusing + * them is visible here. + */ + private def objectOperators(plan: SparkPlan): Seq[SparkPlan] = + collectWithSubqueries(plan) { case p @ (_: ObjectConsumerExec | _: ObjectProducerExec) => p } + + private def assertNoObjectOperators(plan: SparkPlan): Unit = { + val remaining = objectOperators(plan) + if (remaining.nonEmpty) { + fail( + "expected the typed sandwich to be fused away, but plan still has " + + s"${remaining.map(_.nodeName).mkString(", ")}:\n$plan") + } + } test("ds.map produces a fully native plan - single output column") { withFusion() { - withParquetTable((0 until 100).map(i => (i, i.toString)), "tbl") { - val ds = spark.sql("select _1 as a, _2 as b from tbl").as[TypedRec] + withTypedRecs() { ds => val (_, cometPlan) = checkSparkAnswerAndOperator(ds.map(_.a + 1).toDF()) assertNoObjectOperators(cometPlan) } @@ -86,8 +119,7 @@ class CometTypedDatasetSuite extends CometTestBase with AdaptiveSparkPlanHelper test("ds.map produces a fully native plan - multiple output columns") { withFusion() { - withParquetTable((0 until 100).map(i => (i, i.toString)), "tbl") { - val ds = spark.sql("select _1 as a, _2 as b from tbl").as[TypedRec] + withTypedRecs() { ds => val (_, cometPlan) = checkSparkAnswerAndOperator(ds.map(r => TypedRec(r.a + 1, r.b + "!")).toDF()) assertNoObjectOperators(cometPlan) @@ -99,30 +131,35 @@ class CometTypedDatasetSuite extends CometTestBase with AdaptiveSparkPlanHelper } } - test("output schema is unchanged by the rewrite") { - withParquetTable((0 until 20).map(i => (i, i.toString)), "tbl") { - // `withSQLConf` returns Unit on Spark 3.x, so capture rather than return from the block. - def schemaOf(fused: Boolean): String = { - var schema: String = null + test("executed-plan output attributes are unchanged by the rewrite") { + // `Dataset.schema` comes from the analyzed plan, which a physical rule cannot touch, so + // comparing it would be vacuous. The property that matters is that the rewritten projections + // reproduce `SerializeFromObjectExec.output` exactly -- names, types and nullability -- since + // parent operators reference those attributes. + withTypedRecs(20) { _ => + def executedOutput(fused: Boolean): String = { + var out: String = null withSQLConf(CometConf.COMET_EXEC_TYPED_DATASET_MAP_ENABLED.key -> fused.toString) { - schema = spark + out = spark .sql("select _1 as a, _2 as b from tbl") .as[TypedRec] .map(r => TypedRec(r.a + 1, r.b)) .toDF() - .schema - .treeString + .queryExecution + .executedPlan + .output + .map(a => s"${a.name}:${a.dataType.simpleString}:${a.nullable}") + .mkString(",") } - schema + out } - assert(schemaOf(fused = true) === schemaOf(fused = false)) + assert(executedOutput(fused = true) === executedOutput(fused = false)) } } test("closure runs exactly once per row with multiple output columns") { withFusion() { - withParquetTable((0 until 50).map(i => (i, i.toString)), "tbl") { - val ds = spark.sql("select _1 as a, _2 as b from tbl").as[TypedRec] + withTypedRecs(50) { ds => TypedMapCounter.reset() val fused = ds .map { r => @@ -140,6 +177,21 @@ class CometTypedDatasetSuite extends CometTestBase with AdaptiveSparkPlanHelper } } + test("a 12-column record still fuses (maxFields counts each input ordinal once)") { + // Every struct field carries the whole deserializer, so counting BoundReference occurrences + // rather than distinct ordinals made this 12 + 12*12 = 156 fields against a default + // `spark.sql.codegen.maxFields` of 100, and the rule silently declined. + withFusion() { + val cols = (1 to 12).map(i => s"_1 + $i as c$i").mkString(", ") + withParquetTable((0 until 20).map(i => (i, i.toString)), "tbl") { + val ds = spark.sql(s"select $cols from tbl").as[TypedWide12] + val (_, cometPlan) = + checkSparkAnswerAndOperator(ds.map(r => r.copy(c1 = r.c1 + 1)).toDF()) + assertNoObjectOperators(cometPlan) + } + } + } + test("fusion unblocks the aggregate and shuffle above it") { withFusion() { withParquetTable((0 until 100).map(i => (i, (i % 7).toString)), "tbl") { @@ -159,52 +211,39 @@ class CometTypedDatasetSuite extends CometTestBase with AdaptiveSparkPlanHelper s"s$i", new java.math.BigDecimal(s"$i.25"), if (i % 3 == 0) null else Long.box(i * 10L))) - withTempPath { path => - rows.toDF("i", "s", "d", "opt").write.parquet(path.toString) - withParquetTable(path.toString, "tbl") { - val ds = spark.table("tbl").as[TypedWide] - val (_, cometPlan) = checkSparkAnswerAndOperator( - ds.map(r => TypedWide(r.i + 1, r.s, r.d, r.opt.map(_ + 1))).toDF()) - assertNoObjectOperators(cometPlan) - } + withParquetRoundTrip(rows.toDF("i", "s", "d", "opt")) { + val ds = spark.table("tbl").as[TypedWide] + val (_, cometPlan) = checkSparkAnswerAndOperator( + ds.map(r => TypedWide(r.i + 1, r.s, r.d, r.opt.map(_ + 1))).toDF()) + assertNoObjectOperators(cometPlan) } } } test("nested struct and array fields round-trip") { withFusion() { - withTempPath { path => - (1 to 30) - .map(i => (i, (i, s"n$i"), Seq(s"t$i", s"u$i"))) - .toDF("id", "inner", "tags") - .selectExpr("id", "named_struct('a', inner._1, 'b', inner._2) as inner", "tags") - .write - .parquet(path.toString) - withParquetTable(path.toString, "tbl") { - val ds = spark.table("tbl").as[TypedNested] - val (_, cometPlan) = checkSparkAnswerAndOperator( - ds.map(r => TypedNested(r.id + 1, TypedRec(r.inner.a, r.inner.b), r.tags.reverse)) - .toDF()) - assertNoObjectOperators(cometPlan) - } + val nested = (1 to 30) + .map(i => (i, (i, s"n$i"), Seq(s"t$i", s"u$i"))) + .toDF("id", "inner", "tags") + .selectExpr("id", "named_struct('a', inner._1, 'b', inner._2) as inner", "tags") + withParquetRoundTrip(nested) { + val ds = spark.table("tbl").as[TypedNested] + val (_, cometPlan) = checkSparkAnswerAndOperator( + ds.map(r => TypedNested(r.id + 1, TypedRec(r.inner.a, r.inner.b), r.tags.reverse)) + .toDF()) + assertNoObjectOperators(cometPlan) } } } test("null field values in the input and the output") { withFusion() { - withTempPath { path => - (1 to 30) - .map(i => (i, if (i % 4 == 0) null else s"s$i")) - .toDF("a", "b") - .write - .parquet(path.toString) - withParquetTable(path.toString, "tbl") { - val ds = spark.table("tbl").as[TypedRec] - val (_, cometPlan) = checkSparkAnswerAndOperator( - ds.map(r => TypedRec(r.a, if (r.a % 5 == 0) null else r.b)).toDF()) - assertNoObjectOperators(cometPlan) - } + val withNulls = (1 to 30).map(i => (i, if (i % 4 == 0) null else s"s$i")).toDF("a", "b") + withParquetRoundTrip(withNulls) { + val ds = spark.table("tbl").as[TypedRec] + val (_, cometPlan) = checkSparkAnswerAndOperator( + ds.map(r => TypedRec(r.a, if (r.a % 5 == 0) null else r.b)).toDF()) + assertNoObjectOperators(cometPlan) } } } @@ -213,8 +252,7 @@ class CometTypedDatasetSuite extends CometTestBase with AdaptiveSparkPlanHelper // Returning null for a non-nullable top-level product is an error in Spark. The serializer's // `assertnotnull` has to survive the fuse, or Comet would silently emit a null row instead. withFusion() { - withParquetTable((0 until 20).map(i => (i, i.toString)), "tbl") { - val ds = spark.sql("select _1 as a, _2 as b from tbl").as[TypedRec] + withTypedRecs(20) { ds => val df = ds.map(r => if (r.a == 7) null else TypedRec(r.a, r.b)).toDF() assertNoObjectOperators(df.queryExecution.executedPlan) val err = intercept[Exception](df.collect()) @@ -261,8 +299,7 @@ class CometTypedDatasetSuite extends CometTestBase with AdaptiveSparkPlanHelper test("Java MapFunction takes the same path as a Scala closure") { withFusion() { - withParquetTable((0 until 40).map(i => (i, i.toString)), "tbl") { - val ds = spark.sql("select _1 as a, _2 as b from tbl").as[TypedRec] + withTypedRecs(40) { ds => val fn = new MapFunction[TypedRec, TypedRec] { override def call(r: TypedRec): TypedRec = TypedRec(r.a + 100, r.b) } @@ -275,8 +312,7 @@ class CometTypedDatasetSuite extends CometTestBase with AdaptiveSparkPlanHelper test("chained maps fuse into one native pipeline") { withFusion() { - withParquetTable((0 until 40).map(i => (i, i.toString)), "tbl") { - val ds = spark.sql("select _1 as a, _2 as b from tbl").as[TypedRec] + withTypedRecs(40) { ds => val (_, cometPlan) = checkSparkAnswerAndOperator( ds.map(r => TypedRec(r.a + 1, r.b)).map(r => TypedRec(r.a * 2, r.b + "x")).toDF()) assertNoObjectOperators(cometPlan) @@ -285,8 +321,7 @@ class CometTypedDatasetSuite extends CometTestBase with AdaptiveSparkPlanHelper } test("off by default") { - withParquetTable((0 until 20).map(i => (i, i.toString)), "tbl") { - val ds = spark.sql("select _1 as a, _2 as b from tbl").as[TypedRec] + withTypedRecs(20) { ds => val df = ds.map(r => TypedRec(r.a + 1, r.b)).toDF() df.collect() assert( @@ -299,13 +334,12 @@ class CometTypedDatasetSuite extends CometTestBase with AdaptiveSparkPlanHelper withSQLConf( CometConf.COMET_EXEC_TYPED_DATASET_MAP_ENABLED.key -> "true", CometConf.COMET_SCALA_UDF_CODEGEN_ENABLED.key -> "false") { - withParquetTable((0 until 20).map(i => (i, i.toString)), "tbl") { - val ds = spark.sql("select _1 as a, _2 as b from tbl").as[TypedRec] - val df = ds.map(r => TypedRec(r.a + 1, r.b)).toDF() - checkSparkAnswer(df) - assert( - objectOperators(df.queryExecution.executedPlan).nonEmpty, - "without the dispatcher there is nothing to fuse into") + withTypedRecs(20) { ds => + checkSparkAnswerAndFallbackReason( + ds.map(r => TypedRec(r.a + 1, r.b)).toDF(), + s"Cannot fuse typed Dataset map: " + + s"${CometConf.COMET_SCALA_UDF_CODEGEN_ENABLED.key}=false, so there is no dispatcher " + + "to fuse into") } } } @@ -314,21 +348,19 @@ class CometTypedDatasetSuite extends CometTestBase with AdaptiveSparkPlanHelper // Without CSE the single kernel would evaluate the shared closure Invoke once per struct // field, so the rule must decline rather than change how many times user code runs. withFusion(SQLConf.SUBEXPRESSION_ELIMINATION_ENABLED.key -> "false") { - withParquetTable((0 until 20).map(i => (i, i.toString)), "tbl") { - val ds = spark.sql("select _1 as a, _2 as b from tbl").as[TypedRec] - val df = ds.map(r => TypedRec(r.a + 1, r.b)).toDF() - checkSparkAnswer(df) - assert( - objectOperators(df.queryExecution.executedPlan).nonEmpty, - "multi-column fusion needs CSE to keep the closure to one call per row") + withTypedRecs(20) { ds => + checkSparkAnswerAndFallbackReason( + ds.map(r => TypedRec(r.a + 1, r.b)).toDF(), + "Cannot fuse typed Dataset map: 2 output columns require subexpression elimination " + + "to keep the closure to one call per row, but " + + s"${SQLConf.SUBEXPRESSION_ELIMINATION_ENABLED.key}=false") } } } test("mapPartitions is not fused") { withFusion() { - withParquetTable((0 until 20).map(i => (i, i.toString)), "tbl") { - val ds = spark.sql("select _1 as a, _2 as b from tbl").as[TypedRec] + withTypedRecs(20) { ds => val df = ds.mapPartitions(it => it.map(r => TypedRec(r.a + 1, r.b))).toDF() checkSparkAnswer(df) assert( From b9d68038f29d5d73f8371c8c3b63bdd359f038f5 Mon Sep 17 00:00:00 2001 From: Andy Grove Date: Mon, 21 Sep 2026 13:55:01 -0600 Subject: [PATCH 3/6] fix: repair the merge and tighten the typed-map fuse after review Merging main brought in a subtree-naming block in emitJvmCodegenDispatch that reads `target`, which this branch had moved into the new bindForDispatch helper. Return it from the helper rather than recomputing, so the tree that is bound and the tree that is walked for explain naming stay the same one. Also from the self-review pass: - forceDispatch declined silently via logDebug while every other decline path recorded a fallback reason. Route it through decline() so EXPLAIN says why the sandwich stayed on Spark. - Drop the now-unused Logging mixin from the rule. - Cover the maxFields accounting change from the dispatcher side, not just through the typed-map path: it is a change to the shared canHandle gate, so CometCodegenSuite gets a test that reads one column five times under maxFields=3. - Pin the one-kernel property with assertOneKernelForSubtree, which landed on main after this branch was cut. - Add a test that the rewritten plan is still correct when Comet declines the projection, which is the property that justifies a TreeNodeTag over a Comet Expression subclass. It asserts both that the sandwich was fused and that the projection is not native, so it cannot decay into a duplicate of the happy path. --- .../comet/rules/RewriteTypedDatasetMap.scala | 12 +++-- .../apache/comet/serde/CometScalaUDF.scala | 21 ++++----- .../org/apache/comet/CometCodegenSuite.scala | 22 +++++++++ .../apache/comet/CometTypedDatasetSuite.scala | 47 +++++++++++++++---- 4 files changed, 77 insertions(+), 25 deletions(-) diff --git a/spark/src/main/scala/org/apache/comet/rules/RewriteTypedDatasetMap.scala b/spark/src/main/scala/org/apache/comet/rules/RewriteTypedDatasetMap.scala index 8f8488d931e..32d7b2e1d1a 100644 --- a/spark/src/main/scala/org/apache/comet/rules/RewriteTypedDatasetMap.scala +++ b/spark/src/main/scala/org/apache/comet/rules/RewriteTypedDatasetMap.scala @@ -20,7 +20,6 @@ package org.apache.comet.rules import org.apache.spark.api.java.function.MapFunction -import org.apache.spark.internal.Logging import org.apache.spark.sql.catalyst.expressions.{Alias, AttributeSet, BoundReference, CreateNamedStruct, Expression, GetStructField, Literal, NamedExpression} import org.apache.spark.sql.catalyst.expressions.objects.Invoke import org.apache.spark.sql.catalyst.plans.logical.FunctionUtils @@ -74,7 +73,7 @@ import org.apache.comet.serde.{CometScalaUDF, QueryPlanSerde} * of adjacent `MapElementsExec` nodes is fused as a whole, because `ds.map(f).map(g)` leaves both * under a single Serialize/Deserialize pair. */ -object RewriteTypedDatasetMap extends Logging { +object RewriteTypedDatasetMap { /** Name of the single intermediate struct column when there is more than one output column. */ private val FUSED_COLUMN = "comet_fused_object" @@ -192,10 +191,11 @@ object RewriteTypedDatasetMap extends Logging { fused match { // One output column is already one kernel; the struct wrapper would have nothing to dedupe. case Seq(only) => - forceDispatch(only.children.head).map(_ => ProjectExec(fused, child)) + forceDispatch(serialize, only.children.head).map(_ => ProjectExec(fused, child)) case _ => forceDispatch( + serialize, CreateNamedStruct(fused.flatMap(ne => Seq(Literal(ne.name), ne.children.head)))).map { structExpr => val structAlias = Alias(structExpr, FUSED_COLUMN)() @@ -222,10 +222,12 @@ object RewriteTypedDatasetMap extends Logging { * for an unfused one. Delegating to [[CometScalaUDF.canDispatch]] rather than re-deriving the * binding is what keeps the prediction identical to what the serde will really do. */ - private def forceDispatch[T <: Expression](expr: T): Option[T] = + private def forceDispatch[T <: Expression]( + serialize: SerializeFromObjectExec, + expr: T): Option[T] = CometScalaUDF.canDispatch(expr) match { case Some(reason) => - logDebug(s"RewriteTypedDatasetMap: not rewriting because $reason") + decline(serialize, s"the fused expression is not dispatchable ($reason)") None case None => expr.setTagValue(QueryPlanSerde.FORCE_DISPATCH, ()) diff --git a/spark/src/main/scala/org/apache/comet/serde/CometScalaUDF.scala b/spark/src/main/scala/org/apache/comet/serde/CometScalaUDF.scala index bc7cb06719c..9088095d3c1 100644 --- a/spark/src/main/scala/org/apache/comet/serde/CometScalaUDF.scala +++ b/spark/src/main/scala/org/apache/comet/serde/CometScalaUDF.scala @@ -59,13 +59,15 @@ object CometScalaUDF extends CometExpressionSerde[ScalaUDF] { emitJvmCodegenDispatch(expr, inputs, binding) /** - * Bind `expr` the way [[emitJvmCodegenDispatch]] will, returning the attributes it reads in - * ordinal order alongside the bound tree. + * Bind `expr` the way [[emitJvmCodegenDispatch]] will: the tree actually bound (`expr` itself, + * or its `replacement` when it is `RuntimeReplaceable`), the attributes it reads in ordinal + * order, and the bound tree. * * Callers that only need to know whether a tree is dispatchable should use [[canDispatch]] * rather than repeating this. */ - private def bindForDispatch(expr: Expression): (Seq[AttributeReference], Expression) = { + private def bindForDispatch( + expr: Expression): (Expression, Seq[AttributeReference], Expression) = { // `RuntimeReplaceable` expressions (e.g. Spark 4's `StructsToJson`) have a `doGenCode` that // always throws "Cannot generate code for expression". Catalyst's `ReplaceExpressions` rule // normally rewrites them to their `replacement` form before codegen runs. Comet's serde @@ -78,7 +80,7 @@ object CometScalaUDF extends CometExpressionSerde[ScalaUDF] { // Bind against only the AttributeReferences the tree actually reads, so ordinals align with // the data args we ship. val attrs = target.collect { case a: AttributeReference => a }.distinct - (attrs, BindReferences.bindReference(target, AttributeSeq(attrs))) + (target, attrs, BindReferences.bindReference(target, AttributeSeq(attrs))) } /** @@ -94,7 +96,7 @@ object CometScalaUDF extends CometExpressionSerde[ScalaUDF] { return Some( s"${CometConf.COMET_SCALA_UDF_CODEGEN_ENABLED.key}=false; expression has no native path") } - CometBatchKernelCodegen.canHandle(bindForDispatch(expr)._2) + CometBatchKernelCodegen.canHandle(bindForDispatch(expr)._3) } /** @@ -121,12 +123,9 @@ object CometScalaUDF extends CometExpressionSerde[ScalaUDF] { return None } - // `RuntimeReplaceable` expressions (e.g. Spark 4's `StructsToJson`) have a `doGenCode` that - // always throws "Cannot generate code for expression". Catalyst's `ReplaceExpressions` rule - // normally rewrites them to their `replacement` form before codegen runs. Comet's serde - // sometimes works with the pre-rewrite form (via shim reconstruction) for matching purposes, - // so unwrap to the replacement here before binding so the kernel compiles. - val (attrs, boundExpr) = bindForDispatch(expr) + // `target` is `expr` unwrapped past `RuntimeReplaceable`; see `bindForDispatch`. Everything + // below that reasons about the tree that will actually be compiled must use it, not `expr`. + val (target, attrs, boundExpr) = bindForDispatch(expr) // Gate at plan time. Surface the reason via withFallbackReason rather than crashing Janino // at execute. diff --git a/spark/src/test/scala/org/apache/comet/CometCodegenSuite.scala b/spark/src/test/scala/org/apache/comet/CometCodegenSuite.scala index 149ee7a1454..619ce89b7f8 100644 --- a/spark/src/test/scala/org/apache/comet/CometCodegenSuite.scala +++ b/spark/src/test/scala/org/apache/comet/CometCodegenSuite.scala @@ -247,6 +247,28 @@ class CometCodegenSuite } } + test("maxFields counts each input ordinal once, not each read of it") { + // WSCG gates on the operator's *schema* (`plan.schema.map(...).sum`), so a column read more + // than once contributes one field, and the kernel likewise emits one typed field and one + // getter per ordinal rather than per occurrence. Counting `BoundReference` occurrences instead + // made the total scale with how often the tree happens to repeat a column: five reads of one + // Int column plus a 1-field output came to 6 and was refused, where the kernel is really 2 + // fields wide. Same UDF arity as the test above, so the two differ only in how many distinct + // columns they read. + spark.udf.register( + "sumFiveInts", + (a: Int, b: Int, c: Int, d: Int, e: Int) => a + b + c + d + e) + withTable("t") { + sql("CREATE TABLE t (a INT) USING parquet") + sql("INSERT INTO t VALUES (1), (10)") + withSQLConf("spark.sql.codegen.maxFields" -> "3") { + assertCodegenRan { + checkSparkAnswerAndOperator(sql("SELECT sumFiveInts(a, a, a, a, a) FROM t")) + } + } + } + } + test("explain.codegen.enabled surfaces routed expressions in COMET-INFO") { // With the opt-in flag on, `hypot` and `nanvl` (both `CometCodegenDispatch`) roll up // into one `[COMET-INFO: JVM codegen dispatcher: hypot, nanvl]` line on the diff --git a/spark/src/test/scala/org/apache/comet/CometTypedDatasetSuite.scala b/spark/src/test/scala/org/apache/comet/CometTypedDatasetSuite.scala index 82076235b57..b5c7f717a87 100644 --- a/spark/src/test/scala/org/apache/comet/CometTypedDatasetSuite.scala +++ b/spark/src/test/scala/org/apache/comet/CometTypedDatasetSuite.scala @@ -63,7 +63,7 @@ object TypedMapCounter { * `SerializeFromObject` / `MapElements` / `DeserializeToObject` sandwich a typed `Dataset.map` * produces into a Comet projection. See https://github.com/apache/datafusion-comet/issues/5710. */ -class CometTypedDatasetSuite extends CometTestBase { +class CometTypedDatasetSuite extends CometTestBase with CometCodegenAssertions { import testImplicits._ @@ -120,13 +120,18 @@ class CometTypedDatasetSuite extends CometTestBase { test("ds.map produces a fully native plan - multiple output columns") { withFusion() { withTypedRecs() { ds => - val (_, cometPlan) = - checkSparkAnswerAndOperator(ds.map(r => TypedRec(r.a + 1, r.b + "!")).toDF()) - assertNoObjectOperators(cometPlan) - // Inner projection builds the struct, outer one unpacks it. - assert( - collectWithSubqueries(cometPlan) { case p: CometProjectExec => p }.size >= 2, - s"expected stacked projections for the struct fuse:\n$cometPlan") + // The struct wrapper exists so the N serializer expressions compile into one kernel + // rather than N. `assertOneKernelForSubtree` pins that directly; the closure-call counter + // test below pins the consequence that matters to a user. + assertOneKernelForSubtree { + val (_, cometPlan) = + checkSparkAnswerAndOperator(ds.map(r => TypedRec(r.a + 1, r.b + "!")).toDF()) + assertNoObjectOperators(cometPlan) + // Inner projection builds the struct, outer one unpacks it. + assert( + collectWithSubqueries(cometPlan) { case p: CometProjectExec => p }.size >= 2, + s"expected stacked projections for the struct fuse:\n$cometPlan") + } } } } @@ -320,6 +325,30 @@ class CometTypedDatasetSuite extends CometTestBase { } } + test("rewritten plan is still correct when Comet does not take the projection") { + // The rewrite commits to the fused shape before knowing whether the projection will convert. + // Building it from stock Spark expressions (a `TreeNodeTag` rather than a Comet `Expression` + // subclass) is what makes that safe: when the projection stays on Spark, whole-stage codegen + // compiles the same fused tree it would have built from the sandwich anyway. Turning the scan + // off is the cheapest way to strand the projection on Spark. + withFusion(CometConf.COMET_NATIVE_SCAN_ENABLED.key -> "false") { + withTypedRecs(20) { ds => + Seq(ds.map(r => TypedRec(r.a + 1, r.b + "!")).toDF(), ds.map(_.a + 1).toDF()) + .foreach { df => + checkSparkAnswer(df) + val plan = df.queryExecution.executedPlan + // Non-vacuous only if the rewrite really did fire and Comet really did not take the + // result. Both halves are asserted so a future change that makes either untrue shows + // up here rather than silently turning this into a duplicate of the happy path. + assertNoObjectOperators(plan) + assert( + collectWithSubqueries(plan) { case p: CometProjectExec => p }.isEmpty, + s"expected the fused projection to stay on Spark without a native scan:\n$plan") + } + } + } + } + test("off by default") { withTypedRecs(20) { ds => val df = ds.map(r => TypedRec(r.a + 1, r.b)).toDF() @@ -337,7 +366,7 @@ class CometTypedDatasetSuite extends CometTestBase { withTypedRecs(20) { ds => checkSparkAnswerAndFallbackReason( ds.map(r => TypedRec(r.a + 1, r.b)).toDF(), - s"Cannot fuse typed Dataset map: " + + "Cannot fuse typed Dataset map: " + s"${CometConf.COMET_SCALA_UDF_CODEGEN_ENABLED.key}=false, so there is no dispatcher " + "to fuse into") } From 8429ff3a4da6c3fe5e4ccee767ada832bcc25c91 Mon Sep 17 00:00:00 2001 From: Andy Grove Date: Mon, 21 Sep 2026 15:39:15 -0600 Subject: [PATCH 4/6] bench: measure where the typed Dataset map fuse actually pays off The flag is off by default because the PR had no numbers. This is the benchmark that produces them: three arms (fuse off, fuse on, Comet disabled) over six plan shapes, organised by what sits above the map, since that was the hypothesis worth testing. It contradicts the hypothesis. The claimed win -- a typed map between native operators -- does not reproduce, and `map -> filter -> group by` is a large and repeatable loss. The one consistent win is the multi-column top-of-plan shape the PR predicted would be neutral. Numbers and the reading of them are in the PR description. Sized at 4Mi rows by the noise detector, not by taste: at 128Ki the `fuse off (repeat)` row came back up to 27% from the row it duplicates, which is wider than any difference between the arms. --- .../CometTypedDatasetMapBenchmark.scala | 351 ++++++++++++++++++ 1 file changed, 351 insertions(+) create mode 100644 spark/src/test/scala/org/apache/spark/sql/benchmark/CometTypedDatasetMapBenchmark.scala diff --git a/spark/src/test/scala/org/apache/spark/sql/benchmark/CometTypedDatasetMapBenchmark.scala b/spark/src/test/scala/org/apache/spark/sql/benchmark/CometTypedDatasetMapBenchmark.scala new file mode 100644 index 00000000000..44fd5278078 --- /dev/null +++ b/spark/src/test/scala/org/apache/spark/sql/benchmark/CometTypedDatasetMapBenchmark.scala @@ -0,0 +1,351 @@ +/* + * 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 java.nio.charset.StandardCharsets + +import org.apache.spark.benchmark.Benchmark +import org.apache.spark.sql.{DataFrame, Dataset, Encoder, Encoders, Row} +import org.apache.spark.sql.functions.{col, sum} +import org.apache.spark.sql.internal.SQLConf + +import org.apache.comet.CometConf + +/** Top-level so `NewInstance` needs no outer pointer, which is the ordinary user shape. */ +case class TypedMapRec(a: Long, b: String) + +/** + * Benchmark of `RewriteTypedDatasetMap`, which fuses the `SerializeFromObject` / `MapElements` / + * `DeserializeToObject` sandwich a typed `Dataset.map` produces into a Comet projection routed + * through the JVM codegen dispatcher. See https://github.com/apache/datafusion-comet/issues/5710. + * + * The question this exists to answer is where the rewrite pays for itself, because + * `spark.comet.exec.typedDatasetMap.enabled` is off by default until it has an answer. Three + * arms: + * + * - `fuse off` -- today's default. The sandwich falls back to Spark, and the fallback cascades + * to whatever sits above it. + * - `fuse on` -- the rewrite. The user closure still runs on the JVM, once per row, but inside + * a Janino-compiled kernel reading and writing Arrow vectors, so the operators above the map + * stay native. + * - `Spark (Comet disabled)` -- a reference point, not the comparison of interest. + * + * '''The shape of the plan above the map is the whole story,''' so the cases are organised by it + * rather than by the closure. When the map is at the top of the plan there is nothing above it to + * rescue: the kernel writes Arrow only for the sink to read rows straight back out, and the + * rewrite is expected to break even at best. When an aggregate or a shuffle sits above it, those + * operators are what the fallback was costing, and they are what the rewrite buys back. + * + * Both Comet arms are consumed with `noop()`, which reads rows. That is not neutral between them + * and is not meant to be: with the fuse off the plan is already row-based at the top, while with + * it on the top is columnar and pays a columnar-to-row transition at the sink. That transition is + * exactly the top-of-plan cost the feature flag exists to guard against, so charging the fused + * arm for it is the honest measurement, not a confound. + * + * A `fuse off (repeat)` case repeats the baseline at the end of every table. It measures the same + * work as the first row, so the spread between the two is this machine's noise floor for that + * table; ignore any difference between the other rows that is smaller than that spread. This + * matters more than usual here, because the arm that loses operators to Spark is the one most + * sensitive to a busy machine. + * + * Before timing anything, every case is run through all three arms and the rows are compared, so + * a timing cannot come from an arm that computed something else. Each case's plans are also + * checked: a `fuse on` arm that is not fully native, or a `fuse off` arm that is, would not be + * measuring what its name says. + * + * To run this benchmark: + * {{{ + * SPARK_GENERATE_BENCHMARK_FILES=1 make benchmark-org.apache.spark.sql.benchmark.CometTypedDatasetMapBenchmark + * }}} + * Results will be written to "spark/benchmarks/CometTypedDatasetMapBenchmark-**results.txt". + */ +object CometTypedDatasetMapBenchmark extends CometBenchmarkBase { + + /** + * ~512 batches at the default `spark.comet.batchSize`, so per-batch costs are amortised. + * + * Sized by the noise detector rather than by taste. At 128Ki rows an iteration of every case + * here lands in the 15-25ms range on an M3 Max, and the `fuse off (repeat)` row came back up to + * 27% away from the `fuse off` row it duplicates -- a spread wider than any difference between + * the arms, which makes the whole table unreadable. Lower this on a smaller machine only if the + * repeat row still agrees with the baseline afterwards. + */ + private val Rows = 4 * 1024 * 1024 + + // Declared rather than pulled in with `import spark.implicits._`, which would force the session + // to start while this object is still initialising. + private implicit val recEncoder: Encoder[TypedMapRec] = Encoders.product[TypedMapRec] + private implicit val longEncoder: Encoder[Long] = Encoders.scalaLong + + /** + * @param name + * Case name, as it appears in the results table. + * @param build + * Builds the query from the typed Dataset. Kept as a function rather than SQL because there + * is no SQL spelling of a typed `map`. + * @param fusesToNative + * False for a case the rewrite is expected to decline, where the point of the row is to show + * the decline costs nothing rather than to show a speedup. + * @param extraConfigs + * Applied to all three arms, so they never account for a difference between them. + */ + private case class MapCase( + name: String, + build: Dataset[TypedMapRec] => DataFrame, + fusesToNative: Boolean = true, + extraConfigs: Seq[(String, String)] = Nil) + + /** + * AQE off and one shuffle partition for every case that shuffles, so both arms plan the same + * shape on every iteration and the plan check is not looking at an `AQEShuffleRead`. + */ + private val deterministicShuffle = + Seq(SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "false", SQLConf.SHUFFLE_PARTITIONS.key -> "1") + + private def cases: Seq[MapCase] = Seq( + // Nothing above the map. The rewrite has no operator to rescue and pays a columnar-to-row + // transition at the sink that the unfused plan does not, so this is the case the default-off + // decision rests on. One output column takes the direct path: a single fused expression, no + // struct wrapper. + MapCase("map -> sink, 1 col", _.map(_.a + 1).toDF()), + // Two output columns take the `CreateNamedStruct` + `GetStructField` path, which is a second + // projection and a struct round-trip through Arrow. Worth separating from the row above: if + // the struct path is much worse at the top of the plan, that is a cost of the multi-column + // encoding rather than of fusing as such. + MapCase("map -> sink, 2 cols", _.map(r => TypedMapRec(r.a + 1, r.b)).toDF()), + // The case the rewrite exists for. With the fuse off, the fallback island takes the partial + // aggregate, the exchange and the final aggregate down with it. `b` has 100 distinct values, + // so the grouping is cheap and the per-row closure still dominates. Both output columns are + // consumed so that column pruning cannot give the two arms different work to do. + MapCase( + "map -> group by", + _.map(r => TypedMapRec(r.a + 1, r.b)).groupBy("b").agg(sum("a")), + extraConfigs = deterministicShuffle), + // A filter between the map and the aggregate, so the rewrite is rescuing three operator kinds + // rather than two and the fused projection feeds a native filter directly. + MapCase( + "map -> filter -> group by", + _.map(r => TypedMapRec(r.a * 2, r.b)) + .filter(col("a") % 3 === 0) + .groupBy("b") + .agg(sum("a")), + extraConfigs = deterministicShuffle), + // `ds.map(f).map(g)` leaves two adjacent `MapElements` under one Serialize/Deserialize pair, + // which the rule fuses as a whole. Two closure calls per row against one bridge crossing, so + // the bridge is amortised further here than anywhere else in the table. + MapCase( + "map -> map -> group by", + _.map(r => TypedMapRec(r.a + 1, r.b)) + .map(r => TypedMapRec(r.a * 2, r.b)) + .groupBy("b") + .agg(sum("a")), + extraConfigs = deterministicShuffle), + // `mapPartitions` is iterator-shaped, so there is no per-row expression and the rule declines. + // Included as a control: both Comet arms should land on the same plan and the same time, and a + // difference between them would mean the rewrite is reaching something it should not. + MapCase( + "mapPartitions (declined)", + _.mapPartitions(it => it.map(r => TypedMapRec(r.a + 1, r.b))).groupBy("b").agg(sum("a")), + fusesToNative = false, + extraConfigs = deterministicShuffle)) + + private def sparkConfigs(c: MapCase): Seq[(String, String)] = + Seq(CometConf.COMET_ENABLED.key -> "false") ++ c.extraConfigs + + /** Comet on, rewrite on: the behaviour this benchmark is validating. */ + private def fusedConfigs(c: MapCase): Seq[(String, String)] = cometConfigs(c, fuse = true) + + /** Comet on, rewrite off: today's default, where the typed map is a Spark fallback island. */ + private def unfusedConfigs(c: MapCase): Seq[(String, String)] = cometConfigs(c, fuse = false) + + private def cometConfigs(c: MapCase, fuse: Boolean): Seq[(String, String)] = Seq( + CometConf.COMET_ENABLED.key -> "true", + CometConf.COMET_EXEC_ENABLED.key -> "true", + // The rewrite has no dispatcher to fuse into without this, and it is on by default; set it + // explicitly in both Comet arms so the table does not depend on that default. + CometConf.COMET_SCALA_UDF_CODEGEN_ENABLED.key -> "true", + CometConf.COMET_EXEC_TYPED_DATASET_MAP_ENABLED.key -> fuse.toString) ++ c.extraConfigs + + private val FusedCaseName = "Comet, typed map fused" + private val UnfusedCaseName = "Comet, fuse off (Spark fallback island)" + private val SparkCaseName = "Spark (Comet disabled)" + + override def runCometBenchmark(mainArgs: Array[String]): Unit = { + val selected = cases + runBenchmark("Typed Dataset map fusion: environment") { + emitEnvironment(selected) + } + withCorpus(Rows) { + selected.foreach(verifyArmsAgree) + selected.foreach(runSteadyState) + } + } + + /** + * The reader of a results file cannot see the confs the numbers were produced under, and for + * this benchmark the feature flag and the batch size are the whole point. + */ + private def emitEnvironment(selected: Seq[MapCase]): Unit = { + emit(s"Spark version: ${spark.version}") + emit( + s"Java version: ${System.getProperty("java.version")} " + + s"(${System.getProperty("java.vm.name")})") + emit(s"Scala version: ${scala.util.Properties.versionNumberString}") + emit(s"spark.master: ${spark.conf.get("spark.master", "")}") + emit( + s"${CometConf.COMET_BATCH_SIZE.key}: " + + CometConf.COMET_BATCH_SIZE.get(spark.sessionState.conf)) + emit( + s"${SQLConf.WHOLESTAGE_CODEGEN_ENABLED.key}: " + + spark.conf.get(SQLConf.WHOLESTAGE_CODEGEN_ENABLED.key)) + emit(s"Feature flag: ${CometConf.COMET_EXEC_TYPED_DATASET_MAP_ENABLED.key}") + emit(s"Dispatcher conf: ${CometConf.COMET_SCALA_UDF_CODEGEN_ENABLED.key} (true in both arms)") + emit("Steady-state tables: Spark's Benchmark defaults -- 2s of untimed warmup per case, then") + emit( + " at least 2 iterations and at least 2s of timed iterations; the table reports best and") + emit(" average of the timed iterations.") + emit(s"Rows: $Rows. Grouping key `b` has 100 distinct values.") + emit(s"Cases: ${selected.map(_.name).mkString(", ")}") + } + + /** + * Fails if the three arms disagree on a case. Rows are compared as a sorted multiset rather + * than positionally, because the grouped cases shuffle and their output order is a property of + * the plan, which is the one thing that differs between the arms. + */ + private def verifyArmsAgree(c: MapCase): Unit = { + def collect(configs: Seq[(String, String)]): Array[String] = { + // Assigned to a local rather than returned from the block: Spark 3.4 and 3.5 declare + // `SQLHelper.withSQLConf` as returning `Unit`; only Spark 4 has the result-returning form. + var collected: Array[Row] = Array.empty + withSQLConf(configs: _*) { + collected = c.build(typedDataset).collect() + } + collected.map(_.toSeq.map(String.valueOf).mkString("|")).sorted + } + + val expected = collect(sparkConfigs(c)) + Seq(FusedCaseName -> fusedConfigs(c), UnfusedCaseName -> unfusedConfigs(c)).foreach { + case (armName, configs) => + val actual = collect(configs) + assert( + expected.length == actual.length, + s"${c.name}: Spark produced ${expected.length} rows, $armName ${actual.length}") + expected.indices.find(i => expected(i) != actual(i)).foreach { i => + throw new AssertionError( + s"${c.name}: row $i differs -- Spark ${expected(i)}, $armName ${actual(i)}") + } + } + } + + private def runSteadyState(c: MapCase): Unit = { + runBenchmark(s"${c.name} -- $Rows rows") { + val benchmark = new Benchmark(s"${c.name} -- $Rows rows", Rows, output = output) + checkPlans(benchmark, c) + // The unfused arm goes first so the `Relative` column reads as the speedup this change buys + // over the behaviour that ships today. + benchmark.addCase(UnfusedCaseName)(_ => runCase(c, unfusedConfigs(c))) + benchmark.addCase(FusedCaseName)(_ => runCase(c, fusedConfigs(c))) + benchmark.addCase(SparkCaseName)(_ => runCase(c, sparkConfigs(c))) + benchmark.addCase(s"$UnfusedCaseName (repeat)")(_ => runCase(c, unfusedConfigs(c))) + benchmark.run() + } + } + + /** + * Warns rather than fails, so one Spark version planning a case differently degrades the table + * to a note instead of aborting the run. + */ + private def checkPlans(benchmark: Benchmark, c: MapCase): Unit = { + var fusedNonComet: Option[String] = None + var unfusedIsFullyComet = false + withSQLConf(fusedConfigs(c): _*) { + val df = c.build(typedDataset) + df.noop() + fusedNonComet = + findFirstNonCometOperator(stripAQEPlan(df.queryExecution.executedPlan)).map(_.nodeName) + } + withSQLConf(unfusedConfigs(c): _*) { + val df = c.build(typedDataset) + df.noop() + unfusedIsFullyComet = + findFirstNonCometOperator(stripAQEPlan(df.queryExecution.executedPlan)).isEmpty + } + if (c.fusesToNative) { + fusedNonComet.foreach(op => + warn( + benchmark, + "WARNING: the fused plan is not fully Comet native (first non-Comet operator: " + + s"$op), so that case is partly measuring Spark.")) + if (unfusedIsFullyComet) { + warn( + benchmark, + "WARNING: the fuse-off plan is fully Comet native, so this case is not exercising " + + "the fallback island it is meant to be compared against.") + } + } else if (fusedNonComet.isEmpty) { + warn( + benchmark, + "WARNING: this case is supposed to be declined by the rewrite, but its fused plan is " + + "fully Comet native, so the two Comet arms are not the control they claim to be.") + } + } + + private def runCase(c: MapCase, configs: Seq[(String, String)]): Unit = + withSQLConf(configs: _*) { + c.build(typedDataset).noop() + } + + /** + * The corpus read back as a typed `Dataset`. Rebuilt per call rather than cached, because the + * plan it produces has to be built under the arm's own confs. + */ + private def typedDataset: Dataset[TypedMapRec] = + spark.sql("select c_a as a, c_b as b from parquetV1Table").as[TypedMapRec] + + /** Builds `parquetV1Table` with `rows` rows of the corpus and drops it afterwards. */ + private def withCorpus(rows: Int)(f: => Unit): Unit = { + withTempPath { dir => + withTempTable(tbl, "parquetV1Table") { + spark.range(rows).createOrReplaceTempView(tbl) + // `c_a` varies per row so the closure's arithmetic is not loop-invariant; `c_b` has 100 + // distinct values so the grouped cases spend their time in the map, not the aggregate. + prepareTable( + dir, + spark.sql(s"SELECT id AS c_a, CAST(PMOD(id, 100) AS STRING) AS c_b FROM $tbl")) + f + } + } + } + + /** Writes a warning to the results file as well as the console, ordered against the table. */ + private def warn(benchmark: Benchmark, message: String): Unit = { + val border = "=" * 80 + benchmark.out.println(s"\n$border\n$message\n$border") + } + + /** [[Benchmark]] tees console and results file; this benchmark's own tables need the same. */ + private def emit(line: String): Unit = { + // scalastyle:off println + println(line) + // scalastyle:on println + output.foreach(_.write(s"$line\n".getBytes(StandardCharsets.UTF_8))) + } +} From 92c11a8a16dd51eb75271867461a9dc4d4d04dcd Mon Sep 17 00:00:00 2001 From: Andy Grove Date: Mon, 21 Sep 2026 16:38:31 -0600 Subject: [PATCH 5/6] bench: run each grouped shape at both low and high key cardinality The first version of this benchmark grouped on 100 keys so that "the per-row call is still what dominates". That choice decided the result: the partial aggregate is the one operator fusing actually rescues, and at 100 keys it is nearly free, so the table could only ever say the rewrite does not pay. Running the same shapes at 1Mi keys says the opposite -- 2.0x on `map -> group by`, and `map -> filter -> group by` flips from 0.5x to 1.4x. Both reproduce across two runs with tight repeat rows. Also corrects the class comment, which still claimed the fallback cascades to everything above the island. It does not on current main: the fuse-off plan already keeps the exchange and the final aggregate native, so what fusing buys is the partial aggregate plus a columnar-to-native shuffle upgrade. --- .../CometTypedDatasetMapBenchmark.scala | 124 ++++++++++++------ 1 file changed, 86 insertions(+), 38 deletions(-) diff --git a/spark/src/test/scala/org/apache/spark/sql/benchmark/CometTypedDatasetMapBenchmark.scala b/spark/src/test/scala/org/apache/spark/sql/benchmark/CometTypedDatasetMapBenchmark.scala index 44fd5278078..42aea540464 100644 --- a/spark/src/test/scala/org/apache/spark/sql/benchmark/CometTypedDatasetMapBenchmark.scala +++ b/spark/src/test/scala/org/apache/spark/sql/benchmark/CometTypedDatasetMapBenchmark.scala @@ -40,18 +40,22 @@ case class TypedMapRec(a: Long, b: String) * `spark.comet.exec.typedDatasetMap.enabled` is off by default until it has an answer. Three * arms: * - * - `fuse off` -- today's default. The sandwich falls back to Spark, and the fallback cascades - * to whatever sits above it. + * - `fuse off` -- today's default, where the sandwich is a Spark fallback island. * - `fuse on` -- the rewrite. The user closure still runs on the JVM, once per row, but inside * a Janino-compiled kernel reading and writing Arrow vectors, so the operators above the map * stay native. * - `Spark (Comet disabled)` -- a reference point, not the comparison of interest. * - * '''The shape of the plan above the map is the whole story,''' so the cases are organised by it - * rather than by the closure. When the map is at the top of the plan there is nothing above it to - * rescue: the kernel writes Arrow only for the sink to read rows straight back out, and the - * rewrite is expected to break even at best. When an aggregate or a shuffle sits above it, those - * operators are what the fallback was costing, and they are what the rewrite buys back. + * '''Know what is actually at stake before reading a row.''' The premise behind the rewrite was + * that the island's fallback cascades and costs every operator above it. On current main it + * mostly does not: for `map -> group by` the fuse-off plan already keeps the exchange and the + * final aggregate native, because Comet re-enters above the island via `CometColumnarShuffle`. + * What fusing actually buys is the '''partial''' aggregate and an upgrade from + * `CometColumnarShuffle` to `CometNativeShuffle`. So the cases are organised around how much that + * partial aggregate is worth -- each grouped shape is run at both [[LoCardKeys]] and + * [[HiCardKeys]] distinct keys -- and around what fusing has to pay for it: a struct vector built + * and immediately unpacked, and every row materialised into Arrow even when a selective operator + * above is about to discard it. * * Both Comet arms are consumed with `noop()`, which reads rows. That is not neutral between them * and is not meant to be: with the fuse off the plan is already row-based at the top, while with @@ -100,6 +104,10 @@ object CometTypedDatasetMapBenchmark extends CometBenchmarkBase { * @param build * Builds the query from the typed Dataset. Kept as a function rather than SQL because there * is no SQL spelling of a typed `map`. + * @param hiCard + * Groups on a key with [[HiCardKeys]] distinct values instead of [[LoCardKeys]]. The partial + * aggregate is the one operator fusing actually rescues (see the class comment), so its cost + * is the variable that decides whether the rewrite can pay for itself at all. * @param fusesToNative * False for a case the rewrite is expected to decline, where the point of the row is to show * the decline costs nothing rather than to show a speedup. @@ -109,9 +117,16 @@ object CometTypedDatasetMapBenchmark extends CometBenchmarkBase { private case class MapCase( name: String, build: Dataset[TypedMapRec] => DataFrame, + hiCard: Boolean = false, fusesToNative: Boolean = true, extraConfigs: Seq[(String, String)] = Nil) + /** + * Distinct grouping keys. At [[Rows]] rows the high-cardinality key averages 4 rows a group. + */ + private val LoCardKeys = 100 + private val HiCardKeys = 1024 * 1024 + /** * AQE off and one shuffle partition for every case that shuffles, so both arms plan the same * shape on every iteration and the plan check is not looking at an `AQEShuffleRead`. @@ -119,34 +134,60 @@ object CometTypedDatasetMapBenchmark extends CometBenchmarkBase { private val deterministicShuffle = Seq(SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "false", SQLConf.SHUFFLE_PARTITIONS.key -> "1") + private def groupBySum(ds: Dataset[TypedMapRec]): DataFrame = + ds.map(r => TypedMapRec(r.a + 1, r.b)).groupBy("b").agg(sum("a")) + + private def filterGroupBySum(ds: Dataset[TypedMapRec]): DataFrame = + ds.map(r => TypedMapRec(r.a * 2, r.b)).filter(col("a") % 3 === 0).groupBy("b").agg(sum("a")) + + // Ordered so each low-cardinality case sits next to its high-cardinality twin and the pair is + // measured under similar machine conditions, and so the decisive pairs run before the tables + // that contention is most likely to spoil. Spark's `Benchmark` warms every case separately, so + // ordering does not otherwise affect a reading. private def cases: Seq[MapCase] = Seq( + // The case the rewrite exists for. On current main the fuse-off plan already keeps the + // exchange and the *final* aggregate native -- Comet re-enters above the island -- so the only + // operator fusing rescues here is the partial aggregate, plus an upgrade from + // `CometColumnarShuffle` to `CometNativeShuffle`. With 100 keys that partial aggregate is + // nearly free, which is why this row says so little. Both output columns are consumed so + // column pruning cannot give the two arms different work to do. + MapCase( + s"map -> group by, $LoCardKeys keys", + groupBySum, + extraConfigs = deterministicShuffle), + // The same shape with a partial aggregate that actually costs something: a million-entry hash + // table instead of a hundred. This is the fair test of the PR's premise. If fusing cannot win + // here it cannot win on aggregate rescue at all, because this is the most the rescued operator + // can be worth. + MapCase( + s"map -> group by, $HiCardKeys keys", + groupBySum, + hiCard = true, + extraConfigs = deterministicShuffle), + // A filter between the map and the aggregate. The fused kernel must materialise every row into + // Arrow -- including the string column -- before `CometFilter` discards two thirds of them, + // where whole-stage codegen drops them inside the same row loop. + MapCase( + s"map -> filter -> group by, $LoCardKeys keys", + filterGroupBySum, + extraConfigs = deterministicShuffle), + // ...and with an expensive partial aggregate, to see whether rescuing it outweighs paying for + // the rows the filter is about to throw away. + MapCase( + s"map -> filter -> group by, $HiCardKeys keys", + filterGroupBySum, + hiCard = true, + extraConfigs = deterministicShuffle), // Nothing above the map. The rewrite has no operator to rescue and pays a columnar-to-row // transition at the sink that the unfused plan does not, so this is the case the default-off // decision rests on. One output column takes the direct path: a single fused expression, no // struct wrapper. MapCase("map -> sink, 1 col", _.map(_.a + 1).toDF()), - // Two output columns take the `CreateNamedStruct` + `GetStructField` path, which is a second - // projection and a struct round-trip through Arrow. Worth separating from the row above: if - // the struct path is much worse at the top of the plan, that is a cost of the multi-column - // encoding rather than of fusing as such. + // Two output columns take the `CreateNamedStruct` + `GetStructField` path, which builds a + // struct vector and immediately reads the fields back out. Worth separating from the row + // above: if the struct path is much worse at the top of the plan, that is a cost of the + // multi-column encoding rather than of fusing as such. MapCase("map -> sink, 2 cols", _.map(r => TypedMapRec(r.a + 1, r.b)).toDF()), - // The case the rewrite exists for. With the fuse off, the fallback island takes the partial - // aggregate, the exchange and the final aggregate down with it. `b` has 100 distinct values, - // so the grouping is cheap and the per-row closure still dominates. Both output columns are - // consumed so that column pruning cannot give the two arms different work to do. - MapCase( - "map -> group by", - _.map(r => TypedMapRec(r.a + 1, r.b)).groupBy("b").agg(sum("a")), - extraConfigs = deterministicShuffle), - // A filter between the map and the aggregate, so the rewrite is rescuing three operator kinds - // rather than two and the fused projection feeds a native filter directly. - MapCase( - "map -> filter -> group by", - _.map(r => TypedMapRec(r.a * 2, r.b)) - .filter(col("a") % 3 === 0) - .groupBy("b") - .agg(sum("a")), - extraConfigs = deterministicShuffle), // `ds.map(f).map(g)` leaves two adjacent `MapElements` under one Serialize/Deserialize pair, // which the rule fuses as a whole. Two closure calls per row against one bridge crossing, so // the bridge is amortised further here than anywhere else in the table. @@ -221,7 +262,9 @@ object CometTypedDatasetMapBenchmark extends CometBenchmarkBase { emit( " at least 2 iterations and at least 2s of timed iterations; the table reports best and") emit(" average of the timed iterations.") - emit(s"Rows: $Rows. Grouping key `b` has 100 distinct values.") + emit( + s"Rows: $Rows. Grouping key has $LoCardKeys distinct values, or $HiCardKeys in the " + + "high-cardinality cases.") emit(s"Cases: ${selected.map(_.name).mkString(", ")}") } @@ -236,7 +279,7 @@ object CometTypedDatasetMapBenchmark extends CometBenchmarkBase { // `SQLHelper.withSQLConf` as returning `Unit`; only Spark 4 has the result-returning form. var collected: Array[Row] = Array.empty withSQLConf(configs: _*) { - collected = c.build(typedDataset).collect() + collected = c.build(typedDataset(c.hiCard)).collect() } collected.map(_.toSeq.map(String.valueOf).mkString("|")).sorted } @@ -277,13 +320,13 @@ object CometTypedDatasetMapBenchmark extends CometBenchmarkBase { var fusedNonComet: Option[String] = None var unfusedIsFullyComet = false withSQLConf(fusedConfigs(c): _*) { - val df = c.build(typedDataset) + val df = c.build(typedDataset(c.hiCard)) df.noop() fusedNonComet = findFirstNonCometOperator(stripAQEPlan(df.queryExecution.executedPlan)).map(_.nodeName) } withSQLConf(unfusedConfigs(c): _*) { - val df = c.build(typedDataset) + val df = c.build(typedDataset(c.hiCard)) df.noop() unfusedIsFullyComet = findFirstNonCometOperator(stripAQEPlan(df.queryExecution.executedPlan)).isEmpty @@ -310,26 +353,31 @@ object CometTypedDatasetMapBenchmark extends CometBenchmarkBase { private def runCase(c: MapCase, configs: Seq[(String, String)]): Unit = withSQLConf(configs: _*) { - c.build(typedDataset).noop() + c.build(typedDataset(c.hiCard)).noop() } /** * The corpus read back as a typed `Dataset`. Rebuilt per call rather than cached, because the * plan it produces has to be built under the arm's own confs. */ - private def typedDataset: Dataset[TypedMapRec] = - spark.sql("select c_a as a, c_b as b from parquetV1Table").as[TypedMapRec] + private def typedDataset(hiCard: Boolean): Dataset[TypedMapRec] = { + val key = if (hiCard) "c_hi" else "c_lo" + spark.sql(s"select c_a as a, $key as b from parquetV1Table").as[TypedMapRec] + } /** Builds `parquetV1Table` with `rows` rows of the corpus and drops it afterwards. */ private def withCorpus(rows: Int)(f: => Unit): Unit = { withTempPath { dir => withTempTable(tbl, "parquetV1Table") { spark.range(rows).createOrReplaceTempView(tbl) - // `c_a` varies per row so the closure's arithmetic is not loop-invariant; `c_b` has 100 - // distinct values so the grouped cases spend their time in the map, not the aggregate. + // `c_a` varies per row so the closure's arithmetic is not loop-invariant. `c_lo` makes the + // partial aggregate nearly free and `c_hi` makes it expensive, which is the variable the + // high-cardinality cases exist to move. prepareTable( dir, - spark.sql(s"SELECT id AS c_a, CAST(PMOD(id, 100) AS STRING) AS c_b FROM $tbl")) + spark.sql( + s"SELECT id AS c_a, CAST(PMOD(id, $LoCardKeys) AS STRING) AS c_lo, " + + s"CAST(PMOD(id, $HiCardKeys) AS STRING) AS c_hi FROM $tbl")) f } } From 9a8339e96ace759f7bca2022ff7ea951fb7c5d3a Mon Sep 17 00:00:00 2001 From: Andy Grove Date: Tue, 22 Sep 2026 08:01:24 -0600 Subject: [PATCH 6/6] bench: add width diagnostics, which refute the multi-output kernel Two cases at the top of the plan: a single string output, and a four-column record. The four-column row is the finding -- at two columns the fuse is 1.4-1.6x but at four it is 1.0x, so whatever the multi-column encoding costs scales with field count, and wide case classes are the ordinary shape. The obvious explanation was CreateNamedStruct.doGenCode: it allocates an Object[N] per row, boxes every primitive into it and wraps it in a GenericInternalRow, which the output writer then unpacks field by field. I implemented the fix -- a kernel that writes each field's ExprCode straight into its child vector, with one generateExpressions call across all N so CSE still collapses the shared closure Invoke -- and measured it over two runs. It was 15-20% slower on the 2-col and filtered cases and no better anywhere, so it is not in this commit. The row never escapes the loop body, so escape analysis was already scalar-replacing it; there were no allocations to remove, and inlining every field's code into the per-row body only grew the method. The reasoning is recorded next to the cases so the idea is not re-tried blind. Worth noting it needed no proto or native change: the native side already returns a StructArray and GetStructField on a non-nullable struct is an Arc::clone of the child, so the whole question was JVM-side. --- .../CometTypedDatasetMapBenchmark.scala | 24 +++++++++++++++++++ 1 file changed, 24 insertions(+) diff --git a/spark/src/test/scala/org/apache/spark/sql/benchmark/CometTypedDatasetMapBenchmark.scala b/spark/src/test/scala/org/apache/spark/sql/benchmark/CometTypedDatasetMapBenchmark.scala index 42aea540464..584003039c9 100644 --- a/spark/src/test/scala/org/apache/spark/sql/benchmark/CometTypedDatasetMapBenchmark.scala +++ b/spark/src/test/scala/org/apache/spark/sql/benchmark/CometTypedDatasetMapBenchmark.scala @@ -31,6 +31,9 @@ import org.apache.comet.CometConf /** Top-level so `NewInstance` needs no outer pointer, which is the ordinary user shape. */ case class TypedMapRec(a: Long, b: String) +/** Four output columns, to see whether struct-wrapper cost scales with field count. */ +case class TypedMapWide(a: Long, b: String, c: Long, d: String) + /** * Benchmark of `RewriteTypedDatasetMap`, which fuses the `SerializeFromObject` / `MapElements` / * `DeserializeToObject` sandwich a typed `Dataset.map` produces into a Comet projection routed @@ -97,6 +100,8 @@ object CometTypedDatasetMapBenchmark extends CometBenchmarkBase { // to start while this object is still initialising. private implicit val recEncoder: Encoder[TypedMapRec] = Encoders.product[TypedMapRec] private implicit val longEncoder: Encoder[Long] = Encoders.scalaLong + private implicit val stringEncoder: Encoder[String] = Encoders.STRING + private implicit val wideEncoder: Encoder[TypedMapWide] = Encoders.product[TypedMapWide] /** * @param name @@ -188,6 +193,25 @@ object CometTypedDatasetMapBenchmark extends CometBenchmarkBase { // above: if the struct path is much worse at the top of the plan, that is a cost of the // multi-column encoding rather than of fusing as such. MapCase("map -> sink, 2 cols", _.map(r => TypedMapRec(r.a + 1, r.b)).toDF()), + // Two diagnostics for the cost of the multi-column encoding, at the top of the plan so + // nothing else varies. `1 col (string)` isolates what a varchar output costs on its own -- + // `UTF8String.fromString` plus a variable-width Arrow write -- because the 2-col case pays + // that too and it must not be mistaken for struct overhead. `4 cols` says whether the + // per-column cost is linear or whether the struct wrapper degrades with field count; it is + // the row that shows the fuse buying nothing once the record is wide. + // + // These are what refuted the multi-output-kernel idea. The hypothesis was that + // `CreateNamedStruct.doGenCode`'s per-row `Object[N]` + boxed `GenericInternalRow` was the + // overhead, and that writing each field straight into its child vector would remove it. + // Implemented and measured, it was 15-20% *slower* on the 2-col and filtered cases: the row + // never escapes the loop body, so escape analysis was already scalar-replacing it, and + // inlining every field's code into the row body only made the method bigger. The native + // unpack was never a cost either -- `GetStructField` on a non-nullable struct is an + // `Arc::clone` of the child array. + MapCase("map -> sink, 1 col (string)", _.map(_.b).toDF()), + MapCase( + "map -> sink, 4 cols", + _.map(r => TypedMapWide(r.a + 1, r.b, r.a * 2, r.b + "!")).toDF()), // `ds.map(f).map(g)` leaves two adjacent `MapElements` under one Serialize/Deserialize pair, // which the rule fuses as a whole. Two closure calls per row against one bridge crossing, so // the bridge is amortised further here than anywhere else in the table.