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 a709129bac9..5d3464c2d47 100644 --- a/spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala +++ b/spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala @@ -200,6 +200,16 @@ case class CometExecRule(session: SparkSession) private def isCometNative(op: SparkPlan): Boolean = op.isInstanceOf[CometNativeExec] + /** + * Restore a Spark Partial while retaining its current children. The tag prevents reconversion + * when AQE replans the exchange without its Final, and records why the Partial stays in Spark. + */ + private def restoreSparkPartial(agg: CometHashAggregateExec, reason: String): SparkPlan = { + val partial = agg.originalPlan.withNewChildren(agg.children) + partial.setTagValue(CometExecRule.COMET_UNSAFE_PARTIAL, reason) + withFallbackReason(partial, reason) + } + // spotless:off /** @@ -598,7 +608,7 @@ case class CometExecRule(session: SparkSession) // during the bottom-up conversion. Tags persist through AQE stage creation. tagUnsafePartialAggregates(planWithJoinRewritten) - var newPlan = transform(planWithJoinRewritten) + var newPlan = revertUnsafePartialAggregates(transform(planWithJoinRewritten)) // if the plan cannot be run fully natively then explain why (when appropriate // config is enabled) @@ -876,7 +886,7 @@ case class CometExecRule(session: SparkSession) // PartialMerge stages of a distinct-aggregate rewrite. See issues #1389 and #4813. val modes = agg.aggregateExpressions.map(_.mode).distinct if (modes == Seq(Final) && - !QueryPlanSerde.allAggsSupportMixedExecution(agg.aggregateExpressions) && + !QueryPlanSerde.allAggsSupportNativePartialToSparkFinal(agg.aggregateExpressions) && !canAggregateBeConverted(agg, Final)) { findPartialAggInPlan(agg.child).foreach { partial => // Only tag if the Partial would otherwise have been converted. If the Partial itself @@ -918,6 +928,103 @@ case class CometExecRule(session: SparkSession) } } + /** + * Inspect a failed repair's buffer path without rewriting it or materializing any stage. Report + * only a native Partial/PartialMerge whose emitted state is not known to be Spark-compatible. + * Spark Partials and completed aggregates establish new buffers, so stop there rather than + * finding an unrelated native producer below them. Only known aggregate and exchange wrappers + * forward the same buffer path; an arbitrary operator is not evidence of a mixed boundary. + */ + private def hasUnrepairedNativeBuffer(plan: SparkPlan): Boolean = plan match { + case agg: CometHashAggregateExec if agg.aggregateExpressions.isEmpty => + hasUnrepairedNativeBuffer(agg.child) + case agg: CometHashAggregateExec => + agg.modes.forall(m => m == Partial || m == PartialMerge) && + !QueryPlanSerde.allAggsSupportNativePartialToSparkFinal(agg.aggregateExpressions) + case agg: BaseAggregateExec + if agg.aggregateExpressions.nonEmpty && + agg.aggregateExpressions.forall(_.mode == Partial) => + false + case agg: BaseAggregateExec => + agg.aggregateExpressions.forall(e => e.mode == Partial || e.mode == PartialMerge) && + hasUnrepairedNativeBuffer(agg.child) + case placeholder: CometSinkPlaceHolder => hasUnrepairedNativeBuffer(placeholder.child) + case read: AQEShuffleReadExec => hasUnrepairedNativeBuffer(read.child) + case stage: ShuffleQueryStageExec => hasUnrepairedNativeBuffer(stage.plan) + case reused: ReusedExchangeExec => hasUnrepairedNativeBuffer(reused.child) + case shuffle: CometShuffleExchangeExec => hasUnrepairedNativeBuffer(shuffle.child) + case shuffle: ShuffleExchangeExec => hasUnrepairedNativeBuffer(shuffle.child) + case _ => false + } + + /** + * The early tagging pass cannot know whether a Final's child will become native. Check the + * actual conversion result before serialization or AQE stage creation, restoring the feeding + * aggregate/exchange chain while keeping native work below its Partial. Return the repaired + * plan, or preserve an unrepairable path and record one warning on its Spark Final if an unsafe + * native producer remains. Existing stages and their buffers are never rewritten by this pass. + */ + private[rules] def revertUnsafePartialAggregates(plan: SparkPlan): SparkPlan = { + def revertChain(node: SparkPlan): Option[SparkPlan] = node match { + case agg: CometHashAggregateExec if agg.modes == Seq(Partial) => + Some( + restoreSparkPartial( + agg, + "Partial aggregate disabled: corresponding final aggregate " + + "cannot be converted to Comet and intermediate buffer formats are incompatible")) + + case agg: CometHashAggregateExec + if agg.modes.forall(m => m == Partial || m == PartialMerge) => + revertChain(agg.child).map(child => agg.originalPlan.withNewChildren(Seq(child))) + + case agg: BaseAggregateExec + if agg.aggregateExpressions.nonEmpty && + agg.aggregateExpressions.forall(_.mode == Partial) => + // This producer already emits Spark buffers. Do not reach through it to an unrelated + // aggregate below it. + None + + case agg: BaseAggregateExec + if agg.aggregateExpressions.forall(e => e.mode == Partial || e.mode == PartialMerge) => + revertChain(agg.child).map(child => agg.withNewChildren(Seq(child))) + + case CometSinkPlaceHolder(_, _, shuffle: CometShuffleExchangeExec) => + revertChain(shuffle) + case shuffle: CometShuffleExchangeExec => + revertChain(shuffle.child).map(child => shuffle.originalPlan.withNewChildren(Seq(child))) + case shuffle: ShuffleExchangeExec => + revertChain(shuffle.child).map(child => shuffle.withNewChildren(Seq(child))) + + // Stop at materialized stages and operators outside the feeding aggregate/exchange chain. + case _ => None + } + + plan.transformUp { + case agg: BaseAggregateExec + if agg.aggregateExpressions.map(_.mode).distinct == Seq(Final) && + !QueryPlanSerde.allAggsSupportNativePartialToSparkFinal(agg.aggregateExpressions) => + revertChain(agg.child) + // Rebuild native consumers and shuffles from their original Spark operators. Merely + // replacing their children would leave a native protobuf reading the old buffers. + .map(child => transform(agg.withNewChildren(Seq(child)))) + .getOrElse { + if (hasUnrepairedNativeBuffer(agg.child)) { + val reason = "Comet could not restore a native intermediate buffer producer " + + "below Spark final aggregate; the remaining buffer may be incompatible" + // AQE can revisit the same consumer. Record the explanation and warn once, + // regardless of whether general fallback logging is enabled. + if (!agg + .getTagValue(CometExplainInfo.FALLBACK_REASONS) + .exists(_.contains(reason))) { + if (!CometConf.COMET_EXPLAIN_FALLBACK_LOG_ENABLED.get()) logWarning(reason) + withFallbackReason(agg, reason) + } + } + agg + } + } + } + /** * Look for the bottom Partial-mode aggregate that feeds into the given plan (the child of a * Final). Walks through exchanges and AQE stages, and continues down through intermediate @@ -947,8 +1054,8 @@ case class CometExecRule(session: SparkSession) /** * Conservative check for whether an aggregate could be converted to Comet. Checks operator * enablement, grouping expressions, aggregate expressions, and result expressions. - * Intentionally skips the sparkFinalMode / child-native checks since those depend on - * transformation state. + * Intentionally skips the child-native checks since those depend on transformation state; + * [[revertUnsafePartialAggregates]] checks the actual conversion result before execution. * * WARNING: this intentionally mirrors the predicate checks in `CometBaseAggregate.doConvert` * (operators.scala). Any change to the convertibility rules there must be reflected here or diff --git a/spark/src/main/scala/org/apache/comet/rules/RevertNativeForTransitionHeavyStages.scala b/spark/src/main/scala/org/apache/comet/rules/RevertNativeForTransitionHeavyStages.scala index 1e0cfc79e0b..a56e477b035 100644 --- a/spark/src/main/scala/org/apache/comet/rules/RevertNativeForTransitionHeavyStages.scala +++ b/spark/src/main/scala/org/apache/comet/rules/RevertNativeForTransitionHeavyStages.scala @@ -21,14 +21,16 @@ package org.apache.comet.rules import org.apache.spark.internal.Logging import org.apache.spark.sql.SparkSession +import org.apache.spark.sql.catalyst.expressions.aggregate.{Final, Partial, PartialMerge} import org.apache.spark.sql.catalyst.rules.Rule -import org.apache.spark.sql.comet.{CometColumnarToRowExec, CometExec, CometNativeColumnarToRowExec, CometSparkToColumnarExec} +import org.apache.spark.sql.comet.{CometColumnarToRowExec, CometExec, CometHashAggregateExec, CometNativeColumnarToRowExec, CometSparkToColumnarExec} import org.apache.spark.sql.execution.{ColumnarToRowExec, ColumnarToRowTransition, RowToColumnarExec, SparkPlan} import org.apache.spark.sql.execution.adaptive.QueryStageExec import org.apache.spark.sql.execution.exchange.{BroadcastExchangeLike, ShuffleExchangeLike} import org.apache.comet.CometConf import org.apache.comet.CometSparkSessionExtensions.withFallbackReason +import org.apache.comet.serde.QueryPlanSerde /** * Reverts a query stage to Spark row-based execution when it has too many columnar-to-row (C2R) @@ -85,6 +87,13 @@ case class RevertNativeForTransitionHeavyStages(session: SparkSession) val transitionCount = countTransitions(stagePlan) if (transitionCount <= maxTransitions) return None + // Reverting either side of a native aggregate boundary can make one engine consume the + // other's intermediate state. Typed imperative aggregates such as percentile expose a native + // array where Spark expects serialized binary; others, including COUNT, must remain in one + // engine for planner semantics even though their physical buffer types match. Keep both + // producer and consumer stages native when mixed execution is unsafe across a stage boundary. + if (hasUnsafeMixedAggregateAtStageBoundary(stagePlan)) return None + val reason = s"Stage reverted: $transitionCount C2R transitions exceed threshold $maxTransitions" @@ -105,6 +114,33 @@ case class RevertNativeForTransitionHeavyStages(session: SparkSession) case _ => false } + private def hasUnsafeMixedAggregateAtStageBoundary(stagePlan: SparkPlan): Boolean = { + def reachesBoundaryBeforeAggregate(plan: SparkPlan): Boolean = plan match { + case _ if isStageBoundary(plan) => true + case _: CometHashAggregateExec => false + case _ => plan.children.exists(reachesBoundaryBeforeAggregate) + } + + def visit(plan: SparkPlan): Boolean = plan match { + case _ if isStageBoundary(plan) => false + case aggregate: CometHashAggregateExec + if !QueryPlanSerde + .allAggsSupportNativePartialToSparkFinal(aggregate.aggregateExpressions) || + QueryPlanSerde + .aggsNotSupportingSparkPartialToNativeFinal(aggregate.aggregateExpressions) + .nonEmpty => + val producesBuffer = + aggregate.modes.exists(mode => mode == Partial || mode == PartialMerge) + val consumesAcrossBoundary = + aggregate.modes.exists(mode => mode == Final || mode == PartialMerge) && + reachesBoundaryBeforeAggregate(aggregate.child) + producesBuffer || consumesAcrossBoundary || aggregate.children.exists(visit) + case _ => plan.children.exists(visit) + } + + visit(stagePlan) + } + /** * Like `transformDown`, never descends stage-boundary children. */ diff --git a/spark/src/main/scala/org/apache/comet/serde/CometAggregateExpressionSerde.scala b/spark/src/main/scala/org/apache/comet/serde/CometAggregateExpressionSerde.scala index a52d6008211..73630a4e8b8 100644 --- a/spark/src/main/scala/org/apache/comet/serde/CometAggregateExpressionSerde.scala +++ b/spark/src/main/scala/org/apache/comet/serde/CometAggregateExpressionSerde.scala @@ -82,16 +82,21 @@ trait CometAggregateExpressionSerde[T <: AggregateFunction] { def getSupportLevel(expr: T): SupportLevel = Compatible(None) /** - * Whether this aggregate's intermediate buffer format is compatible between Spark and Comet for - * the given function instance, making it safe to run the Partial in one engine and the Final in - * the other. Aggregates with simple single-value buffers (MIN, MAX, bitwise) are always safe; - * SUM and non-decimal AVG match Spark's buffer and are safe except where noted per instance - * (e.g. TRY-mode SUM uses a Comet-internal flag column). COUNT is intentionally excluded - * despite a matching buffer: mixed COUNT partial/final regressed AQE's + * Whether a Comet aggregate can consume this function's Spark intermediate buffer. This covers + * Spark Partial to Comet Final, including intermediate PartialMerge stages. COUNT is excluded + * despite a matching buffer: a Comet Final above a Spark Partial regressed AQE's * PropagateEmptyRelationAfterAQE pattern (which matches BaseAggregateExec only) and the Spark * 4.0 count-bug decorrelation for correlated IN subqueries. */ - def supportsMixedPartialFinal(fn: T): Boolean = false + def supportsSparkPartialToNativeFinal(fn: T): Boolean = false + + /** + * Whether Spark can consume this function's Comet intermediate buffer. Opt in independently + * from the reverse direction: consuming Spark state does not establish that Comet emits state + * Spark can merge, especially from a never-updated or all-null partial accumulator. Remaining + * forward-compatibility audits are tracked in issue #5975. + */ + def supportsNativePartialToSparkFinal(fn: T): Boolean = false /** * Convert a Spark expression into a protocol buffer representation that can be passed into 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 9e867aad2d2..390eef992e1 100644 --- a/spark/src/main/scala/org/apache/comet/serde/QueryPlanSerde.scala +++ b/spark/src/main/scala/org/apache/comet/serde/QueryPlanSerde.scala @@ -417,31 +417,35 @@ object QueryPlanSerde extends Logging with CometExprShim with CometTypeShim { classOf[VarianceSamp] -> CometVarianceSamp) /** - * Returns true if all aggregate expressions in the list have intermediate buffer formats that - * are compatible between Spark and Comet, making it safe to run Partial in one engine and Final - * in the other. + * Returns true if Spark can consume all the intermediate buffers produced by Comet. Used when a + * Spark Final would otherwise consume a native Partial, including after shuffle fallback. */ - def allAggsSupportMixedExecution(aggExprs: Seq[AggregateExpression]): Boolean = { - aggExprs.forall(aggExpr => supportsMixedExecution(aggExpr.aggregateFunction)) + def allAggsSupportNativePartialToSparkFinal(aggExprs: Seq[AggregateExpression]): Boolean = { + aggExprs.forall { aggExpr => + val fn = aggExpr.aggregateFunction + aggrSerdeMap.get(fn.getClass).exists { handler => + handler + .asInstanceOf[CometAggregateExpressionSerde[AggregateFunction]] + .supportsNativePartialToSparkFinal(fn) + } + } } /** - * Returns the aggregate functions in the list whose intermediate buffer formats are not known - * to be compatible between Spark and Comet. These are the functions that prevent a Spark Final - * aggregate (without a Comet Partial) from running, since the buffer produced by one engine - * cannot be safely consumed by the other. + * Returns functions whose Spark intermediate buffers cannot safely be consumed by a Comet Final + * or PartialMerge. This is independent of native Partial to Spark Final compatibility. */ - def aggsNotSupportingMixedExecution( + def aggsNotSupportingSparkPartialToNativeFinal( aggExprs: Seq[AggregateExpression]): Seq[AggregateFunction] = { - aggExprs.map(_.aggregateFunction).filterNot(supportsMixedExecution) + aggExprs.map(_.aggregateFunction).filterNot(supportsSparkPartialToNativeFinal) } - private def supportsMixedExecution(fn: AggregateFunction): Boolean = { + private def supportsSparkPartialToNativeFinal(fn: AggregateFunction): Boolean = { aggrSerdeMap.get(fn.getClass) match { case Some(handler) => handler .asInstanceOf[CometAggregateExpressionSerde[AggregateFunction]] - .supportsMixedPartialFinal(fn) + .supportsSparkPartialToNativeFinal(fn) case None => false } } diff --git a/spark/src/main/scala/org/apache/comet/serde/aggregates.scala b/spark/src/main/scala/org/apache/comet/serde/aggregates.scala index 316744be1f6..9946421d5ac 100644 --- a/spark/src/main/scala/org/apache/comet/serde/aggregates.scala +++ b/spark/src/main/scala/org/apache/comet/serde/aggregates.scala @@ -35,7 +35,10 @@ import org.apache.comet.shims.{CometCollectShim, CometEvalModeUtil} object CometMin extends CometAggregateExpressionSerde[Min] { - override def supportsMixedPartialFinal(fn: Min): Boolean = true + override def supportsSparkPartialToNativeFinal(fn: Min): Boolean = true + + // Native MIN emits one typed null for empty/all-null input; Spark's least merge ignores it. + override def supportsNativePartialToSparkFinal(fn: Min): Boolean = true override def getSupportLevel(expr: Min): SupportLevel = AggSerde.minMaxSupportLevel(expr.dataType) @@ -72,7 +75,10 @@ object CometMin extends CometAggregateExpressionSerde[Min] { object CometMax extends CometAggregateExpressionSerde[Max] { - override def supportsMixedPartialFinal(fn: Max): Boolean = true + override def supportsSparkPartialToNativeFinal(fn: Max): Boolean = true + + // Native MAX emits one typed null for empty/all-null input; Spark's greatest merge ignores it. + override def supportsNativePartialToSparkFinal(fn: Max): Boolean = true override def getSupportLevel(expr: Max): SupportLevel = AggSerde.minMaxSupportLevel(expr.dataType) @@ -108,6 +114,10 @@ object CometMax extends CometAggregateExpressionSerde[Max] { } object CometCount extends CometAggregateExpressionSerde[Count] { + // Both buffers are a single non-null Long. The AQE/count-bug restrictions documented on the + // reverse direction concern a Comet Final; retaining Spark's Final preserves those rewrites. + override def supportsNativePartialToSparkFinal(fn: Count): Boolean = true + override def convert( aggExpr: AggregateExpression, expr: Count, @@ -132,7 +142,10 @@ object CometCount extends CometAggregateExpressionSerde[Count] { object CometAverage extends CometAggregateExpressionSerde[Average] { - override def supportsMixedPartialFinal(fn: Average): Boolean = + // Keep the default native-to-Spark restriction until #5420: an untouched native AVG emits + // (null, 0), but Spark's merge needs (0.0, 0). + + override def supportsSparkPartialToNativeFinal(fn: Average): Boolean = // Non-decimal AVG has a (sum: double, count: long) buffer matching Spark. Decimal AVG is // deferred (overflow nulls count differently) and stays unsafe for mixed execution. !fn.child.dataType.isInstanceOf[DecimalType] @@ -193,7 +206,17 @@ object CometAverage extends CometAggregateExpressionSerde[Average] { object CometSum extends CometAggregateExpressionSerde[Sum] { - override def supportsMixedPartialFinal(fn: Sum): Boolean = + // Non-decimal, non-TRY SUM emits one nullable sum, including null for empty/all-null input; + // Spark's coalesce-based merge accepts it. Decimal SUM has Spark's (sum, isEmpty) layout, + // but native updates make precision overflow sticky (or throw in ANSI mode). Spark's generated + // scalar SUM can recover before emitting its partial: decimal(38,38) inputs 0.6, 0.6, -0.6 + // sum to 0.6. Keep decimal partials in Spark until those update semantics match. Integer TRY + // SUM also remains excluded because its native state contains an extra has_all_nulls column. + override def supportsNativePartialToSparkFinal(fn: Sum): Boolean = + !fn.child.dataType.isInstanceOf[DecimalType] && + CometEvalModeUtil.fromSparkEvalMode(CometEvalModeUtil.sumEvalMode(fn)) != CometEvalMode.TRY + + override def supportsSparkPartialToNativeFinal(fn: Sum): Boolean = // Decimal SUM is excluded: overflow detection (ANSI throw / Legacy null) does not survive a // Spark-partial / Comet-final split, so the required ArithmeticException is never raised. // TRY-mode integer SUM carries a Comet-internal has_all_nulls column that Spark cannot read. @@ -314,7 +337,10 @@ object CometLast extends CometAggregateExpressionSerde[Last] { } object CometBitAndAgg extends CometAggregateExpressionSerde[BitAndAgg] { - override def supportsMixedPartialFinal(fn: BitAndAgg): Boolean = true + override def supportsSparkPartialToNativeFinal(fn: BitAndAgg): Boolean = true + + // The single native buffer is null for empty/all-null input; Spark's merge skips nulls. + override def supportsNativePartialToSparkFinal(fn: BitAndAgg): Boolean = true override def getSupportLevel(expr: BitAndAgg): SupportLevel = if (AggSerde.bitwiseAggTypeSupported(expr.dataType)) { @@ -353,7 +379,10 @@ object CometBitAndAgg extends CometAggregateExpressionSerde[BitAndAgg] { } object CometBitOrAgg extends CometAggregateExpressionSerde[BitOrAgg] { - override def supportsMixedPartialFinal(fn: BitOrAgg): Boolean = true + override def supportsSparkPartialToNativeFinal(fn: BitOrAgg): Boolean = true + + // The single native buffer is null for empty/all-null input; Spark's merge skips nulls. + override def supportsNativePartialToSparkFinal(fn: BitOrAgg): Boolean = true override def getSupportLevel(expr: BitOrAgg): SupportLevel = if (AggSerde.bitwiseAggTypeSupported(expr.dataType)) { @@ -392,7 +421,10 @@ object CometBitOrAgg extends CometAggregateExpressionSerde[BitOrAgg] { } object CometBitXOrAgg extends CometAggregateExpressionSerde[BitXorAgg] { - override def supportsMixedPartialFinal(fn: BitXorAgg): Boolean = true + override def supportsSparkPartialToNativeFinal(fn: BitXorAgg): Boolean = true + + // The single native buffer is null for empty/all-null input; Spark's merge skips nulls. + override def supportsNativePartialToSparkFinal(fn: BitXorAgg): Boolean = true override def getSupportLevel(expr: BitXorAgg): SupportLevel = if (AggSerde.bitwiseAggTypeSupported(expr.dataType)) { @@ -781,7 +813,11 @@ object CometCorr extends CometAggregateExpressionSerde[Corr] { object CometBloomFilterAggregate extends CometAggregateExpressionSerde[BloomFilterAggregate] { - override def supportsMixedPartialFinal(fn: BloomFilterAggregate): Boolean = true + override def supportsSparkPartialToNativeFinal(fn: BloomFilterAggregate): Boolean = true + + // Native state is Spark's serialized filter, non-null even for empty/all-null input; only + // the final result may be null, so Spark's deserialize always receives a valid filter. + override def supportsNativePartialToSparkFinal(fn: BloomFilterAggregate): Boolean = true override def getSupportLevel(expr: BloomFilterAggregate): SupportLevel = expr.child.dataType match { @@ -957,7 +993,10 @@ object CometApproxCountDistinct extends CometAggregateExpressionSerde[HyperLogLo // The register buffer uses Spark's identical packed-`Long` layout (`numWords` `Long` columns), // matching Spark's `aggBufferSchema`, so a Comet partial and Spark final (or the reverse) can // be mixed in one plan. - override def supportsMixedPartialFinal(fn: HyperLogLogPlusPlus): Boolean = true + override def supportsSparkPartialToNativeFinal(fn: HyperLogLogPlusPlus): Boolean = true + + // Native empty/all-null state contains non-null zero Long words, matching Spark's registers. + override def supportsNativePartialToSparkFinal(fn: HyperLogLogPlusPlus): Boolean = true // Types that Comet's native `xxhash64` hashes identically to Spark's `XxHash64Function`. // `StringType` here is the default UTF8_BINARY collation; a collated `StringType(collationId)` diff --git a/spark/src/main/scala/org/apache/spark/sql/comet/operators.scala b/spark/src/main/scala/org/apache/spark/sql/comet/operators.scala index 57db6100554..39eea79045a 100644 --- a/spark/src/main/scala/org/apache/spark/sql/comet/operators.scala +++ b/spark/src/main/scala/org/apache/spark/sql/comet/operators.scala @@ -1679,7 +1679,7 @@ trait CometBaseAggregate { if (missingCometProducer) { val incompatibleAggs = - QueryPlanSerde.aggsNotSupportingMixedExecution(aggregate.aggregateExpressions) + QueryPlanSerde.aggsNotSupportingSparkPartialToNativeFinal(aggregate.aggregateExpressions) if (incompatibleAggs.nonEmpty) { val names = incompatibleAggs.map(_.prettyName).distinct.sorted.mkString(", ") withFallbackReason( diff --git a/spark/src/test/scala/org/apache/comet/exec/CometAggregateSuite.scala b/spark/src/test/scala/org/apache/comet/exec/CometAggregateSuite.scala index dd44000c940..54f49651880 100644 --- a/spark/src/test/scala/org/apache/comet/exec/CometAggregateSuite.scala +++ b/spark/src/test/scala/org/apache/comet/exec/CometAggregateSuite.scala @@ -25,18 +25,22 @@ import org.apache.hadoop.fs.Path import org.apache.spark.SparkConf import org.apache.spark.sql.{CometTestBase, DataFrame, Row} import org.apache.spark.sql.catalyst.expressions.Cast -import org.apache.spark.sql.catalyst.expressions.aggregate.Final +import org.apache.spark.sql.catalyst.expressions.aggregate.{Final, Partial} import org.apache.spark.sql.catalyst.optimizer.EliminateSorts -import org.apache.spark.sql.comet.CometHashAggregateExec -import org.apache.spark.sql.execution.adaptive.{AdaptiveSparkPlanExec, AdaptiveSparkPlanHelper} -import org.apache.spark.sql.execution.exchange.ReusedExchangeExec +import org.apache.spark.sql.catalyst.plans.physical.HashPartitioning +import org.apache.spark.sql.comet.{CometFilterExec, CometHashAggregateExec, CometNativeExec, CometProjectExec} +import org.apache.spark.sql.comet.execution.shuffle.CometShuffleExchangeExec +import org.apache.spark.sql.execution.adaptive.{AdaptiveSparkPlanExec, AdaptiveSparkPlanHelper, ShuffleQueryStageExec} +import org.apache.spark.sql.execution.aggregate.BaseAggregateExec +import org.apache.spark.sql.execution.exchange.{ReusedExchangeExec, ShuffleExchangeExec} import org.apache.spark.sql.functions.{avg, col, count_distinct, sum} import org.apache.spark.sql.internal.SQLConf -import org.apache.spark.sql.types.{DataTypes, StructField, StructType} +import org.apache.spark.sql.types.{ArrayType, DataTypes, StructField, StructType} import org.apache.comet.CometConf import org.apache.comet.CometConf.COMET_EXEC_STRICT_FLOATING_POINT import org.apache.comet.CometSparkSessionExtensions.isSpark41Plus +import org.apache.comet.rules.CometExecRule import org.apache.comet.testing.{DataGenOptions, FuzzDataGenerator, ParquetGenerator, SchemaGenOptions} /** @@ -192,6 +196,290 @@ class CometAggregateSuite extends CometTestBase with AdaptiveSparkPlanHelper { } } + for (adaptive <- Seq(false, true)) { + test(s"decimal AVG falls back across a Spark shuffle (AQE=$adaptive)") { + withTempDir { dir => + val path = s"${dir.getAbsolutePath}/data" + withSQLConf(CometConf.COMET_ENABLED.key -> "false") { + spark + .range(0L, 8L, 1L, 4) + .selectExpr("id", "CAST(200 AS DECIMAL(20, 2)) AS v") + .write + .parquet(path) + } + + withSQLConf( + SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> adaptive.toString, + SQLConf.FILES_MAX_PARTITION_BYTES.key -> "1048576", + SQLConf.PARQUET_FILTER_PUSHDOWN_ENABLED.key -> "false", + CometConf.COMET_NATIVE_SCAN_ENABLED.key -> "false", + CometConf.COMET_CONVERT_FROM_PARQUET_ENABLED.key -> "true", + CometConf.COMET_SHUFFLE_ENABLED.key -> "false") { + withParquetTable(path, "decimal_avg_fallback") { + // The filter leaves three input partitions empty. Decimal AVG is not safe to mix + // between engines: a native empty partial can poison the Spark final's sum buffer. + val df = sql("SELECT AVG(v) FROM decimal_avg_fallback WHERE id = 1") + val initialPlan = stripAQEPlan(df.queryExecution.executedPlan) + checkAnswer(df, Seq(Row(new java.math.BigDecimal("200.000000")))) + for (plan <- Seq(initialPlan, df.queryExecution.executedPlan)) { + assert(collect(plan) { case agg: CometHashAggregateExec => agg }.isEmpty) + val partials = collect(plan) { + case agg: BaseAggregateExec + if agg.aggregateExpressions.forall(_.mode == Partial) => + agg + } + assert(partials.size == 1) + assert(partials.forall(_.getTagValue(CometExecRule.COMET_UNSAFE_PARTIAL).isDefined)) + // Falling back the aggregate must not discard the native filter/scan conversion. + assert(collect(plan) { case filter: CometFilterExec => filter }.nonEmpty) + } + if (adaptive) { + val stages = collect(df.queryExecution.executedPlan) { + case stage: ShuffleQueryStageExec => stage + } + assert(stages.nonEmpty && stages.forall(_.isMaterialized)) + } + + // Compatible buffers may still use a native Partial and a Spark Final. + val safe = sql("SELECT MIN(v), MAX(v) FROM decimal_avg_fallback WHERE id = 1") + checkAnswer( + safe, + Seq(Row(new java.math.BigDecimal("200.00"), new java.math.BigDecimal("200.00")))) + assert(collect(safe.queryExecution.executedPlan) { case agg: CometHashAggregateExec => + agg + }.size == 1) + + // The same unsafe buffer is valid when both aggregate stages execute in Comet. + withSQLConf(CometConf.COMET_SHUFFLE_ENABLED.key -> "true") { + val native = sql("SELECT AVG(v) FROM decimal_avg_fallback WHERE id = 1") + checkAnswer(native, Seq(Row(new java.math.BigDecimal("200.000000")))) + assert(collect(native.queryExecution.executedPlan) { + case agg: CometHashAggregateExec => agg + }.size == 2) + } + } + } + } + } + } + + for (adaptive <- Seq(false, true)) { + test(s"COUNT and AVG fall back together across a Spark shuffle (AQE=$adaptive)") { + withTempDir { dir => + val path = s"${dir.getAbsolutePath}/data" + withSQLConf(CometConf.COMET_ENABLED.key -> "false") { + spark + .range(0L, 8L, 1L, 4) + .selectExpr("id", "CAST(1 AS BIGINT) AS v") + .write + .parquet(path) + } + + withSQLConf( + SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> adaptive.toString, + SQLConf.FILES_MAX_PARTITION_BYTES.key -> "1048576", + SQLConf.PARQUET_FILTER_PUSHDOWN_ENABLED.key -> "false", + CometConf.COMET_NATIVE_SCAN_ENABLED.key -> "false", + CometConf.COMET_CONVERT_FROM_PARQUET_ENABLED.key -> "true", + CometConf.COMET_SHUFFLE_ENABLED.key -> "false", + CometConf.COMET_ENABLE_FINAL_HASH_AGGREGATE.key -> "true") { + withParquetTable(path, "count_avg_fallback") { + assert(sql("SELECT * FROM count_avg_fallback").rdd.getNumPartitions == 4) + // A safe COUNT buffer must not admit an unsafe AVG buffer in the same Partial. + // Three partitions have no surviving rows, so AVG has no update_batch call and + // its native state is (null, 0), which poisons Spark Final's sum. Keeping Final + // enabled exercises repair after the shuffle falls back during conversion. + val df = sql("SELECT COUNT(*), AVG(v) FROM count_avg_fallback WHERE id = 1") + val initialPlan = stripAQEPlan(df.queryExecution.executedPlan) + checkAnswer(df, Seq(Row(1L, 1.0))) + for (plan <- Seq(initialPlan, df.queryExecution.executedPlan)) { + assert(collect(plan) { case agg: CometHashAggregateExec => agg }.isEmpty) + val partials = collect(plan) { + case agg: BaseAggregateExec + if agg.aggregateExpressions.map(_.mode).distinct == Seq(Partial) => + agg + } + assert(partials.size == 1) + assert(partials.head.getTagValue(CometExecRule.COMET_UNSAFE_PARTIAL).isDefined) + assert(collect(plan) { case filter: CometFilterExec => filter }.nonEmpty) + } + withSQLConf( + CometConf.COMET_SHUFFLE_ENABLED.key -> "true", + CometConf.COMET_SHUFFLE_MODE.key -> "native") { + // The native Final can consume its own empty AVG buffers; only the engine split + // is unsafe. Keep the fully native aggregate path enabled. + val native = + sql("SELECT COUNT(*), AVG(v) FROM count_avg_fallback WHERE id = 1") + val initialNativePlan = stripAQEPlan(native.queryExecution.executedPlan) + checkAnswer(native, Seq(Row(1L, 1.0))) + for (plan <- Seq(initialNativePlan, native.queryExecution.executedPlan)) { + assert(collect(plan) { case agg: CometHashAggregateExec => agg }.size == 2) + } + } + } + } + } + } + + test(s"COUNT preserves safe native partials across a Spark shuffle (AQE=$adaptive)") { + val data = Seq((0, None), (0, None), (1, Some(3)), (1, None), (1, Some(4))) + withParquetTable(data, "count_fallback", false) { + withSQLConf( + SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> adaptive.toString, + SQLConf.PARQUET_FILTER_PUSHDOWN_ENABLED.key -> "false", + CometConf.COMET_SHUFFLE_ENABLED.key -> "false", + CometConf.COMET_ENABLE_FINAL_HASH_AGGREGATE.key -> "true") { + for (query <- Seq( + "SELECT _1, COUNT(_2), COUNT(*) FROM count_fallback GROUP BY _1", + "SELECT COUNT(_2), COUNT(*) FROM count_fallback WHERE _1 < 0")) { + val df = sql(query) + val initialPlan = stripAQEPlan(df.queryExecution.executedPlan) + assert(collect(initialPlan) { + case agg: CometHashAggregateExec if agg.modes == Seq(Partial) => agg + }.size == 1) + assert(collect(initialPlan) { + case agg: BaseAggregateExec + if agg.aggregateExpressions.map(_.mode).distinct == Seq(Final) => + agg + }.size == 1) + checkSparkAnswer(df) + } + } + } + } + + for (fn <- Seq("collect_list", "collect_set")) { + test(s"$fn falls back when enabled native shuffle is ineligible (AQE=$adaptive)") { + val data = (0 until 30).map(i => (i % 3, if (i % 7 == 0) None else Some(i % 5))) + withParquetTable(data, "collect_fallback", false) { + withSQLConf( + SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> adaptive.toString, + SQLConf.PARQUET_FILTER_PUSHDOWN_ENABLED.key -> "false", + SQLConf.USE_OBJECT_HASH_AGG.key -> "true", + CometConf.COMET_SHUFFLE_ENABLED.key -> "true", + CometConf.COMET_SHUFFLE_MODE.key -> "native", + CometConf.COMET_SHUFFLE_NATIVE_HASH_PARTITIONING_ENABLED.key -> "false") { + // Integer keys isolate this from the wide-decimal shuffle restriction in #5420. + // The native Partial emits an Array buffer, but Spark's Final expects Binary. + val query = s"SELECT _1, sort_array($fn(_2)), COUNT(*) " + + "FROM collect_fallback WHERE _1 >= 0 GROUP BY _1" + val df = sql(query) + val initialPlan = stripAQEPlan(df.queryExecution.executedPlan) + checkSparkAnswer(df) + for (plan <- Seq(initialPlan, df.queryExecution.executedPlan)) { + assert(collect(plan) { case agg: CometHashAggregateExec => agg }.isEmpty) + val partials = collect(plan) { + case agg: BaseAggregateExec + if agg.aggregateExpressions.map(_.mode).distinct == Seq(Partial) => + agg + } + assert(partials.size == 1) + assert(partials.head.getTagValue(CometExecRule.COMET_UNSAFE_PARTIAL).isDefined) + assert(collect(plan) { case filter: CometFilterExec => filter }.nonEmpty) + } + // A fully native producer/consumer pair can still use its native buffer format. + withSQLConf(CometConf.COMET_SHUFFLE_NATIVE_HASH_PARTITIONING_ENABLED.key -> "true") { + val native = sql(query) + checkSparkAnswer(native) + assert(getNumCometHashAggregate(native) == 2) + } + } + } + } + } + + for (fn <- Seq("percentile", "collect_list", "sum")) { + test( + s"$fn preserves aggregate buffers with an unsupported array hash key (AQE=$adaptive)") { + withSQLConf( + SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> adaptive.toString, + SQLConf.SHUFFLE_PARTITIONS.key -> "4", + CometConf.COMET_SHUFFLE_ENABLED.key -> "true", + CometConf.COMET_SHUFFLE_MODE.key -> "native", + CometConf.COMET_SHUFFLE_NATIVE_HASH_PARTITIONING_ENABLED.key -> "true") { + withTempView("array_key_aggregate") { + // The array key itself makes native shuffle ineligible; no feature is disabled. + // https://github.com/apache/datafusion-comet/issues/5419#issuecomment-5464233245 + spark + .range(0, 18, 1, 4) + .selectExpr("id % 3 AS k", "id % 5 AS v") + .createOrReplaceTempView("array_key_aggregate") + val aggregate = if (fn == "percentile") "percentile(v, 0.5)" else s"$fn(v)" + val query = s"SELECT array(k) AS ak, $aggregate " + + "FROM array_key_aggregate GROUP BY array(k)" + + def normalizedRows(df: DataFrame): Seq[Row] = { + df.collect() + .toSeq + .map { row => + // Keep the reported collect_list SQL unchanged, normalizing its order only + // after execution so another expression cannot cause an earlier fallback. + if (fn == "collect_list") { + Row(row.getSeq[Long](0), row.getSeq[Long](1).sorted) + } else { + row + } + } + .sortBy(_.getSeq[Long](0).head) + } + + // Spark 3's withSQLConf returns Unit, so capture the baseline inside its body. + var expected: Seq[Row] = Seq.empty + withSQLConf(CometConf.COMET_ENABLED.key -> "false") { + expected = normalizedRows(sql(query)) + } + val df = sql(query) + val initialPlan = stripAQEPlan(df.queryExecution.executedPlan) + // Execute this same DataFrame before inspecting its materialized AQE plan. + assert(normalizedRows(df) == expected) + for (plan <- Seq(initialPlan, df.queryExecution.executedPlan)) { + val exchanges = collect(plan) { case exchange: ShuffleExchangeExec => exchange } + assert(exchanges.size == 1, s"$plan") + assert(exchanges.head.outputPartitioning match { + case HashPartitioning(Seq(key), 4) => key.dataType.isInstanceOf[ArrayType] + case _ => false + }) + assert(collect(plan) { case exchange: CometShuffleExchangeExec => + exchange + }.isEmpty) + val partials = collect(plan) { + case agg: BaseAggregateExec + if agg.aggregateExpressions.map(_.mode).distinct == Seq(Partial) => + agg + } + val finals = collect(plan) { + case agg: BaseAggregateExec + if agg.aggregateExpressions.map(_.mode).distinct == Seq(Final) => + agg + } + assert(finals.size == 1, s"$plan") + val nativeAggregates = collect(plan) { case agg: CometHashAggregateExec => agg } + if (fn == "sum") { + // SUM's Long buffer is safe for Spark's final, so retain its native partial. + assert(nativeAggregates.size == 1, s"$plan") + assert(nativeAggregates.head.modes == Seq(Partial)) + assert(partials.isEmpty, s"$plan") + } else { + assert(nativeAggregates.isEmpty, s"$plan") + assert(partials.size == 1, s"$plan") + assert(partials.head.getTagValue(CometExecRule.COMET_UNSAFE_PARTIAL).isDefined) + assert(collect(partials.head.child) { case project: CometProjectExec => + project + }.nonEmpty) + } + } + if (adaptive) { + val stages = collect(df.queryExecution.executedPlan) { + case stage: ShuffleQueryStageExec => stage + } + assert(stages.nonEmpty && stages.forall(_.isMaterialized)) + } + } + } + } + } + } + test("stddev_pop should return NaN for some cases") { withSQLConf(CometConf.COMET_SHUFFLE_ENABLED.key -> "true") { Seq(true, false).foreach { nullOnDivideByZero => @@ -267,15 +555,56 @@ class CometAggregateSuite extends CometTestBase with AdaptiveSparkPlanHelper { } } - test("mixed engine sum/avg: Comet partial + Spark final matches Spark") { + test("decimal SUM partial stays in Spark when a later input cancels precision overflow") { + // Keep all three values in one ordered input partition. Generated scalar Spark SUM can + // retain the temporary 1.2 and return 0.6 after cancellation; native decimal SUM instead + // makes that precision overflow sticky, or throws immediately in ANSI mode. A matching + // (sum, isEmpty) buffer schema therefore does not establish forward interoperability. + withSQLConf( + SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "false", + SQLConf.WHOLESTAGE_CODEGEN_ENABLED.key -> "true", + "spark.sql.files.minPartitionNum" -> "1", + CometConf.COMET_ENABLE_FINAL_HASH_AGGREGATE.key -> "false", + CometConf.COMET_SHUFFLE_ENABLED.key -> "false") { + withTempPath { path => + spark + .range(0, 3, 1, 1) + .selectExpr("CAST(CASE WHEN id < 2 THEN '0.6' ELSE '-0.6' END AS DECIMAL(38,38)) AS v") + .write + .parquet(path.getCanonicalPath) + withParquetTable(path.getCanonicalPath, "decimal_sum_cancellation") { + for (ansi <- Seq(false, true)) { + withSQLConf(SQLConf.ANSI_ENABLED.key -> ansi.toString) { + val df = sql("SELECT SUM(v) FROM decimal_sum_cancellation") + val plan = df.queryExecution.executedPlan + assert(collect(plan) { case agg: CometHashAggregateExec => agg }.isEmpty) + val partials = collect(plan) { + case agg: BaseAggregateExec + if agg.aggregateExpressions.map(_.mode).distinct == Seq(Partial) => + agg + } + assert(partials.size == 1) + assert(partials.head.getTagValue(CometExecRule.COMET_UNSAFE_PARTIAL).isDefined) + assert(collect(plan) { case native: CometNativeExec => native }.nonEmpty) + checkSparkAnswer(df) + checkAnswer(df, Seq(Row(new java.math.BigDecimal("0.6")))) + } + } + } + } + } + } + + test("mixed engine sum/avg falls back when Spark Final would consume native AVG") { val data = (0 until 100).map(i => (i, i.toLong, i.toDouble, i % 7)) withParquetTable(data, "tbl") { withSQLConf( CometConf.COMET_ENABLE_FINAL_HASH_AGGREGATE.key -> "false", CometConf.COMET_SHUFFLE_ENABLED.key -> "true", CometConf.COMET_SHUFFLE_MODE.key -> "jvm") { - checkSparkAnswer( - "SELECT _4, SUM(_1), SUM(_2), SUM(_3), AVG(_1), AVG(_2), AVG(_3) FROM tbl GROUP BY _4") + checkSparkAnswerAndNumOfAggregates( + "SELECT _4, SUM(_1), SUM(_2), SUM(_3), AVG(_1), AVG(_2), AVG(_3) FROM tbl GROUP BY _4", + 0) } } } @@ -518,7 +847,10 @@ class CometAggregateSuite extends CometTestBase with AdaptiveSparkPlanHelper { checkSparkAnswerAndNumOfAggregates("SELECT _2, COUNT(_1) FROM tbl GROUP BY _2", n) checkSparkAnswerAndNumOfAggregates("SELECT _2, MIN(_1) FROM tbl GROUP BY _2", n) checkSparkAnswerAndNumOfAggregates("SELECT _2, MAX(_1) FROM tbl GROUP BY _2", n) - checkSparkAnswerAndNumOfAggregates("SELECT _2, AVG(_1) FROM tbl GROUP BY _2", n) + val avgStages = if (nativeShuffleEnabled) 2 else 0 + checkSparkAnswerAndNumOfAggregates( + "SELECT _2, AVG(_1) FROM tbl GROUP BY _2", + avgStages) } } } @@ -723,26 +1055,29 @@ class CometAggregateSuite extends CometTestBase with AdaptiveSparkPlanHelper { val path = new Path(dir.toURI.toString, "test") makeParquetFile(path, 1000, 20, dictionaryEnabled) withParquetTable(path.toUri.toString, "tbl") { + // Spark rewrites _7's small decimal SUM to Long; _8 and _9 remain decimal and + // cannot use a native Partial when the Final runs in Spark. val expectedNumOfCometAggregates = if (nativeShuffleEnabled) 2 else 1 + val expectedNumOfDecimalAggregates = if (nativeShuffleEnabled) 2 else 0 checkSparkAnswerAndNumOfAggregates( "SELECT _g2, SUM(_7) FROM tbl GROUP BY _g2", expectedNumOfCometAggregates) checkSparkAnswerAndNumOfAggregates( "SELECT _g3, SUM(_8) FROM tbl GROUP BY _g3", - expectedNumOfCometAggregates) + expectedNumOfDecimalAggregates) checkSparkAnswerAndNumOfAggregates( "SELECT _g4, SUM(_9) FROM tbl GROUP BY _g4", - expectedNumOfCometAggregates) + expectedNumOfDecimalAggregates) checkSparkAnswerAndNumOfAggregates( "SELECT SUM(_7) FROM tbl", expectedNumOfCometAggregates) checkSparkAnswerAndNumOfAggregates( "SELECT SUM(_8) FROM tbl", - expectedNumOfCometAggregates) + expectedNumOfDecimalAggregates) checkSparkAnswerAndNumOfAggregates( "SELECT SUM(_9) FROM tbl", - expectedNumOfCometAggregates) + expectedNumOfDecimalAggregates) } } } @@ -1343,14 +1678,14 @@ class CometAggregateSuite extends CometTestBase with AdaptiveSparkPlanHelper { } } - test("test partial avg") { + test("AVG stays in Spark across a Spark shuffle") { Seq(true, false).foreach { dictionaryEnabled => withParquetTable( (0 until 5).map(i => (i.toDouble, i.toDouble % 2)), "tbl", dictionaryEnabled) { withSQLConf(CometConf.COMET_SHUFFLE_ENABLED.key -> "false") { - checkSparkAnswerAndNumOfAggregates("SELECT _2 , AVG(_1) FROM tbl GROUP BY _2", 1) + checkSparkAnswerAndNumOfAggregates("SELECT _2 , AVG(_1) FROM tbl GROUP BY _2", 0) } } } @@ -1387,7 +1722,9 @@ class CometAggregateSuite extends CometTestBase with AdaptiveSparkPlanHelper { val path = new Path(dir.toURI.toString, "test") makeParquetFile(path, 1000, 20, dictionaryEnabled) withParquetTable(path.toUri.toString, "tbl") { - val expectedNumOfCometAggregates = if (nativeShuffleEnabled) 2 else 1 + // Spark rewrites _7 to Long AVG, whose empty native buffer is also unsafe for a + // Spark Final until #5420. Keep all AVG partials in Spark across this boundary. + val expectedNumOfCometAggregates = if (nativeShuffleEnabled) 2 else 0 checkSparkAnswerAndNumOfAggregates( "SELECT _g2, AVG(_7) FROM tbl GROUP BY _g2", diff --git a/spark/src/test/scala/org/apache/comet/rules/CometExecRuleSuite.scala b/spark/src/test/scala/org/apache/comet/rules/CometExecRuleSuite.scala index 8825b8af097..f484de61b42 100644 --- a/spark/src/test/scala/org/apache/comet/rules/CometExecRuleSuite.scala +++ b/spark/src/test/scala/org/apache/comet/rules/CometExecRuleSuite.scala @@ -21,20 +21,23 @@ package org.apache.comet.rules import scala.util.Random +import org.apache.logging.log4j.Level import org.apache.spark.sql._ import org.apache.spark.sql.catalyst.FunctionIdentifier -import org.apache.spark.sql.catalyst.expressions.{Expression, ExpressionInfo} -import org.apache.spark.sql.catalyst.expressions.aggregate.BloomFilterAggregate +import org.apache.spark.sql.catalyst.expressions.{Attribute, Expression, ExpressionInfo, Literal} +import org.apache.spark.sql.catalyst.expressions.aggregate.{AggregateExpression, BloomFilterAggregate, Final, Min, Partial, PartialMerge} import org.apache.spark.sql.comet._ import org.apache.spark.sql.comet.execution.shuffle.CometShuffleExchangeExec import org.apache.spark.sql.execution._ -import org.apache.spark.sql.execution.adaptive.QueryStageExec +import org.apache.spark.sql.execution.adaptive.{QueryStageExec, ShuffleQueryStageExec} import org.apache.spark.sql.execution.aggregate.{HashAggregateExec, ObjectHashAggregateExec} import org.apache.spark.sql.execution.exchange.{BroadcastExchangeExec, ShuffleExchangeExec} +import org.apache.spark.sql.internal.SQLConf import org.apache.spark.sql.types.{DataTypes, StructField, StructType} -import org.apache.comet.CometConf -import org.apache.comet.CometSparkSessionExtensions.{isSpark35Plus, isSpark40Plus, isSpark42Plus} +import org.apache.comet.{CometConf, CometExplainInfo, ExtendedExplainInfo} +import org.apache.comet.CometSparkSessionExtensions.{isSpark35Plus, isSpark40Plus, isSpark42Plus, withFallbackReason} +import org.apache.comet.serde.{CometAggregateExpressionSerde, ExprOuterClass} import org.apache.comet.testing.{DataGenOptions, FuzzDataGenerator} /** @@ -135,8 +138,7 @@ class CometExecRuleSuite extends CometTestBase { } } - // Regression test for https://github.com/apache/datafusion-comet/issues/1389 - test("CometExecRule should not allow Comet partial and Spark final hash aggregate") { + test("CometExecRule should allow COUNT Comet partial and Spark final hash aggregate") { withTempView("test_data") { createTestDataFrame.createOrReplaceTempView("test_data") @@ -152,11 +154,10 @@ class CometExecRuleSuite extends CometTestBase { CometConf.COMET_EXEC_LOCAL_TABLE_SCAN_ENABLED.key -> "true") { val transformedPlan = applyCometExecRule(sparkPlan) - // COUNT is intentionally excluded from mixed execution (AQE / count-bug reasons), so if - // the final aggregate cannot be converted to Comet, neither should the partial. - assert( - countOperators(transformedPlan, classOf[HashAggregateExec]) == originalHashAggCount) - assert(countOperators(transformedPlan, classOf[CometHashAggregateExec]) == 0) + // COUNT's buffer is compatible in this direction. Keeping the Final in Spark also keeps + // the AQE/count-bug rewrites that prevent the reverse direction from being admitted. + assert(countOperators(transformedPlan, classOf[HashAggregateExec]) == 1) + assert(countOperators(transformedPlan, classOf[CometHashAggregateExec]) == 1) } } } @@ -177,8 +178,8 @@ class CometExecRuleSuite extends CometTestBase { CometConf.COMET_EXEC_LOCAL_TABLE_SCAN_ENABLED.key -> "true") { val transformedPlan = applyCometExecRule(sparkPlan) - // COUNT blocks mixed execution, so if the partial cannot be converted, neither should - // the final. + // COUNT still blocks Spark Partial to Comet Final, independently of the safe reverse + // direction, so if the partial cannot be converted, neither should the final. assert( countOperators(transformedPlan, classOf[HashAggregateExec]) == originalHashAggCount) assert(countOperators(transformedPlan, classOf[CometHashAggregateExec]) == 0) @@ -265,7 +266,7 @@ class CometExecRuleSuite extends CometTestBase { } } - test("CometExecRule should allow AVG mixed Comet partial and Spark final") { + test("CometExecRule should not allow AVG Comet partial and Spark final before buffer repair") { withTempView("test_data") { createTestDataFrame.createOrReplaceTempView("test_data") val sparkPlan = @@ -275,8 +276,9 @@ class CometExecRuleSuite extends CometTestBase { CometConf.COMET_ENABLE_FINAL_HASH_AGGREGATE.key -> "false", CometConf.COMET_EXEC_LOCAL_TABLE_SCAN_ENABLED.key -> "true") { val transformedPlan = applyCometExecRule(sparkPlan) - assert(countOperators(transformedPlan, classOf[HashAggregateExec]) == 1) // final - assert(countOperators(transformedPlan, classOf[CometHashAggregateExec]) == 1) // partial + // Matching field types do not make native AVG's empty (null, 0) state safe for Spark. + assert(countOperators(transformedPlan, classOf[HashAggregateExec]) == 2) + assert(countOperators(transformedPlan, classOf[CometHashAggregateExec]) == 0) } } } @@ -323,6 +325,175 @@ class CometExecRuleSuite extends CometTestBase { } } + for (distinct <- Seq(false, true)) { + test( + s"unsafe aggregate buffers fall back when native shuffle is ineligible (distinct=$distinct)") { + withTempView("test_data") { + createTestDataFrame.createOrReplaceTempView("test_data") + val aggregates = "AVG(id)" + (if (distinct) ", SUM(DISTINCT id)" else "") + + for (fallback <- Seq("disabled hash partitioning", "prior shuffle fallback", "none")) { + withSQLConf( + CometConf.COMET_EXEC_LOCAL_TABLE_SCAN_ENABLED.key -> "true", + CometConf.COMET_SHUFFLE_ENABLED.key -> "true", + CometConf.COMET_SHUFFLE_MODE.key -> "native", + CometConf.COMET_SHUFFLE_NATIVE_HASH_PARTITIONING_ENABLED.key -> + (fallback != "disabled hash partitioning").toString) { + val sparkPlan = + createSparkPlan(spark, s"SELECT $aggregates FROM test_data GROUP BY (id % 3)") + val aggregateCount = countOperators(sparkPlan, classOf[HashAggregateExec]) + assert(aggregateCount == (if (distinct) 4 else 2)) + if (fallback == "prior shuffle fallback") { + // Tag only the lowest exchange. A DISTINCT plan's upper exchange must inherit + // the native-only refusal from its now-Spark merge inputs, not from another tag. + val lowerShuffle = stripAQEPlan(sparkPlan).collect { + case shuffle: ShuffleExchangeExec => shuffle + }.last + withFallbackReason(lowerShuffle, fallback) + } + val transformed = applyCometExecRule(sparkPlan) + val nativeExpected = fallback == "none" + + // Shuffle is enabled, but a native-only shuffle can still fall back. The distinct + // rewrite also has intermediate PartialMerge and mixed Partial/PartialMerge stages. + for (plan <- Seq(transformed, applyCometExecRule(transformed))) { + assert( + countOperators(plan, classOf[CometHashAggregateExec]) == + (if (nativeExpected) aggregateCount else 0)) + assert( + countOperators(plan, classOf[HashAggregateExec]) == + (if (nativeExpected) 0 else aggregateCount)) + } + // AQE reapplies the rule to an exchange without its Final aggregate. The tagged + // Partial must remain in Spark in that stage-only pass too. + transformed.collect { case shuffle: ShuffleExchangeExec => shuffle }.foreach { + shuffle => + val stage = applyCometExecRule(shuffle) + assert(countOperators(stage, classOf[CometHashAggregateExec]) == 0) + } + } + } + } + } + } + + test("aggregate buffer direction opt-ins are independent") { + // A policy-only handler opts into consuming Spark state. Its inherited producer policy + // must stay false; serializing any expression is outside the scope of this fixture. + val reverseOnly = new CometAggregateExpressionSerde[Min] { + override def supportsSparkPartialToNativeFinal(fn: Min): Boolean = true + + override def convert( + aggExpr: AggregateExpression, + expr: Min, + inputs: Seq[Attribute], + binding: Boolean, + conf: SQLConf): Option[ExprOuterClass.AggExpr] = None + } + val fn = Min(Literal(1L)) + assert(reverseOnly.supportsSparkPartialToNativeFinal(fn)) + assert(!reverseOnly.supportsNativePartialToSparkFinal(fn)) + } + + test("restored partial records its reason when its current child is not native") { + // Wrap a converted input to prevent a re-entrant serde call from supplying the reason. + // This planner-only fixture never executes its synthetic buffer boundary. + withSQLConf( + SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "false", + SQLConf.WHOLESTAGE_CODEGEN_ENABLED.key -> "false", + CometConf.COMET_EXEC_LOCAL_TABLE_SCAN_ENABLED.key -> "true", + CometConf.COMET_SHUFFLE_MODE.key -> "native") { + withTempView("test_data") { + createTestDataFrame.createOrReplaceTempView("test_data") + val plan = applyCometExecRule( + createSparkPlan(spark, "SELECT AVG(id) FROM test_data GROUP BY (id % 3)")) + val partial = plan.collectFirst { + case agg: CometHashAggregateExec if agg.modes == Seq(Partial) => agg + }.get + val sparkFinal = plan.collectFirst { + case agg: CometHashAggregateExec if agg.modes == Seq(Final) => + agg.originalPlan.asInstanceOf[HashAggregateExec] + }.get + val nonNativeChild = InputAdapter(partial.child) + val restored = CometExecRule(spark).revertUnsafePartialAggregates( + sparkFinal.copy(child = partial.copy(child = nonNativeChild))) + val sparkPartial = restored.children.head + assert(sparkPartial.isInstanceOf[HashAggregateExec]) + assert(sparkPartial.children.head.isInstanceOf[InputAdapter]) + assert(sparkPartial.children.head.output == nonNativeChild.output) + val reason = sparkPartial.getTagValue(CometExecRule.COMET_UNSAFE_PARTIAL).get + assert(sparkPartial.getTagValue(CometExplainInfo.FALLBACK_REASONS).get.contains(reason)) + assert(new ExtendedExplainInfo().getFallbackReasons(sparkPartial).contains(reason)) + } + } + } + + test("unrepaired native aggregate buffers warn once without rewriting query stages") { + // Construct the stage placeholder emitted by CometExchangeSink, including a native merge + // above it. No SQL reproduction or materialization is assumed: this pins the diagnostic + // when repair stops at a stage, and the absence of warnings for unrelated inner producers. + withSQLConf( + SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "false", + SQLConf.WHOLESTAGE_CODEGEN_ENABLED.key -> "false", + CometConf.COMET_EXEC_LOCAL_TABLE_SCAN_ENABLED.key -> "true", + CometConf.COMET_SHUFFLE_MODE.key -> "native") { + withTempView("test_data") { + createTestDataFrame.createOrReplaceTempView("test_data") + val plan = applyCometExecRule( + createSparkPlan(spark, "SELECT AVG(id) FROM test_data GROUP BY (id % 3)")) + val partial = plan.collectFirst { + case agg: CometHashAggregateExec if agg.modes == Seq(Partial) => agg + }.get + val nativeFinal = plan.collectFirst { + case agg: CometHashAggregateExec if agg.modes == Seq(Final) => agg + }.get + val sparkFinal = nativeFinal.originalPlan.asInstanceOf[HashAggregateExec] + val sparkPartial = partial.originalPlan.asInstanceOf[HashAggregateExec] + val exchange = ShuffleExchangeExec( + org.apache.spark.sql.catalyst.plans.physical.SinglePartition, + partial) + val stage = ShuffleQueryStageExec(0, exchange, exchange.canonicalized) + val placeholder = CometSinkPlaceHolder( + org.apache.comet.serde.OperatorOuterClass.Operator.getDefaultInstance, + stage, + stage) + val nativeMerge = partial.copy( + aggregateExpressions = partial.aggregateExpressions.map(_.copy(mode = PartialMerge)), + child = placeholder) + val warning = "Comet could not restore a native intermediate buffer producer" + val rule = CometExecRule(spark) + for { + (child, shouldWarn) <- Seq( + nativeMerge -> true, + placeholder -> true, + sparkPartial.copy(child = nativeFinal) -> false, + nativeFinal -> false) + logFallback <- Seq("false", "true") + } { + withSQLConf(CometConf.COMET_EXPLAIN_FALLBACK_LOG_ENABLED.key -> logFallback) { + val consumer = sparkFinal.copy(child = child) + val appender = new LogAppender("unrepaired aggregate buffers") + withLogAppender(appender, Seq("org.apache.comet"), Some(Level.WARN)) { + assert(rule.revertUnsafePartialAggregates(consumer) eq consumer) + assert(rule.revertUnsafePartialAggregates(consumer) eq consumer) + } + assert(consumer.child eq child) + val warnings = + appender.loggingEvents.count(_.getMessage.getFormattedMessage.contains(warning)) + assert(warnings == (if (shouldWarn) 1 else 0), s"$child: $warnings") + assert( + new ExtendedExplainInfo() + .getFallbackReasons(consumer) + .exists(_.contains(warning)) == + shouldWarn) + } + } + assert(stage.plan eq exchange) + assert(exchange.child eq partial) + } + } + } + test("CometExecRule should not allow decimal SUM mixed execution") { withTempView("test_data") { createTestDataFrame.createOrReplaceTempView("test_data") @@ -338,9 +509,9 @@ class CometExecRuleSuite extends CometTestBase { CometConf.COMET_ENABLE_FINAL_HASH_AGGREGATE.key -> "false", CometConf.COMET_EXEC_LOCAL_TABLE_SCAN_ENABLED.key -> "true") { val transformedPlan = applyCometExecRule(sparkPlan) - // Decimal SUM overflow detection (ANSI throw / Legacy null) does not survive a - // Spark-partial / Comet-final split, so mixed execution is unsafe and the partial - // must also fall back to Spark. + // Native decimal SUM makes precision overflow sticky (or throws eagerly in ANSI), + // while Spark's generated scalar Partial can recover after a later cancelling input. + // Keep the Partial in Spark even though the emitted buffer field types match. assert(countOperators(transformedPlan, classOf[HashAggregateExec]) == 2) assert(countOperators(transformedPlan, classOf[CometHashAggregateExec]) == 0) } diff --git a/spark/src/test/scala/org/apache/comet/rules/RevertNativeForTransitionHeavyStagesSuite.scala b/spark/src/test/scala/org/apache/comet/rules/RevertNativeForTransitionHeavyStagesSuite.scala index fb66ee8193b..de80f493616 100644 --- a/spark/src/test/scala/org/apache/comet/rules/RevertNativeForTransitionHeavyStagesSuite.scala +++ b/spark/src/test/scala/org/apache/comet/rules/RevertNativeForTransitionHeavyStagesSuite.scala @@ -20,9 +20,12 @@ package org.apache.comet.rules import org.apache.spark.sql.CometTestBase +import org.apache.spark.sql.catalyst.expressions.aggregate.{Final, Partial} import org.apache.spark.sql.comet._ import org.apache.spark.sql.comet.execution.shuffle.CometShuffleExchangeExec import org.apache.spark.sql.execution._ +import org.apache.spark.sql.execution.adaptive.QueryStageExec +import org.apache.spark.sql.internal.SQLConf import org.apache.comet.CometConf @@ -55,6 +58,18 @@ class RevertNativeForTransitionHeavyStagesSuite extends CometTestBase { plan.collect { case _: ColumnarToRowTransition => true }.size } + private def collectCometAggregates(plan: SparkPlan): Seq[CometHashAggregateExec] = { + val current = plan match { + case aggregate: CometHashAggregateExec => Seq(aggregate) + case _ => Seq.empty + } + val descendants = plan match { + case stage: QueryStageExec => collectCometAggregates(stage.plan) + case _ => plan.children.flatMap(collectCometAggregates) + } + current ++ descendants + } + /** * Returns every node that produces a columnar output but consumes a row-based child without a * RowToColumnar transition. Such a node is an invalid columnar/row boundary: a columnar parent @@ -170,7 +185,7 @@ class RevertNativeForTransitionHeavyStagesSuite extends CometTestBase { test("revert fires with unsupported UDF producing transitions") { withParquetTable((0 until 100).map(i => (i, i % 10, s"val_$i")), "tbl") { spark.udf.register("identity_udf", (x: Int) => x) - val query = "SELECT _2, identity_udf(_1), count(*) FROM tbl GROUP BY _2, identity_udf(_1)" + val query = "SELECT _2, identity_udf(_1), max(_1) FROM tbl GROUP BY _2, identity_udf(_1)" // Without revert, plan should have transitions due from UDF withSQLConf(CometConf.COMET_EXEC_TRANSITION_REVERT_ENABLED.key -> "false") { @@ -195,7 +210,7 @@ class RevertNativeForTransitionHeavyStagesSuite extends CometTestBase { test("revert fires and produces correct results when transitions exceed threshold") { withParquetTable((0 until 100).map(i => (i, i % 10, s"val_$i")), "tbl") { - val query = "SELECT _2, count(*), sum(_1) FROM tbl GROUP BY _2" + val query = "SELECT _2, min(_1), sum(_1) FROM tbl GROUP BY _2" // Without revert, plan should have CometExec nodes with transitions withSQLConf( @@ -279,4 +294,140 @@ class RevertNativeForTransitionHeavyStagesSuite extends CometTestBase { } } } + + for (adaptive <- Seq(false, true)) { + test(s"transition reversion preserves incompatible aggregate buffers with AQE=$adaptive") { + withSQLConf( + SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> adaptive.toString, + SQLConf.WHOLESTAGE_CODEGEN_ENABLED.key -> "false", + CometConf.COMET_SHUFFLE_MODE.key -> "native", + CometConf.COMET_EXEC_TRANSITION_REVERT_ENABLED.key -> "true", + CometConf.COMET_EXEC_TRANSITION_REVERT_MAX_TRANSITIONS.key -> "0") { + withParquetTable((0 until 256).map(i => (i % 4, i.toDouble)), "tbl") { + val (_, plan) = + checkSparkAnswer("SELECT _1, percentile(_2, 0.5) FROM tbl GROUP BY _1 ORDER BY _1") + val executedPlan = stripAQEPlan(plan) + val aggregates = collectCometAggregates(executedPlan) + assert(aggregates.exists(_.modes == Seq(Partial)), s"$executedPlan") + assert(aggregates.exists(_.modes == Seq(Final)), s"$executedPlan") + } + } + } + } + + test("transition reversion finds an incompatible aggregate below another aggregate") { + withSQLConf( + SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "false", + SQLConf.WHOLESTAGE_CODEGEN_ENABLED.key -> "false", + CometConf.COMET_SHUFFLE_MODE.key -> "native", + CometConf.COMET_EXEC_TRANSITION_REVERT_ENABLED.key -> "true", + CometConf.COMET_EXEC_TRANSITION_REVERT_MAX_TRANSITIONS.key -> "0") { + withParquetTable((0 until 256).map(i => (i % 4, i.toDouble)), "tbl") { + val query = + """SELECT grouping_key, percentile(inner_percentile, 0.5) + |FROM ( + | SELECT _1 AS grouping_key, percentile(_2, 0.5) AS inner_percentile + | FROM tbl + | GROUP BY _1 + |) inner_aggregate + |GROUP BY grouping_key""".stripMargin + val (_, plan) = checkSparkAnswer(query) + val executedPlan = stripAQEPlan(plan) + val aggregates = collectCometAggregates(executedPlan) + assert( + executedPlan.collect { case _: CometShuffleExchangeExec => true }.size == 1, + s"test requires one exchange below the nested aggregates:\n$executedPlan") + assert(aggregates.count(_.modes == Seq(Partial)) == 2, s"$executedPlan") + assert(aggregates.count(_.modes == Seq(Final)) == 2, s"$executedPlan") + } + } + } + + for (adaptive <- Seq(false, true)) { + test(s"transition reversion does not split native COUNT stages with AQE=$adaptive") { + withSQLConf( + SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> adaptive.toString, + SQLConf.WHOLESTAGE_CODEGEN_ENABLED.key -> "false", + CometConf.COMET_SHUFFLE_MODE.key -> "native", + CometConf.COMET_EXEC_TRANSITION_REVERT_ENABLED.key -> "true", + CometConf.COMET_EXEC_TRANSITION_REVERT_MAX_TRANSITIONS.key -> "0") { + withParquetTable((0 until 256).map(i => (i % 4, i)), "tbl") { + val (_, plan) = checkSparkAnswer("SELECT _1, count(*) FROM tbl GROUP BY _1 ORDER BY _1") + val executedPlan = stripAQEPlan(plan) + val aggregates = collectCometAggregates(executedPlan) + assert(aggregates.exists(_.modes == Seq(Partial)), s"$executedPlan") + assert(aggregates.exists(_.modes == Seq(Final)), s"$executedPlan") + } + } + } + } + + test("transition reversion preserves an incompatible native partial producer") { + withSQLConf( + SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "false", + CometConf.COMET_SHUFFLE_MODE.key -> "native", + CometConf.COMET_EXEC_TRANSITION_REVERT_ENABLED.key -> "false") { + withParquetTable((0 until 256).map(i => (i % 4, i.toDouble)), "tbl") { + val nativePlan = + sql("SELECT _1, percentile(_2, 0.5) FROM tbl GROUP BY _1").queryExecution.executedPlan + val partial = nativePlan + .collectFirst { + case aggregate: CometHashAggregateExec if aggregate.modes == Seq(Partial) => aggregate + } + .getOrElse(fail(s"expected a native partial aggregate:\n$nativePlan")) + val producerWithTransition = partial.withNewChildren( + Seq(CometSparkToColumnarExec(CometNativeColumnarToRowExec(partial.child)))) + val reverter = RevertNativeForTransitionHeavyStages(spark) + assert(reverter.countTransitions(producerWithTransition) == 1) + + var reverted: SparkPlan = null + withSQLConf( + CometConf.COMET_EXEC_TRANSITION_REVERT_ENABLED.key -> "true", + CometConf.COMET_EXEC_TRANSITION_REVERT_MAX_TRANSITIONS.key -> "0") { + reverted = reverter(producerWithTransition) + } + assert(reverted eq producerWithTransition) + } + } + } + + for (adaptive <- Seq(false, true)) { + test(s"transition reversion keeps AVG and collect_list Finals native with AQE=$adaptive") { + withTempDir { dir => + val path = s"${dir.getAbsolutePath}/data" + withSQLConf(CometConf.COMET_ENABLED.key -> "false") { + spark + .range(0L, 8L, 1L, 4) + .selectExpr("id", "id % 2 AS k", "CAST(1 AS BIGINT) AS v") + .write + .parquet(path) + } + withSQLConf( + SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> adaptive.toString, + SQLConf.FILES_MAX_PARTITION_BYTES.key -> "1048576", + SQLConf.PARQUET_FILTER_PUSHDOWN_ENABLED.key -> "false", + CometConf.COMET_SHUFFLE_MODE.key -> "native", + CometConf.COMET_EXEC_TRANSITION_REVERT_ENABLED.key -> "true", + CometConf.COMET_EXEC_TRANSITION_REVERT_MAX_TRANSITIONS.key -> "0") { + withParquetTable(path, "tbl") { + assert(sql("SELECT * FROM tbl").rdd.getNumPartitions == 4) + // Threshold 0 makes the stage holding the Final a revert candidate, while the + // Partial stage below it has no transitions and stays native. The filter leaves + // three scan partitions empty, and their native AVG buffers are (null, 0), so a + // Spark Final would return null. A Spark Final also cannot read the native + // collect_list buffer. + for (query <- Seq( + "SELECT AVG(v) FROM tbl WHERE id = 1", + "SELECT k, sort_array(collect_list(id)) FROM tbl GROUP BY k")) { + val (_, plan) = checkSparkAnswer(query) + val executedPlan = stripAQEPlan(plan) + val aggregates = collectCometAggregates(executedPlan) + assert(aggregates.exists(_.modes == Seq(Partial)), s"$executedPlan") + assert(aggregates.exists(_.modes == Seq(Final)), s"$executedPlan") + } + } + } + } + } + } }