diff --git a/.github/workflows/pr_build_linux.yml b/.github/workflows/pr_build_linux.yml index a9c184a20b0..18a33c8fddf 100644 --- a/.github/workflows/pr_build_linux.yml +++ b/.github/workflows/pr_build_linux.yml @@ -615,6 +615,7 @@ jobs: org.apache.comet.serde.CometScalarFunctionSuite org.apache.comet.serde.CometLiteralSuite 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 b47ed5a46f6..8193f43d4d5 100644 --- a/.github/workflows/pr_build_macos.yml +++ b/.github/workflows/pr_build_macos.yml @@ -263,6 +263,7 @@ jobs: org.apache.comet.serde.CometScalarFunctionSuite org.apache.comet.serde.CometLiteralSuite 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 2182e7de584..2b9d68c8b35 100644 --- a/docs/source/user-guide/latest/operators.md +++ b/docs/source/user-guide/latest/operators.md @@ -122,6 +122,23 @@ omitted from the tables below and may be reconsidered based on demand: | `WriteFilesExec` | ⚠️ | Spark 4.0+. Experimental native Parquet writes, disabled by default (opt-in). Non-partitioned, non-bucketed writes only, and not when `spark.sql.files.maxRecordsPerFile` is set. | | `DataWritingCommandExec` | ⚠️ | Spark 3.4/3.5 only. Experimental native Parquet writes, disabled by default (opt-in). Replaced by `WriteFilesExec` on Spark 4.0+ and removed with Spark 3.x support. | +## 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 d371f1ba44c..9cb54107e50 100644 --- a/spark/src/main/scala/org/apache/comet/CometConf.scala +++ b/spark/src/main/scala/org/apache/comet/CometConf.scala @@ -411,6 +411,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/codegen/CometBatchKernelCodegen.scala b/spark/src/main/scala/org/apache/comet/codegen/CometBatchKernelCodegen.scala index 83fbca6b635..94c872cd49e 100644 --- a/spark/src/main/scala/org/apache/comet/codegen/CometBatchKernelCodegen.scala +++ b/spark/src/main/scala/org/apache/comet/codegen/CometBatchKernelCodegen.scala @@ -134,9 +134,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 54f88916245..a8f32b85547 100644 --- a/spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala +++ b/spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala @@ -750,7 +750,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 => 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..32d7b2e1d1a --- /dev/null +++ b/spark/src/main/scala/org/apache/comet/rules/RewriteTypedDatasetMap.scala @@ -0,0 +1,246 @@ +/* + * 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.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} +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.serde.{CometScalaUDF, QueryPlanSerde} + +/** + * 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 { + + /** 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) + .collect { case deserialize: DeserializeToObjectExec => deserialize } + .flatMap(fuse(serialize, chain, _)) + .getOrElse(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 + + // 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: " + + nonAliases.map(_.getClass.getSimpleName).mkString(", ")) + } + + // 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 decline( + 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. 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 b: BoundReference if b.ordinal == 0 && b.dataType == objType => callFunc + }.asInstanceOf[NamedExpression] + } + + // 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 decline( + serialize, + s"fused expression references attributes outside the child: ${dangling.mkString(", ")}") + } + + // 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") + } + + fused match { + // One output column is already one kernel; the struct wrapper would have nothing to dedupe. + case Seq(only) => + 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)() + 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) + } + } + } + + /** + * 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 forceDispatch[T <: Expression]( + serialize: SerializeFromObjectExec, + expr: T): Option[T] = + CometScalaUDF.canDispatch(expr) match { + case Some(reason) => + decline(serialize, s"the fused expression is not dispatchable ($reason)") + None + 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 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 9249ab280f0..9088095d3c1 100644 --- a/spark/src/main/scala/org/apache/comet/serde/CometScalaUDF.scala +++ b/spark/src/main/scala/org/apache/comet/serde/CometScalaUDF.scala @@ -58,6 +58,47 @@ object CometScalaUDF extends CometExpressionSerde[ScalaUDF] { override def convert(expr: ScalaUDF, inputs: Seq[Attribute], binding: Boolean): Option[Expr] = emitJvmCodegenDispatch(expr, inputs, binding) + /** + * 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): (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 + (target, attrs, BindReferences.bindReference(target, AttributeSeq(attrs))) + } + + /** + * 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)._3) + } + /** * Bind `expr`, closure-serialize it, and emit a `JvmScalarUdf` proto routed through * [[CometScalaUDFCodegen]] so that native execution evaluates the expression inside the @@ -82,20 +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 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)) + // `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/main/scala/org/apache/comet/serde/QueryPlanSerde.scala b/spark/src/main/scala/org/apache/comet/serde/QueryPlanSerde.scala index 43d65fd418e..b183a64bca7 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.{Invoke, 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 @@ -1027,6 +1053,16 @@ 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 { 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 new file mode 100644 index 00000000000..b5c7f717a87 --- /dev/null +++ b/spark/src/test/scala/org/apache/comet/CometTypedDatasetSuite.scala @@ -0,0 +1,403 @@ +/* + * 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, DataFrame, Dataset, Encoders} +import org.apache.spark.sql.comet.CometProjectExec +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. */ +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) + +/** 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) + + 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 CometCodegenAssertions { + + 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) + + /** 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]) + } + + /** + * 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() { + withTypedRecs() { ds => + val (_, cometPlan) = checkSparkAnswerAndOperator(ds.map(_.a + 1).toDF()) + assertNoObjectOperators(cometPlan) + } + } + } + + test("ds.map produces a fully native plan - multiple output columns") { + withFusion() { + withTypedRecs() { ds => + // 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") + } + } + } + } + + 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) { + out = spark + .sql("select _1 as a, _2 as b from tbl") + .as[TypedRec] + .map(r => TypedRec(r.a + 1, r.b)) + .toDF() + .queryExecution + .executedPlan + .output + .map(a => s"${a.name}:${a.dataType.simpleString}:${a.nullable}") + .mkString(",") + } + out + } + assert(executedOutput(fused = true) === executedOutput(fused = false)) + } + } + + test("closure runs exactly once per row with multiple output columns") { + withFusion() { + withTypedRecs(50) { ds => + 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("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") { + 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))) + 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() { + 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() { + 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) + } + } + } + + 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() { + 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()) + // 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() { + withTypedRecs(40) { ds => + 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() { + 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) + } + } + } + + 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() + 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") { + withTypedRecs(20) { ds => + checkSparkAnswerAndFallbackReason( + ds.map(r => TypedRec(r.a + 1, r.b)).toDF(), + "Cannot fuse typed Dataset map: " + + s"${CometConf.COMET_SCALA_UDF_CODEGEN_ENABLED.key}=false, so there is no dispatcher " + + "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") { + 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() { + withTypedRecs(20) { ds => + 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") + } + } + } +} 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..584003039c9 --- /dev/null +++ b/spark/src/test/scala/org/apache/spark/sql/benchmark/CometTypedDatasetMapBenchmark.scala @@ -0,0 +1,423 @@ +/* + * 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) + +/** 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 + * 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, 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. + * + * '''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 + * 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 + private implicit val stringEncoder: Encoder[String] = Encoders.STRING + private implicit val wideEncoder: Encoder[TypedMapWide] = Encoders.product[TypedMapWide] + + /** + * @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 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. + * @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, + 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`. + */ + 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 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()), + // 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. + 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 has $LoCardKeys distinct values, or $HiCardKeys in the " + + "high-cardinality cases.") + 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(c.hiCard)).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(c.hiCard)) + df.noop() + fusedNonComet = + findFirstNonCometOperator(stripAQEPlan(df.queryExecution.executedPlan)).map(_.nodeName) + } + withSQLConf(unfusedConfigs(c): _*) { + val df = c.build(typedDataset(c.hiCard)) + 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(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(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_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, $LoCardKeys) AS STRING) AS c_lo, " + + s"CAST(PMOD(id, $HiCardKeys) AS STRING) AS c_hi 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))) + } +}