From bef0b9f6a35a24f0c0f52be323b311e73d05c1a1 Mon Sep 17 00:00:00 2001 From: msaf Date: Tue, 6 Oct 2026 11:45:09 +0100 Subject: [PATCH 01/13] fix: price the join condition of a sort-merge join in the cost-based engine choice A native sort-merge join with a join condition builds every pair of rows of equal keys, with all their output columns, before the condition drops them, while Spark tests the condition first. fbj_order_type's left band joins took 23 to 28 us per output row natively against 6 to 8 in Spark, and a local micro-benchmark of left band joins shows the same 3.6 to 4 times ratio, but the model priced only the smj line, which favours Comet, and kept them native. A class smjCondition, Line(3000, 80, 0, 850, 23), is added on top of smj over the output leaves of a sort-merge join that has a condition; joins without one keep their prices. Co-Authored-By: Claude Opus 5.5 (1M context) --- docs/source/user-guide/latest/tuning.md | 4 +- .../scala/org/apache/comet/CometConf.scala | 10 +-- .../comet/rules/CostBasedEngineChoice.scala | 3 + .../apache/comet/rules/EngineCostTable.scala | 17 ++++- .../rules/CostBasedEngineChoiceSuite.scala | 65 +++++++++++++++++++ 5 files changed, 91 insertions(+), 8 deletions(-) diff --git a/docs/source/user-guide/latest/tuning.md b/docs/source/user-guide/latest/tuning.md index ce3a8006050..b599bb74f8c 100644 --- a/docs/source/user-guide/latest/tuning.md +++ b/docs/source/user-guide/latest/tuning.md @@ -602,7 +602,9 @@ prices per row from a table of measurements: `c0 + k0*L + k1*L*min(L, 600)` ns n Spark one. - `sort` over every leaf of the sorted rows (Comet also 0.15 ns per byte beyond 12 per leaf); `sortSpill` adds the price of a spill to a fraction `sortSpillFraction` of rows, none by default. `smj` prices a sort-merge join and `bhj` - the probe side of a broadcast hash join, over every output leaf. + the probe side of a broadcast hash join, over every output leaf. A sort-merge join with a join condition adds + `smjCondition`, also over every output leaf: Comet builds every pair of rows of equal keys before the condition drops + them, about 3.5 times Spark's price on band joins. - `predicate` for filters, over the leaves their predicate references, plus a pass-through of 1.5 ns per output leaf natively and none in Spark. A native filter over a native scan, and the native projects over it, stay native whatever their prices (`keepFiltersOverNativeScans=false` lets them move): the rows a filter drops are not estimated, diff --git a/spark/src/main/scala/org/apache/comet/CometConf.scala b/spark/src/main/scala/org/apache/comet/CometConf.scala index acc6e57215f..9c10c3e88ee 100644 --- a/spark/src/main/scala/org/apache/comet/CometConf.scala +++ b/spark/src/main/scala/org/apache/comet/CometConf.scala @@ -689,11 +689,11 @@ object CometConf extends ShimCometConf { "(or `k0,k1`, keeping c0), or spark, with `c0,k` for c0 + k*L ns per row, where L is " + "the number of leaf columns the class prices (for the functions of an aggregate or a " + "window, the number of functions). The classes are shuffleWrite, shuffleRead, sort, " + - "sortSpill, smj, bhj, predicate, projectPassThrough, expression, agg, aggObjectHash, " + - "aggDeclarative, aggCollectList, aggCollectSet, aggPercentile, aggPercentileApprox, " + - "aggOther, aggArrayKey, codegenDispatch, window, windowAggregate, windowOffset, " + - "windowRank, wglPartial, wglFinal, expand, generate, rowLocal, the comet-only c2r and " + - "r2c, and the spark-only " + + "sortSpill, smj, smjCondition, bhj, predicate, projectPassThrough, expression, agg, " + + "aggObjectHash, aggDeclarative, aggCollectList, aggCollectSet, aggPercentile, " + + "aggPercentileApprox, aggOther, aggArrayKey, codegenDispatch, window, " + + "windowAggregate, windowOffset, windowRank, wglPartial, wglFinal, expand, generate, " + + "rowLocal, the comet-only c2r and r2c, and the spark-only " + "expressionOverScan, aggDeclarativeNoCodegen, expandNoCodegen and " + "generateNoCodegen. A row whose leaves are a fraction f inside structs, arrays or " + "maps costs (1 - f) times the flat price plus f times the nested one. The scalars " + diff --git a/spark/src/main/scala/org/apache/comet/rules/CostBasedEngineChoice.scala b/spark/src/main/scala/org/apache/comet/rules/CostBasedEngineChoice.scala index e9a89cff70c..5051354a1cd 100644 --- a/spark/src/main/scala/org/apache/comet/rules/CostBasedEngineChoice.scala +++ b/spark/src/main/scala/org/apache/comet/rules/CostBasedEngineChoice.scala @@ -33,6 +33,7 @@ import org.apache.spark.sql.comet.{CometExec, CometFilterExec, CometHashAggregat import org.apache.spark.sql.execution.{ColumnarToRowTransition, ExpandExec, FilterExec, ProjectExec, SortExec, SparkPlan} import org.apache.spark.sql.execution.aggregate.BaseAggregateExec import org.apache.spark.sql.execution.exchange.ShuffleExchangeLike +import org.apache.spark.sql.execution.joins.SortMergeJoinExec import org.apache.spark.sql.execution.window.WindowExec import org.apache.spark.sql.internal.SQLConf @@ -170,6 +171,8 @@ class EngineCostModel( functions.distinct.map { c => Term(c, Width(functions.size, 0), share * functions.count(_ == c)) } + case (SmjCondition, join: SortMergeJoinExec) => + if (join.condition.isDefined) Seq(Term(SmjCondition, out)) else Nil case (AggObjectHash, agg: BaseAggregateExec) => Seq(Term(AggObjectHash, Width(0, 0), aggregateShare(agg))) case _ => Seq(Term(costClass, out)) diff --git a/spark/src/main/scala/org/apache/comet/rules/EngineCostTable.scala b/spark/src/main/scala/org/apache/comet/rules/EngineCostTable.scala index 1c4810f2cf5..f0a74f4f571 100644 --- a/spark/src/main/scala/org/apache/comet/rules/EngineCostTable.scala +++ b/spark/src/main/scala/org/apache/comet/rules/EngineCostTable.scala @@ -172,6 +172,7 @@ object EngineCostTable { case object Sort extends CostClass("sort") case object SortSpill extends CostClass("sortSpill") case object Smj extends CostClass("smj") + case object SmjCondition extends CostClass("smjCondition") case object Bhj extends CostClass("bhj") case object Predicate extends CostClass("predicate") case object ProjectPassThrough extends CostClass("projectPassThrough") @@ -207,6 +208,7 @@ object EngineCostTable { Sort, SortSpill, Smj, + SmjCondition, Bhj, Predicate, ProjectPassThrough, @@ -301,7 +303,17 @@ object EngineCostTable { * cannot be predicted. * - `smj`: the join over its sorted inputs, every output leaf, noisy. `bhj`: the probe side, * every output leaf; the nested `bhj` is noisy (Comet 0 to 350, Spark 50 to 4200 ns) and - * takes the flat line. + * takes the flat line. `smjCondition`: what a join condition adds to `smj`, every output + * leaf. Comet joins every pair of rows of equal keys into a batch of all its output columns + * before the condition drops most of them; Spark tests the condition on the pair first. In + * a local micro-benchmark of left joins on a 14-day band (1 to 4% of 20 to 100 million + * pairs passing), Comet cost what the join without the condition cost, 0.44 to 1.2 us per + * output row and 21 ns per output leaf, Spark 3.6 to 4 times less (0.12 to 0.30 us, 4.5 ns + * per leaf); inner joins and small key groups ran 1 to 1.8 times slower. fbj_order_type's + * left band joins took 23 to 28 us per output row natively against 6 to 8 in Spark. The + * pairs per output row are not estimated, so the line keeps the ratio of 3.5 at a price + * between the two, high enough to outweigh `smj` and the conversions around the join at any + * width. * - `predicate`: a filter over the leaves its predicate reads, Spark noisy (maxrel 0.5 to * 0.8); passing rows costs the filter scalars. * - `projectPassThrough`: a project over every output leaf. Spark copies the row only when @@ -385,6 +397,7 @@ object EngineCostTable { (R2C, Flat) -> Line(3.4, 7.67, 0.036, 0, 0), (R2C, Nested) -> Line(0, 9.74, 0.040, 0, 0)) ++ both(ExprOverScan, Line(0, 0, 0, 2.5, 0)) ++ + both(SmjCondition, Line(3000, 80, 0, 850, 23)) ++ both(AggObjectHash, Line(1000, 0, 0, 2500, 0)) ++ both(AggDeclarative, Line(6, 0, 0, 15, 0)) ++ both(AggDeclarativeNoCodegen, Line(0, 0, 0, 170, 0)) ++ @@ -447,7 +460,7 @@ object EngineCostTable { */ val operatorClasses: Map[String, Seq[CostClass]] = Map( "SortExec" -> Seq(Sort), - "SortMergeJoinExec" -> Seq(Smj), + "SortMergeJoinExec" -> Seq(Smj, SmjCondition), "BroadcastHashJoinExec" -> Seq(Bhj), "WindowExec" -> Seq(Window), "WindowGroupLimitExec" -> Seq(WglPartial, WglFinal), diff --git a/spark/src/test/scala/org/apache/comet/rules/CostBasedEngineChoiceSuite.scala b/spark/src/test/scala/org/apache/comet/rules/CostBasedEngineChoiceSuite.scala index 16a0ec8d8a0..a3a987b7bab 100644 --- a/spark/src/test/scala/org/apache/comet/rules/CostBasedEngineChoiceSuite.scala +++ b/spark/src/test/scala/org/apache/comet/rules/CostBasedEngineChoiceSuite.scala @@ -878,6 +878,71 @@ class CostBasedEngineChoiceSuite extends CometTestBase { } } + test("a sort-merge join adds smjCondition only with a join condition") { + withTables { + withSQLConf( + SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "false", + SQLConf.AUTO_BROADCASTJOIN_THRESHOLD.key -> "-1", + flag -> "false") { + def join(query: String): CometSortMergeJoinExec = { + val plan = run(query) + val joins = nodes(plan).collect { case j: CometSortMergeJoinExec => j } + assert(joins.size == 1, s"plan:\n$plan") + joins.head + } + val out = Width(4, 0) + val equi = join("SELECT a.k, a.v, b.v, b.k FROM t a JOIN t b ON a.k = b.k") + assert(model.terms(equi, Engine.Comet) == Seq(Term(Smj, out))) + assert(model.terms(equi, Engine.Spark) == Seq(Term(Smj, out))) + same(model.operatorPrice(equi, Engine.Comet), EngineCostTable.default.comet(Smj, out)) + + val band = join( + "SELECT a.k, a.v, b.v, b.k FROM t a JOIN t b ON a.k = b.k " + + "AND a.v <= b.v AND b.v <= a.v + 13") + assert(model.terms(band, Engine.Comet) == Seq(Term(Smj, out), Term(SmjCondition, out))) + assert(model.terms(band, Engine.Spark) == Seq(Term(Smj, out), Term(SmjCondition, out))) + val table = EngineCostTable.default + for (engine <- Engine.all) { + val price: (CostClass, Width) => Double = + if (engine == Engine.Comet) table.comet else table.spark + same(model.operatorPrice(band, engine), price(Smj, out) + price(SmjCondition, out)) + } + } + } + } + + for (aqe <- Seq("false", "true")) { + test( + s"a sort-merge join with a join condition runs in Spark, without one natively (AQE=$aqe)") { + withTables { + withAqe( + aqe, + SQLConf.AUTO_BROADCASTJOIN_THRESHOLD.key -> "-1", + SQLConf.ADAPTIVE_AUTO_BROADCASTJOIN_THRESHOLD.key -> "-1") { + def joins(plan: SparkPlan): (Int, Int) = + ( + count(plan) { case j: CometSortMergeJoinExec => j }, + count(plan) { case j: SortMergeJoinExec => j }) + val equi = "SELECT a.k, a.v, b.s FROM t a JOIN t b ON a.k = b.k" + val band = equi + " AND a.v <= b.v AND b.v <= a.v + 13" + val (equiOff, equiOn) = offAndOn(run(equi)) + assert(joins(equiOff) == (1, 0), s"plan:\n$equiOff") + assert(joins(equiOn) == (1, 0), s"plan:\n$equiOn") + assert(cometOperatorNames(equiOff) == cometOperatorNames(equiOn), s"$equiOff\n$equiOn") + val (bandOff, bandOn) = offAndOn(run(band)) + assert(joins(bandOff) == (1, 0), s"plan:\n$bandOff") + assert(joins(bandOn) == (0, 1), s"plan:\n$bandOn") + withSQLConf( + flag -> "true", + costTable -> "smjCondition.comet=0,0,0;smjCondition.spark=0,0") { + val plan = run(band) + assert(joins(plan) == (1, 0), s"plan:\n$plan") + } + } + } + } + } + test("a window costs its line and the classes of its functions") { withTables { withSQLConf(SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "false", flag -> "false") { From 250bfd2e189153ba62a54408e88c23ea0215a684 Mon Sep 17 00:00:00 2001 From: msaf Date: Tue, 6 Oct 2026 12:22:29 +0100 Subject: [PATCH 02/13] fix: leave validity-interval conditions of sort-merge joins out of smjCondition A join on an SCD validity interval, L <= V < U with V from one input and L, U from the other, runs about as fast natively as in Spark: few pairs per key fail the condition. smjCondition is now skipped when the condition is exactly one such interval, recognized by expression structure only: comparisons in either order, BETWEEN, casts and date truncations, COALESCE(U, literal) and U IS NULL OR V < U. U must not be L shifted by a constant through the aliases of the plan below the join; a band such as v BETWEEN ts - 30 days AND ts, a CTE projecting l + 30 days as the upper bound, a one-sided bound, two lower bounds, bounds from different inputs or any extra conjunct keep the charge. Co-Authored-By: Claude Opus 5.5 (1M context) --- docs/source/user-guide/latest/tuning.md | 4 +- .../comet/rules/CostBasedEngineChoice.scala | 14 +- .../apache/comet/rules/EngineCostTable.scala | 3 +- .../comet/rules/JoinConditionShape.scala | 146 ++++++++++++++++++ .../rules/CostBasedEngineChoiceSuite.scala | 108 +++++++++++++ 5 files changed, 272 insertions(+), 3 deletions(-) create mode 100644 spark/src/main/scala/org/apache/comet/rules/JoinConditionShape.scala diff --git a/docs/source/user-guide/latest/tuning.md b/docs/source/user-guide/latest/tuning.md index b599bb74f8c..13fabf00bb5 100644 --- a/docs/source/user-guide/latest/tuning.md +++ b/docs/source/user-guide/latest/tuning.md @@ -604,7 +604,9 @@ prices per row from a table of measurements: `c0 + k0*L + k1*L*min(L, 600)` ns n price of a spill to a fraction `sortSpillFraction` of rows, none by default. `smj` prices a sort-merge join and `bhj` the probe side of a broadcast hash join, over every output leaf. A sort-merge join with a join condition adds `smjCondition`, also over every output leaf: Comet builds every pair of rows of equal keys before the condition drops - them, about 3.5 times Spark's price on band joins. + them, about 3.5 times Spark's price on band joins. A condition that is one validity interval, `L <= V < U` with `V` + from one input and `L` and `U` from the other (casts, date truncations, `COALESCE(U, literal)` and `U IS NULL OR` + allowed), adds nothing when `U` is not `L` shifted by a constant through the projections below the join. - `predicate` for filters, over the leaves their predicate references, plus a pass-through of 1.5 ns per output leaf natively and none in Spark. A native filter over a native scan, and the native projects over it, stay native whatever their prices (`keepFiltersOverNativeScans=false` lets them move): the rows a filter drops are not estimated, diff --git a/spark/src/main/scala/org/apache/comet/rules/CostBasedEngineChoice.scala b/spark/src/main/scala/org/apache/comet/rules/CostBasedEngineChoice.scala index 5051354a1cd..5292c914a05 100644 --- a/spark/src/main/scala/org/apache/comet/rules/CostBasedEngineChoice.scala +++ b/spark/src/main/scala/org/apache/comet/rules/CostBasedEngineChoice.scala @@ -124,6 +124,14 @@ class EngineCostModel( private def aggregateShare(agg: BaseAggregateExec): Double = if (agg.aggregateExpressions.exists(_.mode == Complete)) 1.0 else 0.5 + /** Whether `condition` of `join`, run as `plan`, is one validity interval. */ + def validityInterval(condition: Expression, join: SortMergeJoinExec, plan: SparkPlan): Boolean = + JoinConditionShape.isValidityInterval( + condition, + join.left.outputSet, + join.right.outputSet, + JoinConditionShape.aliases(plan.children ++ join.children)) + private def classTerms(costClass: CostClass, plan: SparkPlan, engine: Engine): Seq[Term] = { val op = sparkOperator(plan) lazy val out = widthOf(op.output) @@ -172,7 +180,11 @@ class EngineCostModel( Term(c, Width(functions.size, 0), share * functions.count(_ == c)) } case (SmjCondition, join: SortMergeJoinExec) => - if (join.condition.isDefined) Seq(Term(SmjCondition, out)) else Nil + if (join.condition.exists(c => !validityInterval(c, join, plan))) { + Seq(Term(SmjCondition, out)) + } else { + Nil + } case (AggObjectHash, agg: BaseAggregateExec) => Seq(Term(AggObjectHash, Width(0, 0), aggregateShare(agg))) case _ => Seq(Term(costClass, out)) diff --git a/spark/src/main/scala/org/apache/comet/rules/EngineCostTable.scala b/spark/src/main/scala/org/apache/comet/rules/EngineCostTable.scala index f0a74f4f571..7d1c5eaae7a 100644 --- a/spark/src/main/scala/org/apache/comet/rules/EngineCostTable.scala +++ b/spark/src/main/scala/org/apache/comet/rules/EngineCostTable.scala @@ -313,7 +313,8 @@ object EngineCostTable { * left band joins took 23 to 28 us per output row natively against 6 to 8 in Spark. The * pairs per output row are not estimated, so the line keeps the ratio of 3.5 at a price * between the two, high enough to outweigh `smj` and the conversions around the join at any - * width. + * width. A condition that is one validity interval ([[JoinConditionShape]]), where Comet + * costs about what Spark does, adds nothing. * - `predicate`: a filter over the leaves its predicate reads, Spark noisy (maxrel 0.5 to * 0.8); passing rows costs the filter scalars. * - `projectPassThrough`: a project over every output leaf. Spark copies the row only when diff --git a/spark/src/main/scala/org/apache/comet/rules/JoinConditionShape.scala b/spark/src/main/scala/org/apache/comet/rules/JoinConditionShape.scala new file mode 100644 index 00000000000..d046b280931 --- /dev/null +++ b/spark/src/main/scala/org/apache/comet/rules/JoinConditionShape.scala @@ -0,0 +1,146 @@ +/* + * 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 scala.collection.mutable + +import org.apache.spark.sql.catalyst.expressions.{Add, AddMonths, Alias, And, Attribute, AttributeSet, Cast, Coalesce, DateAdd, DateAddInterval, DateAddYMInterval, DateSub, Expression, ExprId, GreaterThan, GreaterThanOrEqual, IsNull, LessThan, LessThanOrEqual, Or, Subtract, TimeAdd, TimestampAddYMInterval, TruncDate, TruncTimestamp} +import org.apache.spark.sql.comet.CometExec +import org.apache.spark.sql.execution.SparkPlan +import org.apache.spark.sql.execution.adaptive.QueryStageExec + +/** + * Recognizes a join condition that is a single validity interval: `L <= V < U` with `V` from one + * side, `L` and `U` from the other, and `U` not derived from `L` by a constant offset. + */ +object JoinConditionShape { + + private case class Bound(small: Expression, big: Expression, nullCheck: Option[Expression]) + + private def unwrap(e: Expression): Expression = e match { + case c: Cast => unwrap(c.child) + case t: TruncTimestamp => unwrap(t.timestamp) + case t: TruncDate => unwrap(t.date) + case c: Coalesce if c.children.size > 1 && c.children.tail.forall(_.foldable) => + unwrap(c.children.head) + case other => other + } + + private def offsetBase(e: Expression): Option[Expression] = e match { + case a: Add if a.right.foldable => Some(a.left) + case a: Add if a.left.foldable => Some(a.right) + case s: Subtract if s.right.foldable => Some(s.left) + case d: DateAdd if d.days.foldable => Some(d.startDate) + case d: DateSub if d.days.foldable => Some(d.startDate) + case t: TimeAdd if t.interval.foldable => Some(t.start) + case d: DateAddInterval if d.interval.foldable => Some(d.start) + case d: DateAddYMInterval if d.interval.foldable => Some(d.date) + case t: TimestampAddYMInterval if t.interval.foldable => Some(t.timestamp) + case m: AddMonths if m.numMonths.foldable => Some(m.startDate) + case _ => None + } + + private def strip(e: Expression): Expression = { + val u = unwrap(e) + offsetBase(u).map(strip).getOrElse(u) + } + + private def comparison(e: Expression): Option[(Expression, Expression)] = e match { + case LessThan(a, b) => Some((a, b)) + case LessThanOrEqual(a, b) => Some((a, b)) + case GreaterThan(a, b) => Some((b, a)) + case GreaterThanOrEqual(a, b) => Some((b, a)) + case _ => None + } + + private def bound(e: Expression): Option[Bound] = e match { + case Or(IsNull(x), c) => comparison(c).map { case (s, b) => Bound(s, b, Some(x)) } + case Or(c, IsNull(x)) => comparison(c).map { case (s, b) => Bound(s, b, Some(x)) } + case c => comparison(c).map { case (s, b) => Bound(s, b, None) } + } + + private def conjuncts(e: Expression): Seq[Expression] = e match { + case And(a, b) => conjuncts(a) ++ conjuncts(b) + case other => Seq(other) + } + + private def same(a: Expression, b: Expression): Boolean = unwrap(a).semanticEquals(unwrap(b)) + + /** The aliases defined in `plans` and below, by the id of the attribute each defines. */ + def aliases(plans: Seq[SparkPlan]): Map[ExprId, Expression] = { + val result = mutable.Map[ExprId, Expression]() + def visit(node: SparkPlan): Unit = { + val operators = node match { + case c: CometExec => Seq(c, c.originalPlan) + case other => Seq(other) + } + operators + .flatMap(_.expressions) + .foreach(_.foreach { + case a: Alias => result.getOrElseUpdate(a.exprId, a.child) + case _ => + }) + node match { + case s: QueryStageExec => visit(s.plan) + case _ => node.children.foreach(visit) + } + } + plans.foreach(visit) + result.toMap + } + + private def sources(e: Expression, aliases: Map[ExprId, Expression], depth: Int): Set[ExprId] = + strip(e) match { + case a: Attribute => + aliases.get(a.exprId) match { + case Some(child) if depth < 64 => sources(child, aliases, depth + 1) + case _ => Set(a.exprId) + } + case other => other.references.map(_.exprId).toSet + } + + /** + * Whether `condition` of a join between `left` and `right` outputs is one validity interval + * whose bounds `aliases` does not derive from one another. + */ + def isValidityInterval( + condition: Expression, + left: AttributeSet, + right: AttributeSet, + aliases: Map[ExprId, Expression]): Boolean = { + conjuncts(condition).map(bound) match { + case Seq(Some(first), Some(second)) => + val shapes = Seq((first, second), (second, first)).collect { + case (lower, upper) if same(lower.big, upper.small) && lower.nullCheck.isEmpty => + (upper.small, lower.small, upper.big, upper.nullCheck) + } + shapes.exists { case (v, l, u, nullCheck) => + val (vRefs, lRefs, uRefs) = (v.references, l.references, u.references) + def onOtherSides(side: AttributeSet, other: AttributeSet): Boolean = + vRefs.subsetOf(side) && lRefs.subsetOf(other) && uRefs.subsetOf(other) + vRefs.nonEmpty && lRefs.nonEmpty && uRefs.nonEmpty && + (onOtherSides(left, right) || onOtherSides(right, left)) && + nullCheck.forall(same(_, u)) && + sources(l, aliases, 0).intersect(sources(u, aliases, 0)).isEmpty + } + case _ => false + } + } +} diff --git a/spark/src/test/scala/org/apache/comet/rules/CostBasedEngineChoiceSuite.scala b/spark/src/test/scala/org/apache/comet/rules/CostBasedEngineChoiceSuite.scala index a3a987b7bab..92e350904b0 100644 --- a/spark/src/test/scala/org/apache/comet/rules/CostBasedEngineChoiceSuite.scala +++ b/spark/src/test/scala/org/apache/comet/rules/CostBasedEngineChoiceSuite.scala @@ -943,6 +943,114 @@ class CostBasedEngineChoiceSuite extends CometTestBase { } } + private def withIntervals(f: => Unit): Unit = { + withTempPath { dir => + spark + .range(2000) + .selectExpr( + "cast(id % 97 AS int) AS k", + "timestamp_seconds(id * 3600) AS v", + "id * 3600 AS vs", + "concat('s', cast(id % 7 AS string)) AS s") + .write + .parquet(s"${dir.getCanonicalPath}/ev") + spark + .range(500) + .selectExpr( + "cast(id % 97 AS int) AS k", + "timestamp_seconds(id * 7200) AS l", + "id * 7200 AS ls", + "IF(id % 5 = 0, NULL, timestamp_seconds(id * 7200 + 259200)) AS u", + "timestamp_seconds(id * 7200) AS created_at", + "IF(id % 5 = 0, NULL, timestamp_seconds(id * 7200 + 259200)) AS completed_dt", + "concat('s', cast(id % 5 AS string)) AS s") + .write + .parquet(s"${dir.getCanonicalPath}/dim") + spark.read.parquet(s"${dir.getCanonicalPath}/ev").createOrReplaceTempView("ev") + spark.read.parquet(s"${dir.getCanonicalPath}/dim").createOrReplaceTempView("dim") + withTempView("ev", "dim")(f) + } + } + + private val validityIntervals = Seq( + "e.v >= d.l AND e.v < d.u", + "e.v BETWEEN d.l AND d.u", + "e.v > d.l AND e.v <= COALESCE(d.u, timestamp '9999-12-31')", + "(d.u IS NULL OR e.v < d.u) AND e.v >= d.l", + "to_date(e.v) >= to_date(d.l) AND CAST(e.v AS date) < CAST(d.u AS date)", + "date_trunc('day', e.v) >= d.l AND e.v < date_trunc('day', d.u)") + + private val chargedConditions = Seq( + "SELECT e.k, e.v, d.l FROM ev e JOIN dim d ON e.k = d.k " + + "AND e.vs BETWEEN d.ls - 2592000 AND d.ls", + "SELECT e.k, e.v, d.l FROM ev e JOIN dim d ON e.k = d.k " + + "AND e.v BETWEEN d.l - INTERVAL 30 DAYS AND d.l", + "WITH p AS (SELECT k, l, l + INTERVAL 30 DAYS AS l_end FROM dim) " + + "SELECT e.k, e.v, p.l FROM ev e JOIN p ON e.k = p.k AND e.v >= p.l AND e.v < p.l_end", + "SELECT e.k, e.v, d.l FROM ev e LEFT JOIN dim d ON e.k = d.k AND e.v > d.created_at " + + "AND COALESCE(d.completed_dt, timestamp '5999-12-31') <= e.v", + "SELECT e.k, e.v, d.l FROM ev e JOIN dim d ON e.k = d.k " + + "JOIN dim x ON e.k = x.k AND e.v >= d.l AND e.v < x.u", + "SELECT e.k, e.v, d.l FROM ev e JOIN dim d ON e.k = d.k AND e.v >= d.l", + "SELECT e.k, e.v, d.l FROM ev e JOIN dim d ON e.k = d.k AND e.v >= d.l AND e.v < d.u " + + "AND e.s <> d.s") + + private def conditionTerms(query: String): Seq[Seq[Term]] = { + val plan = runUnordered(sql(query)) + val joins = nodes(plan).collect { + case j: CometSortMergeJoinExec + if j.originalPlan.asInstanceOf[SortMergeJoinExec].condition.isDefined => + j + case j: SortMergeJoinExec if j.condition.isDefined => j + } + assert(joins.nonEmpty, s"plan:\n$plan") + joins.map(j => model.terms(j, Engine.Comet).filter(_.costClass == SmjCondition)) + } + + for (aqe <- Seq("false", "true")) { + test( + s"a validity interval adds no smjCondition, a band or another condition does (AQE=$aqe)") { + withIntervals { + withAqe( + aqe, + flag -> "false", + SQLConf.AUTO_BROADCASTJOIN_THRESHOLD.key -> "-1", + SQLConf.ADAPTIVE_AUTO_BROADCASTJOIN_THRESHOLD.key -> "-1") { + for (condition <- validityIntervals; joinType <- Seq("JOIN", "LEFT JOIN")) { + val query = + s"SELECT e.k, e.v, d.l FROM ev e $joinType dim d ON e.k = d.k AND $condition" + assert(conditionTerms(query).forall(_.isEmpty), query) + } + for (query <- chargedConditions) { + assert(conditionTerms(query).exists(_.nonEmpty), query) + } + } + } + } + + test(s"a validity interval join stays native, a band join runs in Spark (AQE=$aqe)") { + withIntervals { + withAqe( + aqe, + SQLConf.AUTO_BROADCASTJOIN_THRESHOLD.key -> "-1", + SQLConf.ADAPTIVE_AUTO_BROADCASTJOIN_THRESHOLD.key -> "-1") { + def joins(plan: SparkPlan): (Int, Int) = + ( + count(plan) { case j: CometSortMergeJoinExec => j }, + count(plan) { case j: SortMergeJoinExec => j }) + val interval = "SELECT e.k, e.v, d.l FROM ev e LEFT JOIN dim d ON e.k = d.k " + + "AND e.v >= d.l AND e.v < COALESCE(d.u, timestamp '9999-12-31')" + val (intervalOff, intervalOn) = offAndOn(run(interval)) + assert(joins(intervalOff) == (1, 0), s"plan:\n$intervalOff") + assert(joins(intervalOn) == (1, 0), s"plan:\n$intervalOn") + val (bandOff, bandOn) = offAndOn(run(chargedConditions.head)) + assert(joins(bandOff) == (1, 0), s"plan:\n$bandOff") + assert(joins(bandOn) == (0, 1), s"plan:\n$bandOn") + } + } + } + } + test("a window costs its line and the classes of its functions") { withTables { withSQLConf(SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "false", flag -> "false") { From 61590a84a7aca79762d87fb81f60e62c535ed821 Mon Sep 17 00:00:00 2001 From: msaf Date: Tue, 6 Oct 2026 12:29:15 +0100 Subject: [PATCH 03/13] fix: keep native operator protobufs out of Java-serialized plans Every CometNativeExec held its native Operator (the protobuf of its whole subtree, with QueryContext SQL text per expression) and the block's serialized plan as ordinary fields. Any task closure that captures a plan node, such as Spark's SortExec above a Comet subtree in a dynamic-partition write, Java-serialized all of them, growing quadratically with plan depth and query text: 352.7 MiB task binaries for jms_orders against 1.15 MiB on vanilla Spark. nativeOp and serializedPlanOpt are now @transient constructor fields. They stay in the product, so makeCopy, copy, transform, withNewChildren, equality and canonicalization keep them on the driver; executors never read them, since CometExecRDD, the native shuffle spec and the native write paths carry the plan bytes they need explicitly. Co-Authored-By: Claude Opus 5.5 (1M context) --- .../sql/comet/CometDeltaNativeScanExec.scala | 4 +- .../sql/comet/CometCsvNativeScanExec.scala | 4 +- .../comet/CometIcebergNativeScanExec.scala | 4 +- .../sql/comet/CometIcebergWriteExec.scala | 2 +- .../spark/sql/comet/CometNativeScanExec.scala | 4 +- .../sql/comet/CometNativeWriteExec.scala | 2 +- .../spark/sql/comet/CometSampleExec.scala | 4 +- .../spark/sql/comet/CometWindowExec.scala | 4 +- .../sql/comet/CometWindowGroupLimitExec.scala | 4 +- .../spark/sql/comet/CometWriteFilesExec.scala | 2 +- .../apache/spark/sql/comet/operators.scala | 54 +++--- .../comet/exec/CometTaskBinarySizeSuite.scala | 159 +++++++++++++++++- 12 files changed, 203 insertions(+), 44 deletions(-) diff --git a/contrib/delta-spark/src/main/scala/org/apache/spark/sql/comet/CometDeltaNativeScanExec.scala b/contrib/delta-spark/src/main/scala/org/apache/spark/sql/comet/CometDeltaNativeScanExec.scala index f8a922f8138..0ac30149005 100644 --- a/contrib/delta-spark/src/main/scala/org/apache/spark/sql/comet/CometDeltaNativeScanExec.scala +++ b/contrib/delta-spark/src/main/scala/org/apache/spark/sql/comet/CometDeltaNativeScanExec.scala @@ -47,14 +47,14 @@ import org.apache.comet.serde.OperatorOuterClass.Operator * `TreeNode.makeCopy` on MERGE re-planning (the CometIcebergNativeScanExec lesson). */ case class CometDeltaNativeScanExec( - override val nativeOp: Operator, + @transient override val nativeOp: Operator, override val output: Seq[Attribute], requiredSchema: StructType, runtimeFilters: Seq[Expression], dataFilters: Seq[Expression], @transient relation: HadoopFsRelation, originalPlan: FileSourceScanExec, - override val serializedPlanOpt: SerializedPlan, + @transient override val serializedPlanOpt: SerializedPlan, sourceKey: String) extends CometLeafExec with CometScanWithPlanData { diff --git a/spark/src/main/scala/org/apache/spark/sql/comet/CometCsvNativeScanExec.scala b/spark/src/main/scala/org/apache/spark/sql/comet/CometCsvNativeScanExec.scala index dc92648b8d5..170f48389d8 100644 --- a/spark/src/main/scala/org/apache/spark/sql/comet/CometCsvNativeScanExec.scala +++ b/spark/src/main/scala/org/apache/spark/sql/comet/CometCsvNativeScanExec.scala @@ -41,10 +41,10 @@ import org.apache.comet.serde.operator.{partition2Proto, schema2Proto} * Native CSV scan operator that delegates file reading to datafusion. */ case class CometCsvNativeScanExec( - override val nativeOp: Operator, + @transient override val nativeOp: Operator, override val output: Seq[Attribute], @transient override val originalPlan: BatchScanExec, - override val serializedPlanOpt: SerializedPlan) + @transient override val serializedPlanOpt: SerializedPlan) extends CometLeafExec { override val supportsColumnar: Boolean = true diff --git a/spark/src/main/scala/org/apache/spark/sql/comet/CometIcebergNativeScanExec.scala b/spark/src/main/scala/org/apache/spark/sql/comet/CometIcebergNativeScanExec.scala index 5dfea10809e..97de84dee7f 100644 --- a/spark/src/main/scala/org/apache/spark/sql/comet/CometIcebergNativeScanExec.scala +++ b/spark/src/main/scala/org/apache/spark/sql/comet/CometIcebergNativeScanExec.scala @@ -55,11 +55,11 @@ import org.apache.comet.serde.operator.CometIcebergNativeScan * `PlanDataInjector.findAllPlanData` before `commonData` is read. */ case class CometIcebergNativeScanExec( - override val nativeOp: Operator, + @transient override val nativeOp: Operator, override val output: Seq[Attribute], runtimeFilters: Seq[Expression], @transient override val originalPlan: BatchScanExec, - override val serializedPlanOpt: SerializedPlan, + @transient override val serializedPlanOpt: SerializedPlan, metadataLocation: String, scanHashCode: Int, @transient nativeIcebergScanMetadata: CometIcebergNativeScanMetadata) diff --git a/spark/src/main/scala/org/apache/spark/sql/comet/CometIcebergWriteExec.scala b/spark/src/main/scala/org/apache/spark/sql/comet/CometIcebergWriteExec.scala index 00a5047625f..870cdd7ccf5 100644 --- a/spark/src/main/scala/org/apache/spark/sql/comet/CometIcebergWriteExec.scala +++ b/spark/src/main/scala/org/apache/spark/sql/comet/CometIcebergWriteExec.scala @@ -63,7 +63,7 @@ import org.apache.comet.serde.OperatorOuterClass.Operator * with this spec id; required because iceberg-rust's `DataFile` is spec-agnostic at the wire. */ case class CometIcebergWriteExec( - nativeOp: Operator, + @transient nativeOp: Operator, child: SparkPlan, @transient batchWrite: BatchWrite, @transient table: AnyRef, diff --git a/spark/src/main/scala/org/apache/spark/sql/comet/CometNativeScanExec.scala b/spark/src/main/scala/org/apache/spark/sql/comet/CometNativeScanExec.scala index b1ab7f4a014..53eb85ab7bd 100644 --- a/spark/src/main/scala/org/apache/spark/sql/comet/CometNativeScanExec.scala +++ b/spark/src/main/scala/org/apache/spark/sql/comet/CometNativeScanExec.scala @@ -59,7 +59,7 @@ import org.apache.comet.shims.ShimFileFormat * each executor task receives only its partition's file list rather than all files. */ case class CometNativeScanExec( - override val nativeOp: Operator, + @transient override val nativeOp: Operator, @transient relation: HadoopFsRelation, override val output: Seq[Attribute], requiredSchema: StructType, @@ -70,7 +70,7 @@ case class CometNativeScanExec( tableIdentifier: Option[TableIdentifier], disableBucketedScan: Boolean = false, originalPlan: FileSourceScanExec, - override val serializedPlanOpt: SerializedPlan, + @transient override val serializedPlanOpt: SerializedPlan, @transient scan: CometScanExec, // Lazy access to file partitions without serializing with plan sourceKey: String) // Key for PlanDataInjector to match common+partition data at runtime extends CometLeafExec diff --git a/spark/src/main/scala/org/apache/spark/sql/comet/CometNativeWriteExec.scala b/spark/src/main/scala/org/apache/spark/sql/comet/CometNativeWriteExec.scala index f0d10b17667..ab71be00162 100644 --- a/spark/src/main/scala/org/apache/spark/sql/comet/CometNativeWriteExec.scala +++ b/spark/src/main/scala/org/apache/spark/sql/comet/CometNativeWriteExec.scala @@ -68,7 +68,7 @@ import org.apache.comet.serde.OperatorOuterClass.Operator * Unique identifier for this write job */ case class CometNativeWriteExec( - nativeOp: Operator, + @transient nativeOp: Operator, child: SparkPlan, outputPath: String, mode: SaveMode, diff --git a/spark/src/main/scala/org/apache/spark/sql/comet/CometSampleExec.scala b/spark/src/main/scala/org/apache/spark/sql/comet/CometSampleExec.scala index dcdae2c779d..e49d9c8f6d7 100644 --- a/spark/src/main/scala/org/apache/spark/sql/comet/CometSampleExec.scala +++ b/spark/src/main/scala/org/apache/spark/sql/comet/CometSampleExec.scala @@ -84,13 +84,13 @@ object CometSampleExec extends CometOperatorSerde[SampleExec] { * order. */ case class CometSampleExec( - override val nativeOp: Operator, + @transient override val nativeOp: Operator, override val originalPlan: SparkPlan, lowerBound: Double, upperBound: Double, seed: Long, child: SparkPlan, - override val serializedPlanOpt: SerializedPlan) + @transient override val serializedPlanOpt: SerializedPlan) extends CometUnaryExec { override def output: Seq[Attribute] = child.output diff --git a/spark/src/main/scala/org/apache/spark/sql/comet/CometWindowExec.scala b/spark/src/main/scala/org/apache/spark/sql/comet/CometWindowExec.scala index 4d4abf3d2e0..410402d114e 100644 --- a/spark/src/main/scala/org/apache/spark/sql/comet/CometWindowExec.scala +++ b/spark/src/main/scala/org/apache/spark/sql/comet/CometWindowExec.scala @@ -641,14 +641,14 @@ object CometWindowExec extends CometOperatorSerde[WindowExec] { * executions separated by a Comet shuffle exchange. */ case class CometWindowExec( - override val nativeOp: Operator, + @transient override val nativeOp: Operator, override val originalPlan: SparkPlan, override val output: Seq[Attribute], windowExpression: Seq[NamedExpression], partitionSpec: Seq[Expression], orderSpec: Seq[SortOrder], child: SparkPlan, - override val serializedPlanOpt: SerializedPlan) + @transient override val serializedPlanOpt: SerializedPlan) extends CometUnaryExec { override def nodeName: String = "CometWindowExec" diff --git a/spark/src/main/scala/org/apache/spark/sql/comet/CometWindowGroupLimitExec.scala b/spark/src/main/scala/org/apache/spark/sql/comet/CometWindowGroupLimitExec.scala index 265fe4d8d51..9925c77f058 100644 --- a/spark/src/main/scala/org/apache/spark/sql/comet/CometWindowGroupLimitExec.scala +++ b/spark/src/main/scala/org/apache/spark/sql/comet/CometWindowGroupLimitExec.scala @@ -146,7 +146,7 @@ object CometWindowGroupLimitExec extends CometOperatorSerde[SparkPlan] { * would otherwise show identical labels for both). */ case class CometWindowGroupLimitExec( - override val nativeOp: Operator, + @transient override val nativeOp: Operator, override val originalPlan: SparkPlan, override val output: Seq[Attribute], partitionSpec: Seq[Expression], @@ -155,7 +155,7 @@ case class CometWindowGroupLimitExec( limit: Int, mode: String, child: SparkPlan, - override val serializedPlanOpt: SerializedPlan) + @transient override val serializedPlanOpt: SerializedPlan) extends CometUnaryExec { override def nodeName: String = "CometWindowGroupLimitExec" diff --git a/spark/src/main/scala/org/apache/spark/sql/comet/CometWriteFilesExec.scala b/spark/src/main/scala/org/apache/spark/sql/comet/CometWriteFilesExec.scala index 10b0b9db15b..24949aa79e5 100644 --- a/spark/src/main/scala/org/apache/spark/sql/comet/CometWriteFilesExec.scala +++ b/spark/src/main/scala/org/apache/spark/sql/comet/CometWriteFilesExec.scala @@ -80,7 +80,7 @@ import org.apache.comet.shims.ShimCometWriteFilesExec * The Comet native operator producing the batches to write. */ case class CometWriteFilesExec( - nativeOp: Operator, + @transient nativeOp: Operator, override val originalPlan: SparkPlan, child: SparkPlan) extends CometNativeExec 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 99a19be31eb..0718697a794 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 @@ -1384,12 +1384,12 @@ object CometProjectExec extends CometOperatorSerde[ProjectExec] { } case class CometProjectExec( - override val nativeOp: Operator, + @transient override val nativeOp: Operator, override val originalPlan: SparkPlan, override val output: Seq[Attribute], projectList: Seq[NamedExpression], child: SparkPlan, - override val serializedPlanOpt: SerializedPlan) + @transient override val serializedPlanOpt: SerializedPlan) extends CometUnaryExec with PartitioningPreservingUnaryExecNode { override def producedAttributes: AttributeSet = outputSet @@ -1443,12 +1443,12 @@ object CometFilterExec extends CometOperatorSerde[FilterExec] { } case class CometFilterExec( - override val nativeOp: Operator, + @transient override val nativeOp: Operator, override val originalPlan: SparkPlan, override val output: Seq[Attribute], condition: Expression, child: SparkPlan, - override val serializedPlanOpt: SerializedPlan) + @transient override val serializedPlanOpt: SerializedPlan) extends CometUnaryExec { override def outputPartitioning: Partitioning = child.outputPartitioning @@ -1523,13 +1523,13 @@ object CometSortExec extends CometOperatorSerde[SortExec] { } case class CometSortExec( - override val nativeOp: Operator, + @transient override val nativeOp: Operator, override val originalPlan: SparkPlan, override val output: Seq[Attribute], override val outputOrdering: Seq[SortOrder], sortOrder: Seq[SortOrder], child: SparkPlan, - override val serializedPlanOpt: SerializedPlan) + @transient override val serializedPlanOpt: SerializedPlan) extends CometUnaryExec { override def outputPartitioning: Partitioning = child.outputPartitioning @@ -1589,11 +1589,11 @@ object CometLocalLimitExec extends CometOperatorSerde[LocalLimitExec] { } case class CometLocalLimitExec( - override val nativeOp: Operator, + @transient override val nativeOp: Operator, override val originalPlan: SparkPlan, limit: Int, child: SparkPlan, - override val serializedPlanOpt: SerializedPlan) + @transient override val serializedPlanOpt: SerializedPlan) extends CometUnaryExec { override def output: Seq[Attribute] = child.output @@ -1650,12 +1650,12 @@ object CometGlobalLimitExec extends CometOperatorSerde[GlobalLimitExec] { } case class CometGlobalLimitExec( - override val nativeOp: Operator, + @transient override val nativeOp: Operator, override val originalPlan: SparkPlan, limit: Int, offset: Int, child: SparkPlan, - override val serializedPlanOpt: SerializedPlan) + @transient override val serializedPlanOpt: SerializedPlan) extends CometUnaryExec { override def output: Seq[Attribute] = child.output @@ -1714,12 +1714,12 @@ object CometExpandExec extends CometOperatorSerde[ExpandExec] { } case class CometExpandExec( - override val nativeOp: Operator, + @transient override val nativeOp: Operator, override val originalPlan: SparkPlan, override val output: Seq[Attribute], projections: Seq[Seq[Expression]], child: SparkPlan, - override val serializedPlanOpt: SerializedPlan) + @transient override val serializedPlanOpt: SerializedPlan) extends CometUnaryExec { override def outputPartitioning: Partitioning = UnknownPartitioning(0) @@ -1823,14 +1823,14 @@ object CometExplodeExec extends CometOperatorSerde[GenerateExec] { } case class CometExplodeExec( - override val nativeOp: Operator, + @transient override val nativeOp: Operator, override val originalPlan: SparkPlan, override val output: Seq[Attribute], generator: Generator, generatorOutput: Seq[Attribute], outer: Boolean, child: SparkPlan, - override val serializedPlanOpt: SerializedPlan) + @transient override val serializedPlanOpt: SerializedPlan) extends CometUnaryExec { override def outputPartitioning: Partitioning = child.outputPartitioning @@ -2427,7 +2427,7 @@ object CometObjectHashAggregateExec } case class CometHashAggregateExec( - override val nativeOp: Operator, + @transient override val nativeOp: Operator, override val originalPlan: SparkPlan, override val output: Seq[Attribute], groupingExpressions: Seq[NamedExpression], @@ -2436,7 +2436,7 @@ case class CometHashAggregateExec( resultExpressions: Seq[NamedExpression], input: Seq[Attribute], child: SparkPlan, - override val serializedPlanOpt: SerializedPlan) + @transient override val serializedPlanOpt: SerializedPlan) extends CometUnaryExec with PartitioningPreservingUnaryExecNode { @@ -2601,7 +2601,7 @@ trait CometHashJoin { } case class CometBroadcastNestedLoopJoinExec( - override val nativeOp: Operator, + @transient override val nativeOp: Operator, override val originalPlan: SparkPlan, override val output: Seq[Attribute], override val outputOrdering: Seq[SortOrder], @@ -2610,7 +2610,7 @@ case class CometBroadcastNestedLoopJoinExec( buildSide: BuildSide, override val left: SparkPlan, override val right: SparkPlan, - override val serializedPlanOpt: SerializedPlan) + @transient override val serializedPlanOpt: SerializedPlan) extends CometBinaryExec { // Mirror Spark's BroadcastNestedLoopJoinExec: output partitioning derives from the streamed @@ -2812,7 +2812,7 @@ object CometHashJoinExec extends CometOperatorSerde[HashJoin] with CometHashJoin } case class CometHashJoinExec( - override val nativeOp: Operator, + @transient override val nativeOp: Operator, override val originalPlan: SparkPlan, override val output: Seq[Attribute], override val outputOrdering: Seq[SortOrder], @@ -2823,7 +2823,7 @@ case class CometHashJoinExec( buildSide: BuildSide, override val left: SparkPlan, override val right: SparkPlan, - override val serializedPlanOpt: SerializedPlan) + @transient override val serializedPlanOpt: SerializedPlan) extends CometBinaryExec { override def outputPartitioning: Partitioning = joinType match { @@ -2874,7 +2874,7 @@ case class CometHashJoinExec( } case class CometBroadcastHashJoinExec( - override val nativeOp: Operator, + @transient override val nativeOp: Operator, override val originalPlan: SparkPlan, override val output: Seq[Attribute], override val outputOrdering: Seq[SortOrder], @@ -2886,7 +2886,7 @@ case class CometBroadcastHashJoinExec( isNullAwareAntiJoin: Boolean, override val left: SparkPlan, override val right: SparkPlan, - override val serializedPlanOpt: SerializedPlan) + @transient override val serializedPlanOpt: SerializedPlan) extends CometBinaryExec { // The following logic of `outputPartitioning` is copied from Spark `BroadcastHashJoinExec`. @@ -3166,7 +3166,7 @@ object CometSortMergeJoinExec extends CometOperatorSerde[SortMergeJoinExec] { } case class CometSortMergeJoinExec( - override val nativeOp: Operator, + @transient override val nativeOp: Operator, override val originalPlan: SparkPlan, override val output: Seq[Attribute], override val outputOrdering: Seq[SortOrder], @@ -3176,7 +3176,7 @@ case class CometSortMergeJoinExec( condition: Option[Expression], override val left: SparkPlan, override val right: SparkPlan, - override val serializedPlanOpt: SerializedPlan) + @transient override val serializedPlanOpt: SerializedPlan) extends CometBinaryExec { override def outputPartitioning: Partitioning = joinType match { @@ -3225,7 +3225,9 @@ object CometScanWrapper extends CometSink[SparkPlan] { } } -case class CometScanWrapper(override val nativeOp: Operator, override val originalPlan: SparkPlan) +case class CometScanWrapper( + @transient override val nativeOp: Operator, + override val originalPlan: SparkPlan) extends CometNativeExec with LeafExecNode { override val serializedPlanOpt: SerializedPlan = SerializedPlan(None) @@ -3241,7 +3243,7 @@ case class CometScanWrapper(override val nativeOp: Operator, override val origin * This is very similar to `CometScanWrapper` above except it has child. */ case class CometSinkPlaceHolder( - override val nativeOp: Operator, // Must be a Scan + @transient override val nativeOp: Operator, // Must be a Scan override val originalPlan: SparkPlan, child: SparkPlan) extends CometUnaryExec { diff --git a/spark/src/test/scala/org/apache/comet/exec/CometTaskBinarySizeSuite.scala b/spark/src/test/scala/org/apache/comet/exec/CometTaskBinarySizeSuite.scala index cbcb80f32c7..53158f378a2 100644 --- a/spark/src/test/scala/org/apache/comet/exec/CometTaskBinarySizeSuite.scala +++ b/spark/src/test/scala/org/apache/comet/exec/CometTaskBinarySizeSuite.scala @@ -24,6 +24,9 @@ import scala.collection.mutable import org.apache.spark.{ShuffleDependency, SparkEnv, TaskContext} import org.apache.spark.rdd.RDD import org.apache.spark.sql.{CometTestBase, DataFrame} +import org.apache.spark.sql.comet.CometNativeExec +import org.apache.spark.sql.execution.{CommandResultExec, SortExec, SparkPlan} +import org.apache.spark.sql.execution.datasources.WriteFilesExec import org.apache.spark.sql.internal.SQLConf import org.apache.comet.CometConf @@ -47,6 +50,10 @@ class CometTaskBinarySizeSuite extends CometTestBase { private def stageBinaries(df: DataFrame): Seq[StageBinary] = { if (spark.conf.get(SQLConf.ADAPTIVE_EXECUTION_ENABLED.key).toBoolean) df.collect() + rddBinaries(df.queryExecution.executedPlan.execute()) + } + + private def rddBinaries(root: RDD[_]): Seq[StageBinary] = { val serializer = SparkEnv.get.closureSerializer.newInstance() val func = (_: TaskContext, it: Iterator[_]) => it.size val out = mutable.ArrayBuffer.empty[StageBinary] @@ -66,7 +73,7 @@ class CometTaskBinarySizeSuite extends CometTestBase { case _ => } } - visit(df.queryExecution.executedPlan.execute(), None) + visit(root, None) out.toSeq } @@ -198,4 +205,154 @@ class CometTaskBinarySizeSuite extends CometTestBase { } } } + + private def javaSize(obj: AnyRef): Long = + SparkEnv.get.closureSerializer.newInstance().serialize(obj).limit().toLong + + private case class WriteStage(planBytes: Long, stageBytes: Long, nativeNodes: Int) + + private def writeStage(insert: String): WriteStage = { + val command = spark.sql(insert).queryExecution.executedPlan match { + case c: CommandResultExec => c.commandPhysicalPlan + case p => p + } + val write = collectFirst(command) { case w: WriteFilesExec => w } + .getOrElse(fail(s"no WriteFilesExec in\n$command")) + assert(write.child.exists(_.isInstanceOf[SortExec]), write) + val nativeNodes = collectWithSubqueries(write) { case n: CometNativeExec => n }.size + val stage = rddBinaries(write.child.execute()).head + WriteStage(javaSize(write.child), stage.bytes, nativeNodes) + } + + private def longViewSql(events: String, dim: String, cols: Int): String = { + val derived = (0 until cols) + .map(i => + s"CASE WHEN c$i > ${i * 7} THEN c$i * ${i + 3} WHEN s${i % 4} LIKE 'x$i%' THEN " + + s"length(s${i % 4}) ELSE coalesce(c${(i + 1) % cols}, 0) END AS d$i") + .mkString(",\n ") + val outs = (0 until cols) + .map(i => s"coalesce(b.d$i, 0) + dim.w * ${i + 1} - abs(b.d${(i + 3) % cols}) AS o$i") + .mkString(",\n ") + val finals = (0 until cols) + .map(i => s"CASE WHEN o$i % 3 = 0 THEN o$i ELSE o$i + ${i + 11} END AS f$i") + .mkString(",\n ") + s"""WITH base AS ( + | SELECT k, s0, s1, $derived + | FROM parquet.`$events` + | WHERE c0 IS NOT NULL AND s1 NOT LIKE '%zzz%' + |), joined AS ( + | SELECT b.k, b.s0, b.s1, $outs + | FROM base b JOIN parquet.`$dim` dim ON b.k = dim.k + | WHERE b.s0 NOT LIKE '%qqq%' + |) + |SELECT k, s1, $finals, s0 AS p + |FROM joined + |WHERE k % 11 <> 5""".stripMargin + } + + test("a dynamic partition overwrite does not carry every native operator's plan") { + withTempDir { dir => + val cols = 40 + val events = new java.io.File(dir, "events").getCanonicalPath + spark + .range(5000) + .selectExpr( + Seq("id % 97 AS k") ++ (0 until cols).map(i => s"(id * ${i + 1}) % 1013 AS c$i") ++ + (0 until 4).map(i => s"concat('x', cast(id % ${i + 5} AS string)) AS s$i"): _*) + .repartition(4) + .write + .parquet(events) + val dim = new java.io.File(dir, "dim").getCanonicalPath + spark.range(97).selectExpr("id AS k", "id % 7 AS w").write.parquet(dim) + val viewSql = longViewSql(events, dim, cols) + withView("jms_orders_src") { + withTable("jms_orders") { + spark.sql(s"CREATE VIEW jms_orders_src AS $viewSql") + spark.sql( + "CREATE TABLE jms_orders (k BIGINT, s1 STRING, " + + (0 until cols).map(i => s"f$i BIGINT").mkString(", ") + + ", p STRING) USING parquet PARTITIONED BY (p)") + val insert = + "INSERT OVERWRITE TABLE jms_orders PARTITION (p) SELECT * FROM jms_orders_src" + withSQLConf( + CometConf.COMET_EXEC_SORT_ENABLED.key -> "false", + SQLConf.SHUFFLE_PARTITIONS.key -> "8") { + Seq(false, true).foreach { aqe => + def run(cometEnabled: Boolean): WriteStage = { + var result: WriteStage = null + withSQLConf( + CometConf.COMET_ENABLED.key -> cometEnabled.toString, + SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> aqe.toString) { + result = writeStage(insert) + } + result + } + val vanilla = run(cometEnabled = false) + val expected = spark.table("jms_orders").collect().toSet + val comet = run(cometEnabled = true) + assert(spark.table("jms_orders").collect().toSet == expected) + withClue(s"aqe=$aqe vanilla=$vanilla comet=$comet") { + assert(comet.nativeNodes >= 8) + val bound = 2 * viewSql.length + assert(comet.planBytes < 2 * vanilla.planBytes + bound) + assert(comet.stageBytes < vanilla.stageBytes + bound) + } + } + } + } + } + } + } + + test("native operators keep their plan through copies and drop it when serialized") { + withTempDir { dir => + val path = new java.io.File(dir, "t").getCanonicalPath + spark + .range(1000) + .selectExpr("id % 13 AS k", "id AS v", "cast(id AS string) AS s") + .write + .parquet(path) + val sql = + s"SELECT k, sum(v + 1) AS sv, max(length(s)) AS ms FROM parquet.`$path` " + + "WHERE v % 3 <> 0 GROUP BY k" + withSQLConf(SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "false") { + val df = spark.sql(sql) + checkSparkAnswer(df) + val plan = df.queryExecution.executedPlan + val natives = collect(plan) { case n: CometNativeExec => n } + assert(natives.size >= 4, plan) + assert(natives.exists(_.serializedPlanOpt.isDefined), plan) + natives.foreach(n => assert(n.nativeOp != null, n)) + + def assertKept(copy: SparkPlan): Unit = { + val copies = collect(copy) { case n: CometNativeExec => n } + assert(copies.size == natives.size) + natives.zip(copies).foreach { case (n, c) => + assert(c ne n) + assert(c.nativeOp == n.nativeOp) + assert(c.serializedPlanOpt.plan.map(_.toSeq) == n.serializedPlanOpt.plan.map(_.toSeq)) + } + assert(copy.canonicalized == plan.canonicalized) + assert(copy.sameResult(plan)) + } + assertKept(plan.clone()) + def remake(p: SparkPlan): SparkPlan = + p.makeCopy(p.productIterator.map { + case c: SparkPlan if p.children.exists(_ eq c) => remake(c) + case other => other.asInstanceOf[AnyRef] + }.toArray) + assertKept(remake(plan)) + + val serializer = SparkEnv.get.closureSerializer.newInstance() + val restored = serializer.deserialize[SparkPlan](serializer.serialize(plan)) + assert(restored.output == plan.output) + val restoredNatives = collect(restored) { case n: CometNativeExec => n } + assert(restoredNatives.map(_.getClass) == natives.map(_.getClass)) + restoredNatives.foreach { n => + assert(n.nativeOp == null, n) + assert(n.serializedPlanOpt == null, n) + } + } + } + } } From 2fa930bc690e73be1250d37a7012bf70a479d128 Mon Sep 17 00:00:00 2001 From: msaf Date: Tue, 6 Oct 2026 12:47:26 +0100 Subject: [PATCH 04/13] perf: native groups accumulator for Spark first/last and boolean min/max Spark First/Last without ordering ran through DataFusion's GroupsAccumulatorAdapter (one row Accumulator per group). SparkFirstLast keeps the FirstValue/LastValue state layout (value, is_set) and adds a GroupsAccumulator for boolean, integer, float, decimal, date, timestamp, string and binary inputs, honouring ignore nulls, FILTER, merge, EmitTo::First and convert_to_state. Nested types keep the row accumulator. Spark Min/Max over booleans map to bool_and/bool_or, which have a groups accumulator and the same null semantics. Partial aggregate of 10M rows into 1M groups (bench first_last): first ignore nulls with filter over longs 5.83 s -> 0.27 s, last ignore nulls with filter over strings 7.05 s -> 0.64 s, first over strings 4.41 s -> 0.29 s, max over booleans 4.17 s -> 0.25 s. Co-Authored-By: Claude Opus 5.5 (1M context) --- native/core/src/execution/planner.rs | 24 +- native/spark-expr/Cargo.toml | 4 + native/spark-expr/benches/first_last.rs | 233 +++++ native/spark-expr/src/agg_funcs/first_last.rs | 884 ++++++++++++++++++ native/spark-expr/src/agg_funcs/mod.rs | 2 + .../comet/exec/CometAggregateSuite.scala | 68 ++ 6 files changed, 1208 insertions(+), 7 deletions(-) create mode 100644 native/spark-expr/benches/first_last.rs create mode 100644 native/spark-expr/src/agg_funcs/first_last.rs diff --git a/native/core/src/execution/planner.rs b/native/core/src/execution/planner.rs index 3bc8e5a6915..fb691cdc904 100644 --- a/native/core/src/execution/planner.rs +++ b/native/core/src/execution/planner.rs @@ -61,6 +61,7 @@ use arrow::datatypes::{ }; use arrow::ffi_stream::FFI_ArrowArrayStream; use datafusion::functions_aggregate::bit_and_or_xor::{bit_and_udaf, bit_or_udaf, bit_xor_udaf}; +use datafusion::functions_aggregate::bool_and_or::{bool_and_udaf, bool_or_udaf}; use datafusion::functions_aggregate::count::count_udaf; use datafusion::functions_aggregate::min_max::max_udaf; use datafusion::functions_aggregate::min_max::min_udaf; @@ -73,7 +74,6 @@ use datafusion::{ common::DataFusionError, config::ConfigOptions, execution::FunctionRegistry, - functions_aggregate::first_last::{FirstValue, LastValue}, logical_expr::Operator as DataFusionOperator, physical_expr::{ expressions::{ @@ -155,8 +155,8 @@ use datafusion_comet_spark_expr::{ jvm_udf::JvmScalarUdfExpr, spark_in_list, ApproxPercentile, ArrayInsert, Avg, AvgDecimal, Cast, CheckOverflow, Correlation, Covariance, CreateNamedStruct, DecimalRescaleCheckOverflow, GetArrayStructFields, GetStructField, HllPlusPlus, IfExpr, ListExtract, MaxMinBy, Mode, - NormalizeNaNAndZero, Regr, RegrType, SparkCastOptions, Stddev, SumDecimal, ToJson, - UnboundColumn, Variance, WideDecimalBinaryExpr, WideDecimalOp, + NormalizeNaNAndZero, Regr, RegrType, SparkCastOptions, SparkFirstLast, Stddev, SumDecimal, + ToJson, UnboundColumn, Variance, WideDecimalBinaryExpr, WideDecimalOp, }; use itertools::Itertools; use jni::objects::{Global, JObject}; @@ -2937,8 +2937,13 @@ impl PhysicalPlanner { let child = self.create_expr(expr.child.as_ref().unwrap(), Arc::clone(&schema))?; let datatype = to_arrow_datatype(expr.datatype.as_ref().unwrap()); let child = Arc::new(CastExpr::new(child, datatype.clone(), None)); + let func = if datatype == DataType::Boolean { + bool_and_udaf() + } else { + min_udaf() + }; - AggregateExprBuilder::new(min_udaf(), vec![child]) + AggregateExprBuilder::new(func, vec![child]) .schema(schema) .alias("min") .with_ignore_nulls(false) @@ -2950,8 +2955,13 @@ impl PhysicalPlanner { let child = self.create_expr(expr.child.as_ref().unwrap(), Arc::clone(&schema))?; let datatype = to_arrow_datatype(expr.datatype.as_ref().unwrap()); let child = Arc::new(CastExpr::new(child, datatype.clone(), None)); + let func = if datatype == DataType::Boolean { + bool_or_udaf() + } else { + max_udaf() + }; - AggregateExprBuilder::new(max_udaf(), vec![child]) + AggregateExprBuilder::new(func, vec![child]) .schema(schema) .alias("max") .with_ignore_nulls(false) @@ -3032,7 +3042,7 @@ impl PhysicalPlanner { } AggExprStruct::First(expr) => { let child = self.create_expr(expr.child.as_ref().unwrap(), Arc::clone(&schema))?; - let func = AggregateUDF::new_from_impl(FirstValue::new()); + let func = AggregateUDF::new_from_impl(SparkFirstLast::first()); AggregateExprBuilder::new(Arc::new(func), vec![child]) .schema(schema) @@ -3044,7 +3054,7 @@ impl PhysicalPlanner { } AggExprStruct::Last(expr) => { let child = self.create_expr(expr.child.as_ref().unwrap(), Arc::clone(&schema))?; - let func = AggregateUDF::new_from_impl(LastValue::new()); + let func = AggregateUDF::new_from_impl(SparkFirstLast::last()); AggregateExprBuilder::new(Arc::new(func), vec![child]) .schema(schema) diff --git a/native/spark-expr/Cargo.toml b/native/spark-expr/Cargo.toml index fc216a64560..3e07bdda05b 100644 --- a/native/spark-expr/Cargo.toml +++ b/native/spark-expr/Cargo.toml @@ -95,6 +95,10 @@ harness = false name = "aggregate" harness = false +[[bench]] +name = "first_last" +harness = false + [[bench]] name = "approx_percentile" harness = false diff --git a/native/spark-expr/benches/first_last.rs b/native/spark-expr/benches/first_last.rs new file mode 100644 index 00000000000..6ac5cf1bf0f --- /dev/null +++ b/native/spark-expr/benches/first_last.rs @@ -0,0 +1,233 @@ +// 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. + +use arrow::array::builder::{BooleanBuilder, Int64Builder, StringBuilder}; +use arrow::array::{ArrayRef, RecordBatch}; +use arrow::datatypes::{DataType, Field, Schema}; +use criterion::{criterion_group, criterion_main, Criterion}; +use datafusion::datasource::memory::MemorySourceConfig; +use datafusion::datasource::source::DataSourceExec; +use datafusion::execution::TaskContext; +use datafusion::functions_aggregate::bool_and_or::bool_or_udaf; +use datafusion::functions_aggregate::first_last::{FirstValue, LastValue}; +use datafusion::functions_aggregate::min_max::max_udaf; +use datafusion::logical_expr::AggregateUDF; +use datafusion::physical_expr::aggregate::AggregateExprBuilder; +use datafusion::physical_expr::expressions::Column; +use datafusion::physical_expr::PhysicalExpr; +use datafusion::physical_plan::aggregates::{AggregateExec, AggregateMode, PhysicalGroupBy}; +use datafusion::physical_plan::ExecutionPlan; +use datafusion::prelude::SessionConfig; +use datafusion_comet_spark_expr::SparkFirstLast; +use futures::StreamExt; +use std::sync::Arc; +use std::time::Duration; +use tokio::runtime::Runtime; + +const NUM_ROWS: usize = 10_000_000; +const NUM_GROUPS: u64 = 1_000_000; +const BATCH_SIZE: usize = 8192; + +fn batches() -> Vec { + let schema = Arc::new(Schema::new(vec![ + Field::new("k", DataType::Int64, false), + Field::new("i", DataType::Int64, true), + Field::new("s", DataType::Utf8, true), + Field::new("b", DataType::Boolean, true), + Field::new("f", DataType::Boolean, false), + ])); + let mut state = 0x9E3779B97F4A7C15u64; + let mut next = move || { + state ^= state << 13; + state ^= state >> 7; + state ^= state << 17; + state + }; + let mut out = Vec::new(); + let mut row = 0; + while row < NUM_ROWS { + let n = BATCH_SIZE.min(NUM_ROWS - row); + let mut k = Int64Builder::with_capacity(n); + let mut i = Int64Builder::with_capacity(n); + let mut s = StringBuilder::with_capacity(n, n * 16); + let mut b = BooleanBuilder::with_capacity(n); + let mut f = BooleanBuilder::with_capacity(n); + for _ in 0..n { + let r = next(); + k.append_value((r % NUM_GROUPS) as i64); + let null = (r >> 40) % 5 == 0; + if null { + i.append_null(); + s.append_null(); + b.append_null(); + } else { + i.append_value((r >> 20) as i64); + s.append_value(format!("user-{}", (r >> 24) % 100_000)); + b.append_value((r >> 50) % 2 == 0); + } + f.append_value((r >> 33) % 2 == 0); + } + let columns: Vec = vec![ + Arc::new(k.finish()), + Arc::new(i.finish()), + Arc::new(s.finish()), + Arc::new(b.finish()), + Arc::new(f.finish()), + ]; + out.push(RecordBatch::try_new(Arc::clone(&schema), columns).unwrap()); + row += n; + } + out +} + +fn task_context() -> TaskContext { + let mut config = SessionConfig::new(); + config + .options_mut() + .execution + .skip_partial_aggregation_probe_ratio_threshold = 1.1; + TaskContext::default().with_session_config(config) +} + +async fn run( + partitions: &[Vec], + udaf: Arc, + column: &str, + ignore_nulls: bool, + filtered: bool, +) -> usize { + let schema = partitions[0][0].schema(); + let scan: Arc = Arc::new(DataSourceExec::new(Arc::new( + MemorySourceConfig::try_new(partitions, Arc::clone(&schema), None).unwrap(), + ))); + let index = schema.index_of(column).unwrap(); + let value: Arc = Arc::new(Column::new(column, index)); + let key: Arc = Arc::new(Column::new("k", 0)); + let filter: Arc = Arc::new(Column::new("f", 4)); + let aggr = AggregateExprBuilder::new(udaf, vec![value]) + .schema(Arc::clone(&schema)) + .alias("a") + .with_ignore_nulls(ignore_nulls) + .build() + .unwrap(); + let aggregate = AggregateExec::try_new( + AggregateMode::Partial, + PhysicalGroupBy::new_single(vec![(key, "k".to_string())]), + vec![aggr.into()], + vec![filtered.then_some(filter)], + scan, + schema, + ) + .unwrap(); + let mut stream = aggregate.execute(0, Arc::new(task_context())).unwrap(); + let mut rows = 0; + while let Some(batch) = stream.next().await { + rows += batch.unwrap().num_rows(); + } + rows +} + +fn criterion_benchmark(c: &mut Criterion) { + let partitions = vec![batches()]; + let rt = Runtime::new().unwrap(); + let mut group = c.benchmark_group("first_last_10m_rows_1m_groups"); + + let cases: Vec<(&str, Arc, &str, bool, bool)> = vec![ + ( + "first_ignore_nulls_filter_i64_adapter", + Arc::new(AggregateUDF::new_from_impl(FirstValue::new())), + "i", + true, + true, + ), + ( + "first_ignore_nulls_filter_i64_groups", + Arc::new(AggregateUDF::new_from_impl(SparkFirstLast::first())), + "i", + true, + true, + ), + ( + "last_ignore_nulls_filter_utf8_adapter", + Arc::new(AggregateUDF::new_from_impl(LastValue::new())), + "s", + true, + true, + ), + ( + "last_ignore_nulls_filter_utf8_groups", + Arc::new(AggregateUDF::new_from_impl(SparkFirstLast::last())), + "s", + true, + true, + ), + ( + "first_utf8_adapter", + Arc::new(AggregateUDF::new_from_impl(FirstValue::new())), + "s", + false, + false, + ), + ( + "first_utf8_groups", + Arc::new(AggregateUDF::new_from_impl(SparkFirstLast::first())), + "s", + false, + false, + ), + ("max_boolean_adapter", max_udaf(), "b", false, false), + ("max_boolean_bool_or", bool_or_udaf(), "b", false, false), + ]; + + for (name, udaf, column, ignore_nulls, filtered) in cases { + let rows = rt.block_on(run( + &partitions, + Arc::clone(&udaf), + column, + ignore_nulls, + filtered, + )); + assert!(rows > 0, "{name} produced no rows"); + println!("{name}: {rows} output rows"); + group.bench_function(name, |b| { + b.to_async(&rt).iter(|| { + run( + &partitions, + Arc::clone(&udaf), + column, + ignore_nulls, + filtered, + ) + }) + }); + } + group.finish(); +} + +fn config() -> Criterion { + Criterion::default() + .sample_size(10) + .measurement_time(Duration::from_secs(20)) + .warm_up_time(Duration::from_secs(1)) +} + +criterion_group! { + name = benches; + config = config(); + targets = criterion_benchmark +} +criterion_main!(benches); diff --git a/native/spark-expr/src/agg_funcs/first_last.rs b/native/spark-expr/src/agg_funcs/first_last.rs new file mode 100644 index 00000000000..788f2332349 --- /dev/null +++ b/native/spark-expr/src/agg_funcs/first_last.rs @@ -0,0 +1,884 @@ +// 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. + +use std::marker::PhantomData; +use std::mem::size_of; +use std::sync::Arc; + +use arrow::array::{ + Array, ArrayRef, ArrowPrimitiveType, AsArray, BooleanArray, GenericByteArray, PrimitiveArray, +}; +use arrow::buffer::{BooleanBuffer, Buffer, NullBuffer, OffsetBuffer, ScalarBuffer}; +use arrow::compute::nullif; +use arrow::datatypes::ArrowNativeType; +use arrow::datatypes::{ + BinaryType, ByteArrayType, DataType, Date32Type, Date64Type, Decimal128Type, Decimal256Type, + Field, FieldRef, Float32Type, Float64Type, Int16Type, Int32Type, Int64Type, Int8Type, + LargeBinaryType, LargeUtf8Type, TimeUnit, TimestampMicrosecondType, TimestampMillisecondType, + TimestampNanosecondType, TimestampSecondType, UInt16Type, UInt32Type, UInt64Type, UInt8Type, + Utf8Type, +}; +use datafusion::common::{exec_err, internal_err, Result, ScalarValue}; +use datafusion::functions_aggregate::first_last::{FirstValue, LastValue}; +use datafusion::logical_expr::function::{AccumulatorArgs, StateFieldsArgs}; +use datafusion::logical_expr::{ + Accumulator, AggregateUDFImpl, EmitTo, GroupsAccumulator, Signature, Volatility, +}; + +#[derive(Debug, PartialEq, Eq, Hash)] +pub struct SparkFirstLast { + name: String, + signature: Signature, + is_first: bool, +} + +impl SparkFirstLast { + pub fn first() -> Self { + Self::new("first", true) + } + + pub fn last() -> Self { + Self::new("last", false) + } + + fn new(name: &str, is_first: bool) -> Self { + Self { + name: name.to_string(), + signature: Signature::any(1, Volatility::Immutable), + is_first, + } + } + + fn inner(&self) -> Box { + if self.is_first { + Box::new(FirstValue::new()) + } else { + Box::new(LastValue::new()) + } + } + + pub fn groups_supported_type(data_type: &DataType) -> bool { + matches!( + data_type, + DataType::Boolean + | DataType::Int8 + | DataType::Int16 + | DataType::Int32 + | DataType::Int64 + | DataType::UInt8 + | DataType::UInt16 + | DataType::UInt32 + | DataType::UInt64 + | DataType::Float32 + | DataType::Float64 + | DataType::Decimal128(_, _) + | DataType::Decimal256(_, _) + | DataType::Date32 + | DataType::Date64 + | DataType::Timestamp(_, _) + | DataType::Utf8 + | DataType::LargeUtf8 + | DataType::Binary + | DataType::LargeBinary + ) + } +} + +impl AggregateUDFImpl for SparkFirstLast { + fn name(&self) -> &str { + &self.name + } + + fn signature(&self) -> &Signature { + &self.signature + } + + fn return_type(&self, arg_types: &[DataType]) -> Result { + Ok(arg_types[0].clone()) + } + + fn return_field(&self, arg_fields: &[FieldRef]) -> Result { + Ok(Arc::new( + Field::new(self.name(), arg_fields[0].data_type().clone(), true) + .with_metadata(arg_fields[0].metadata().clone()), + )) + } + + fn accumulator(&self, acc_args: AccumulatorArgs) -> Result> { + self.inner().accumulator(acc_args) + } + + fn state_fields(&self, args: StateFieldsArgs) -> Result> { + self.inner().state_fields(args) + } + + fn default_value(&self, data_type: &DataType) -> Result { + ScalarValue::try_from(data_type) + } + + fn supports_null_handling_clause(&self) -> bool { + true + } + + fn groups_accumulator_supported(&self, args: AccumulatorArgs) -> bool { + args.order_bys.is_empty() + && !args.is_distinct + && Self::groups_supported_type(args.return_field.data_type()) + } + + fn create_groups_accumulator( + &self, + args: AccumulatorArgs, + ) -> Result> { + let data_type = args.return_field.data_type().clone(); + let is_first = self.is_first; + let ignore_nulls = args.ignore_nulls; + + macro_rules! primitive { + ($t:ty) => { + Ok(Box::new(FirstLastGroupsAccumulator::new( + PrimitiveStore::<$t>::new(data_type), + is_first, + ignore_nulls, + ))) + }; + } + macro_rules! bytes { + ($t:ty) => { + Ok(Box::new(FirstLastGroupsAccumulator::new( + BytesStore::<$t>::new(), + is_first, + ignore_nulls, + ))) + }; + } + + match &data_type { + DataType::Boolean => Ok(Box::new(FirstLastGroupsAccumulator::new( + BooleanStore::default(), + is_first, + ignore_nulls, + ))), + DataType::Int8 => primitive!(Int8Type), + DataType::Int16 => primitive!(Int16Type), + DataType::Int32 => primitive!(Int32Type), + DataType::Int64 => primitive!(Int64Type), + DataType::UInt8 => primitive!(UInt8Type), + DataType::UInt16 => primitive!(UInt16Type), + DataType::UInt32 => primitive!(UInt32Type), + DataType::UInt64 => primitive!(UInt64Type), + DataType::Float32 => primitive!(Float32Type), + DataType::Float64 => primitive!(Float64Type), + DataType::Decimal128(_, _) => primitive!(Decimal128Type), + DataType::Decimal256(_, _) => primitive!(Decimal256Type), + DataType::Date32 => primitive!(Date32Type), + DataType::Date64 => primitive!(Date64Type), + DataType::Timestamp(TimeUnit::Second, _) => primitive!(TimestampSecondType), + DataType::Timestamp(TimeUnit::Millisecond, _) => primitive!(TimestampMillisecondType), + DataType::Timestamp(TimeUnit::Microsecond, _) => primitive!(TimestampMicrosecondType), + DataType::Timestamp(TimeUnit::Nanosecond, _) => primitive!(TimestampNanosecondType), + DataType::Utf8 => bytes!(Utf8Type), + DataType::LargeUtf8 => bytes!(LargeUtf8Type), + DataType::Binary => bytes!(BinaryType), + DataType::LargeBinary => bytes!(LargeBinaryType), + other => internal_err!("{} groups accumulator does not support {other}", self.name), + } + } +} + +trait ValueStore: Send + 'static { + type Input: Array + 'static; + + fn downcast(array: &ArrayRef) -> Result<&Self::Input> { + match array.as_any().downcast_ref::() { + Some(input) => Ok(input), + None => internal_err!("unexpected input type {}", array.data_type()), + } + } + + fn resize(&mut self, total_num_groups: usize); + + fn set(&mut self, group: usize, input: &Self::Input, row: usize); + + fn emit(&mut self, emit_to: EmitTo) -> Result; + + fn size(&self) -> usize; +} + +struct PrimitiveStore { + data_type: DataType, + values: Vec, + valid: Vec, +} + +impl PrimitiveStore { + fn new(data_type: DataType) -> Self { + Self { + data_type, + values: Vec::new(), + valid: Vec::new(), + } + } +} + +impl ValueStore for PrimitiveStore { + type Input = PrimitiveArray; + + fn resize(&mut self, total_num_groups: usize) { + if total_num_groups > self.values.len() { + self.values.resize(total_num_groups, T::Native::default()); + self.valid.resize(total_num_groups, false); + } + } + + #[inline] + fn set(&mut self, group: usize, input: &Self::Input, row: usize) { + self.values[group] = input.values()[row]; + self.valid[group] = input.is_valid(row); + } + + fn emit(&mut self, emit_to: EmitTo) -> Result { + let values = emit_to.take_needed(&mut self.values); + let valid = emit_to.take_needed(&mut self.valid); + Ok(Arc::new( + PrimitiveArray::::new(ScalarBuffer::from(values), Some(NullBuffer::from(valid))) + .with_data_type(self.data_type.clone()), + )) + } + + fn size(&self) -> usize { + self.values.capacity() * size_of::() + self.valid.capacity() + } +} + +#[derive(Default)] +struct BooleanStore { + values: Vec, + valid: Vec, +} + +impl ValueStore for BooleanStore { + type Input = BooleanArray; + + fn resize(&mut self, total_num_groups: usize) { + if total_num_groups > self.values.len() { + self.values.resize(total_num_groups, false); + self.valid.resize(total_num_groups, false); + } + } + + #[inline] + fn set(&mut self, group: usize, input: &Self::Input, row: usize) { + self.values[group] = input.values().value(row); + self.valid[group] = input.is_valid(row); + } + + fn emit(&mut self, emit_to: EmitTo) -> Result { + let values = emit_to.take_needed(&mut self.values); + let valid = emit_to.take_needed(&mut self.valid); + Ok(Arc::new(BooleanArray::new( + BooleanBuffer::from(values), + Some(NullBuffer::from(valid)), + ))) + } + + fn size(&self) -> usize { + self.values.capacity() + self.valid.capacity() + } +} + +struct BytesStore { + values: Vec>, + valid: Vec, + heap: usize, + phantom: PhantomData, +} + +impl BytesStore { + fn new() -> Self { + Self { + values: Vec::new(), + valid: Vec::new(), + heap: 0, + phantom: PhantomData, + } + } +} + +impl ValueStore for BytesStore { + type Input = GenericByteArray; + + fn resize(&mut self, total_num_groups: usize) { + if total_num_groups > self.values.len() { + self.values.resize_with(total_num_groups, Vec::new); + self.valid.resize(total_num_groups, false); + } + } + + #[inline] + fn set(&mut self, group: usize, input: &Self::Input, row: usize) { + if input.is_null(row) { + self.valid[group] = false; + return; + } + let bytes: &[u8] = input.value(row).as_ref(); + let slot = &mut self.values[group]; + let before = slot.capacity(); + slot.clear(); + slot.extend_from_slice(bytes); + self.heap += slot.capacity() - before; + self.valid[group] = true; + } + + fn emit(&mut self, emit_to: EmitTo) -> Result { + let values = emit_to.take_needed(&mut self.values); + let valid = emit_to.take_needed(&mut self.valid); + + let mut total = 0usize; + let mut released = 0usize; + for (value, &ok) in values.iter().zip(valid.iter()) { + if ok { + total += value.len(); + } + released += value.capacity(); + } + self.heap -= released; + if T::Offset::from_usize(total).is_none() { + return exec_err!( + "{} values of {total} bytes overflow the offsets of {}", + values.len(), + T::DATA_TYPE + ); + } + + let mut offsets: Vec = Vec::with_capacity(values.len() + 1); + let mut data: Vec = Vec::with_capacity(total); + offsets.push(T::Offset::usize_as(0)); + for (value, &ok) in values.iter().zip(valid.iter()) { + if ok { + data.extend_from_slice(value); + } + offsets.push(T::Offset::usize_as(data.len())); + } + + let offsets = unsafe { OffsetBuffer::new_unchecked(ScalarBuffer::from(offsets)) }; + let array = unsafe { + GenericByteArray::::new_unchecked( + offsets, + Buffer::from_vec(data), + Some(NullBuffer::from(valid)), + ) + }; + Ok(Arc::new(array)) + } + + fn size(&self) -> usize { + self.values.capacity() * size_of::>() + self.heap + self.valid.capacity() + } +} + +struct FirstLastGroupsAccumulator { + store: S, + is_set: Vec, + is_first: bool, + ignore_nulls: bool, +} + +impl FirstLastGroupsAccumulator { + fn new(store: S, is_first: bool, ignore_nulls: bool) -> Self { + Self { + store, + is_set: Vec::new(), + is_first, + ignore_nulls, + } + } + + fn resize(&mut self, total_num_groups: usize) { + if total_num_groups > self.is_set.len() { + self.is_set.resize(total_num_groups, false); + } + self.store.resize(total_num_groups); + } + + fn apply( + &mut self, + input: &S::Input, + group_indices: &[usize], + mask: Option<&BooleanBuffer>, + skip_nulls: bool, + ) { + let skip_nulls = skip_nulls && input.null_count() > 0; + for (row, &group) in group_indices.iter().enumerate() { + if self.is_first && self.is_set[group] { + continue; + } + if let Some(mask) = mask { + if !mask.value(row) { + continue; + } + } + if skip_nulls && input.is_null(row) { + continue; + } + self.store.set(group, input, row); + self.is_set[group] = true; + } + } +} + +fn true_mask(array: &BooleanArray) -> BooleanBuffer { + match array.nulls() { + Some(nulls) => array.values() & nulls.inner(), + None => array.values().clone(), + } +} + +impl GroupsAccumulator for FirstLastGroupsAccumulator { + fn update_batch( + &mut self, + values: &[ArrayRef], + group_indices: &[usize], + opt_filter: Option<&BooleanArray>, + total_num_groups: usize, + ) -> Result<()> { + self.resize(total_num_groups); + let input = S::downcast(&values[0])?; + let mask = opt_filter.map(true_mask); + self.apply(input, group_indices, mask.as_ref(), self.ignore_nulls); + Ok(()) + } + + fn merge_batch( + &mut self, + values: &[ArrayRef], + group_indices: &[usize], + total_num_groups: usize, + ) -> Result<()> { + self.resize(total_num_groups); + let input = S::downcast(&values[0])?; + let mask = true_mask(values[1].as_boolean()); + self.apply(input, group_indices, Some(&mask), false); + Ok(()) + } + + fn evaluate(&mut self, emit_to: EmitTo) -> Result { + let values = self.store.emit(emit_to)?; + emit_to.take_needed(&mut self.is_set); + Ok(values) + } + + fn state(&mut self, emit_to: EmitTo) -> Result> { + let values = self.store.emit(emit_to)?; + let is_set = emit_to.take_needed(&mut self.is_set); + Ok(vec![ + values, + Arc::new(BooleanArray::new(BooleanBuffer::from(is_set), None)), + ]) + } + + fn convert_to_state( + &self, + values: &[ArrayRef], + opt_filter: Option<&BooleanArray>, + ) -> Result> { + let input = &values[0]; + let mut is_set = match opt_filter { + Some(filter) => true_mask(filter), + None => BooleanBuffer::new_set(input.len()), + }; + if self.ignore_nulls { + if let Some(nulls) = input.logical_nulls() { + is_set = &is_set & nulls.inner(); + } + } + let value = if is_set.count_set_bits() == input.len() { + Arc::clone(input) + } else { + nullif(input.as_ref(), &BooleanArray::new(!&is_set, None))? + }; + Ok(vec![value, Arc::new(BooleanArray::new(is_set, None))]) + } + + fn size(&self) -> usize { + self.store.size() + self.is_set.capacity() + } +} + +#[cfg(test)] +mod tests { + use super::*; + use arrow::array::{Int32Array, Int64Array, StringArray}; + use arrow::datatypes::Schema; + use datafusion::physical_expr::expressions::col; + use datafusion::physical_expr::PhysicalExpr; + + fn accumulator( + udaf: &SparkFirstLast, + data_type: DataType, + ignore_nulls: bool, + ) -> Box { + let schema = Schema::new(vec![Field::new("a", data_type.clone(), true)]); + let expr: Arc = col("a", &schema).unwrap(); + let return_field: FieldRef = Arc::new(Field::new("f", data_type.clone(), true)); + let expr_field: FieldRef = Arc::new(Field::new("a", data_type, true)); + let args = AccumulatorArgs { + return_field, + schema: &schema, + expr_fields: &[expr_field], + ignore_nulls, + order_bys: &[], + is_reversed: false, + name: "f", + is_distinct: false, + exprs: &[expr], + }; + assert!(udaf.groups_accumulator_supported(args.clone())); + udaf.create_groups_accumulator(args).unwrap() + } + + fn ints(values: Vec>) -> ArrayRef { + Arc::new(Int32Array::from(values)) + } + + fn strings(values: Vec>) -> ArrayRef { + Arc::new(StringArray::from(values)) + } + + fn filter(values: Vec>) -> BooleanArray { + BooleanArray::from(values) + } + + fn as_ints(array: &ArrayRef) -> Vec> { + array.as_primitive::().iter().collect() + } + + fn as_strings(array: &ArrayRef) -> Vec> { + array + .as_string::() + .iter() + .map(|v| v.map(|s| s.to_string())) + .collect() + } + + fn as_bools(array: &ArrayRef) -> Vec> { + array.as_boolean().iter().collect() + } + + fn reference( + is_first: bool, + ignore_nulls: bool, + values: &[Option], + groups: &[usize], + mask: &[Option], + num_groups: usize, + ) -> Vec> { + let mut out: Vec<(Option, bool)> = vec![(None, false); num_groups]; + for ((v, &g), m) in values.iter().zip(groups).zip(mask) { + if *m != Some(true) || (ignore_nulls && v.is_none()) { + continue; + } + if is_first && out[g].1 { + continue; + } + out[g] = (*v, true); + } + out.into_iter().map(|(v, _)| v).collect() + } + + #[test] + fn first_and_last_respect_nulls() { + let values = ints(vec![None, Some(1), Some(2), None, None, Some(5)]); + let groups = [0, 0, 1, 1, 2, 2]; + let cases = [ + (true, false, vec![None, Some(2), None]), + (true, true, vec![Some(1), Some(2), Some(5)]), + (false, false, vec![Some(1), None, Some(5)]), + (false, true, vec![Some(1), Some(2), Some(5)]), + ]; + for (is_first, ignore_nulls, expected) in cases { + let udaf = if is_first { + SparkFirstLast::first() + } else { + SparkFirstLast::last() + }; + let mut acc = accumulator(&udaf, DataType::Int32, ignore_nulls); + acc.update_batch(&[Arc::clone(&values)], &groups, None, 4) + .unwrap(); + let out = acc.evaluate(EmitTo::All).unwrap(); + let mut expected = expected; + expected.push(None); + assert_eq!( + as_ints(&out), + expected, + "first={is_first} ignore={ignore_nulls}" + ); + } + } + + #[test] + fn filter_skips_rows_including_null_filter_values() { + let values = ints(vec![Some(1), Some(2), Some(3), Some(4)]); + let groups = [0, 0, 0, 0]; + let mask = filter(vec![Some(false), None, Some(true), Some(false)]); + let mut first = accumulator(&SparkFirstLast::first(), DataType::Int32, true); + first + .update_batch(&[Arc::clone(&values)], &groups, Some(&mask), 1) + .unwrap(); + assert_eq!( + as_ints(&first.evaluate(EmitTo::All).unwrap()), + vec![Some(3)] + ); + let mut last = accumulator(&SparkFirstLast::last(), DataType::Int32, false); + last.update_batch(&[values], &groups, Some(&mask), 1) + .unwrap(); + assert_eq!(as_ints(&last.evaluate(EmitTo::All).unwrap()), vec![Some(3)]); + } + + #[test] + fn state_and_merge_follow_spark_buffers() { + for is_first in [true, false] { + for ignore_nulls in [true, false] { + let udaf = if is_first { + SparkFirstLast::first() + } else { + SparkFirstLast::last() + }; + let mut left = accumulator(&udaf, DataType::Utf8, ignore_nulls); + left.update_batch(&[strings(vec![Some("a"), None, None])], &[0, 1, 1], None, 3) + .unwrap(); + let left_state = left.state(EmitTo::All).unwrap(); + assert_eq!( + as_bools(&left_state[1]), + vec![Some(true), Some(!ignore_nulls), Some(false)] + ); + + let mut right = accumulator(&udaf, DataType::Utf8, ignore_nulls); + right + .update_batch( + &[strings(vec![Some("b"), Some("c"), Some("d")])], + &[0, 1, 2], + None, + 3, + ) + .unwrap(); + let right_state = right.state(EmitTo::All).unwrap(); + + let mut merged = accumulator(&udaf, DataType::Utf8, ignore_nulls); + merged.merge_batch(&left_state, &[0, 1, 2], 3).unwrap(); + merged.merge_batch(&right_state, &[0, 1, 2], 3).unwrap(); + let out = as_strings(&merged.evaluate(EmitTo::All).unwrap()); + let expected: Vec> = match (is_first, ignore_nulls) { + (true, true) => vec![Some("a"), Some("c"), Some("d")], + (true, false) => vec![Some("a"), None, Some("d")], + (false, _) => vec![Some("b"), Some("c"), Some("d")], + } + .into_iter() + .map(|v| v.map(String::from)) + .collect(); + assert_eq!(out, expected, "first={is_first} ignore={ignore_nulls}"); + } + } + } + + #[test] + fn merge_ignores_unset_partial_states() { + let mut acc = accumulator(&SparkFirstLast::last(), DataType::Int32, false); + let state: Vec = vec![ + ints(vec![Some(1), None, Some(7)]), + Arc::new(BooleanArray::from(vec![Some(true), Some(false), None])), + ]; + acc.merge_batch(&state, &[0, 0, 0], 1).unwrap(); + assert_eq!(as_ints(&acc.evaluate(EmitTo::All).unwrap()), vec![Some(1)]); + } + + #[test] + fn emit_first_shifts_remaining_groups() { + for data_type in [DataType::Int32, DataType::Utf8] { + let mut acc = accumulator(&SparkFirstLast::first(), data_type.clone(), true); + let input = match data_type { + DataType::Int32 => ints(vec![Some(0), Some(1), Some(2), None]), + _ => strings(vec![Some("0"), Some("1"), Some("2"), None]), + }; + acc.update_batch(&[input], &[0, 1, 2, 3], None, 4).unwrap(); + let size_before = acc.size(); + let head = acc.state(EmitTo::First(2)).unwrap(); + assert_eq!(head[0].len(), 2); + assert_eq!(as_bools(&head[1]), vec![Some(true), Some(true)]); + assert!(acc.size() <= size_before); + let next = match data_type { + DataType::Int32 => ints(vec![Some(10), Some(11), Some(12)]), + _ => strings(vec![Some("10"), Some("11"), Some("12")]), + }; + acc.update_batch(&[next], &[0, 1, 2], None, 3).unwrap(); + let rest = acc.evaluate(EmitTo::All).unwrap(); + let rendered: Vec> = match data_type { + DataType::Int32 => as_ints(&rest) + .into_iter() + .map(|v| v.map(|v| v.to_string())) + .collect(), + _ => as_strings(&rest), + }; + assert_eq!( + rendered, + vec![ + Some("2".to_string()), + Some("11".to_string()), + Some("12".to_string()) + ] + ); + } + } + + #[test] + fn many_groups_match_reference() { + let num_rows = 200_000usize; + let num_groups = 30_000usize; + let mut seed = 0x2545F4914F6CDD1Du64; + let mut next = || { + seed ^= seed << 13; + seed ^= seed >> 7; + seed ^= seed << 17; + seed + }; + let values: Vec> = (0..num_rows) + .map(|_| { + let r = next(); + (r % 4 != 0).then_some((r >> 8) as i32) + }) + .collect(); + let groups: Vec = (0..num_rows) + .map(|_| (next() % num_groups as u64) as usize) + .collect(); + let mask: Vec> = (0..num_rows) + .map(|_| match next() % 5 { + 0 => None, + 1 => Some(false), + _ => Some(true), + }) + .collect(); + + for is_first in [true, false] { + for ignore_nulls in [true, false] { + let udaf = if is_first { + SparkFirstLast::first() + } else { + SparkFirstLast::last() + }; + let expected = + reference(is_first, ignore_nulls, &values, &groups, &mask, num_groups); + + let mut partials = Vec::new(); + for chunk in 0..4 { + let range = chunk * num_rows / 4..(chunk + 1) * num_rows / 4; + let mut acc = accumulator(&udaf, DataType::Int32, ignore_nulls); + for batch in range.clone().step_by(8192) { + let end = (batch + 8192).min(range.end); + acc.update_batch( + &[ints(values[batch..end].to_vec())], + &groups[batch..end], + Some(&filter(mask[batch..end].to_vec())), + num_groups, + ) + .unwrap(); + } + partials.push(acc.state(EmitTo::All).unwrap()); + } + + let identity: Vec = (0..num_groups).collect(); + let mut merged = accumulator(&udaf, DataType::Int32, ignore_nulls); + for state in &partials { + merged.merge_batch(state, &identity, num_groups).unwrap(); + } + let mut out = as_ints(&merged.state(EmitTo::First(1000)).unwrap()[0]); + out.extend(as_ints(&merged.evaluate(EmitTo::All).unwrap())); + assert_eq!(out, expected, "first={is_first} ignore={ignore_nulls}"); + + let converted = accumulator(&udaf, DataType::Int32, ignore_nulls) + .convert_to_state(&[ints(values.clone())], Some(&filter(mask.clone()))) + .unwrap(); + let mut from_rows = accumulator(&udaf, DataType::Int32, ignore_nulls); + from_rows + .merge_batch(&converted, &groups, num_groups) + .unwrap(); + assert_eq!( + as_ints(&from_rows.evaluate(EmitTo::All).unwrap()), + expected, + "convert_to_state first={is_first} ignore={ignore_nulls}" + ); + } + } + } + + #[test] + fn typed_outputs_keep_their_data_type() { + let decimal = DataType::Decimal128(20, 3); + let mut acc = accumulator(&SparkFirstLast::last(), decimal.clone(), true); + let input: ArrayRef = Arc::new( + PrimitiveArray::::from(vec![Some(1234), None]) + .with_data_type(decimal.clone()), + ); + acc.update_batch(&[input], &[0, 0], None, 1).unwrap(); + let out = acc.evaluate(EmitTo::All).unwrap(); + assert_eq!(out.data_type(), &decimal); + assert_eq!(out.as_primitive::().value(0), 1234); + + let ts = DataType::Timestamp(TimeUnit::Microsecond, Some("UTC".into())); + let mut acc = accumulator(&SparkFirstLast::first(), ts.clone(), false); + let input: ArrayRef = Arc::new( + PrimitiveArray::::from(vec![Some(7), Some(8)]) + .with_data_type(ts.clone()), + ); + acc.update_batch(&[input], &[0, 0], None, 1).unwrap(); + assert_eq!(acc.evaluate(EmitTo::All).unwrap().data_type(), &ts); + + let mut acc = accumulator(&SparkFirstLast::last(), DataType::Boolean, true); + let input: ArrayRef = Arc::new(BooleanArray::from(vec![Some(true), Some(false), None])); + acc.update_batch(&[input], &[0, 0, 1], None, 2).unwrap(); + assert_eq!( + as_bools(&acc.evaluate(EmitTo::All).unwrap()), + vec![Some(false), None] + ); + + let mut acc = accumulator(&SparkFirstLast::first(), DataType::Int64, false); + let input: ArrayRef = Arc::new(Int64Array::from(vec![Some(1), Some(2)])); + acc.update_batch(&[input], &[1, 1], None, 2).unwrap(); + let out = acc.evaluate(EmitTo::All).unwrap(); + assert_eq!( + out.as_primitive::().iter().collect::>(), + vec![None, Some(1)] + ); + } + + #[test] + fn nested_types_keep_the_row_accumulator() { + let list = DataType::List(Arc::new(Field::new_list_field(DataType::Int32, true))); + assert!(!SparkFirstLast::groups_supported_type(&list)); + assert!(SparkFirstLast::groups_supported_type(&DataType::Utf8)); + } + + #[test] + fn size_tracks_string_heap() { + let mut acc = accumulator(&SparkFirstLast::last(), DataType::Utf8, false); + let long = "x".repeat(1000); + acc.update_batch(&[strings(vec![Some(&long)])], &[0], None, 1) + .unwrap(); + assert!(acc.size() >= 1000); + acc.evaluate(EmitTo::All).unwrap(); + assert!(acc.size() < 1000); + } +} diff --git a/native/spark-expr/src/agg_funcs/mod.rs b/native/spark-expr/src/agg_funcs/mod.rs index cfe06565193..ec8cbdfdca0 100644 --- a/native/spark-expr/src/agg_funcs/mod.rs +++ b/native/spark-expr/src/agg_funcs/mod.rs @@ -21,6 +21,7 @@ mod avg_decimal; mod collect; mod correlation; mod covariance; +mod first_last; mod hll_plus_plus; mod hll_plus_plus_const; mod list_agg; @@ -41,6 +42,7 @@ pub use avg_decimal::AvgDecimal; pub use collect::{CometCollectList, CometCollectSet}; pub use correlation::Correlation; pub use covariance::Covariance; +pub use first_last::SparkFirstLast; pub use hll_plus_plus::{hllpp_precision, HllPlusPlus}; pub use list_agg::SparkListAgg; pub use max_min_by::MaxMinBy; 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 b4ef6f3cc8d..dca5f3f05f5 100644 --- a/spark/src/test/scala/org/apache/comet/exec/CometAggregateSuite.scala +++ b/spark/src/test/scala/org/apache/comet/exec/CometAggregateSuite.scala @@ -1542,6 +1542,74 @@ class CometAggregateSuite extends CometTestBase with AdaptiveSparkPlanHelper { } } + private def withFirstLastTable(f: => Unit): Unit = { + withTempDir { dir => + val path = s"${dir.getAbsolutePath}/first_last_groups.parquet" + spark + .range(0, 200000, 1, 4) + .selectExpr( + "id % 20011 AS k", + "CAST(id % 20011 AS INT) * 3 AS c", + "concat('k', id % 20011) AS cs", + "IF(id DIV 20011 = 1, id % 2 = 0, NULL) AS one_b", + "IF(id DIV 20011 = 2, timestamp_seconds(id), NULL) AS one_ts", + "IF(id DIV 20011 = 3, CAST(id AS INT), NULL) AS one_i", + "IF(id DIV 20011 = 4, date_add(DATE'2020-01-01', CAST(id % 1000 AS INT)), NULL) AS one_date", + "IF(id DIV 20011 = 5, concat('s', id), NULL) AS one_s", + "IF(id DIV 20011 = 6, CAST(id AS DOUBLE) / 7, NULL) AS one_f", + "IF(id DIV 20011 = 7, CAST(id AS DECIMAL(20, 3)) / 7, NULL) AS one_d", + "IF(id % 7 = 0, NULL, id % 3 = 0) AS b", + "IF(id % 20011 % 5 = 0, NULL, id % 4 = 0) AS bn") + .write + .parquet(path) + withSQLConf(CometConf.COMET_BATCH_SIZE.key -> "1000") { + spark.read.parquet(path).createOrReplaceTempView("first_last_groups") + f + } + } + } + + private val firstLastQueries = Seq( + """SELECT k, count(DISTINCT cs), count(DISTINCT one_s), first(one_i, true), sum(c), + | max(b), min(bn), first(one_ts, true) FILTER (WHERE k % 2 = 0), + | last(one_d, true) FILTER (WHERE c > 30) + |FROM first_last_groups GROUP BY k""".stripMargin, + """SELECT k, count(DISTINCT one_i) FILTER (WHERE c % 3 = 0), count(DISTINCT cs), + | first(one_ts, true), last(one_date, true), max(one_b), min(one_b) + |FROM first_last_groups GROUP BY k""".stripMargin, + """SELECT k, first(c), last(c), first(one_b, true), last(one_b, true), first(one_ts, true), + | last(one_i, true), first(one_date, true), last(one_f, true), first(one_f, true), + | last(one_d, true), max(b), min(b), max(bn), min(bn) + |FROM first_last_groups GROUP BY k""".stripMargin, + """SELECT k % 3, first(one_d, true) FILTER (WHERE k = 7), + | last(one_i, true) FILTER (WHERE k = 9), first(c) FILTER (WHERE k = 11), + | max(bn) FILTER (WHERE k % 5 = 0), min(b) FILTER (WHERE k < 100) + |FROM first_last_groups GROUP BY k % 3""".stripMargin, + "SELECT max(b), min(b), max(bn) FILTER (WHERE k % 5 = 0), min(one_b) FROM first_last_groups") + + test("first/last with ignore nulls and filters, and boolean min/max, match Spark") { + withFirstLastTable { + firstLastQueries.foreach(q => checkSparkAnswerAndOperator(sql(q))) + } + } + + test("first/last and boolean min/max match Spark when the aggregate spills") { + withFirstLastTable { + withSQLConf( + CometConf.COMET_OFFHEAP_MEMORY_POOL_TYPE.key -> "fair_unified", + CometConf.COMET_OFFHEAP_MEMORY_POOL_FRACTION.key -> "0.0003", + SQLConf.SHUFFLE_PARTITIONS.key -> "2") { + val spills = firstLastQueries.map { q => + val (_, plan) = checkSparkAnswerAndOperator(sql(q)) + stripAQEPlan(plan).collect { case a: CometHashAggregateExec => + a.metrics.get("spill_count").map(_.value).getOrElse(0L) + }.sum + } + assert(spills.count(_ > 0) >= 3, s"spills per query: $spills") + } + } + } + test("first/last with ignore null") { val data = Range(0, 8192).flatMap(n => Seq((n, 1), (n, 2))).toDF("a", "b") withTempDir { dir => From 25f52bfce376ca913bfa74bd4bc5fefddca16757 Mon Sep 17 00:00:00 2001 From: msaf Date: Tue, 6 Oct 2026 13:21:42 +0100 Subject: [PATCH 05/13] test: differential fuzz suite for first/last and boolean min/max aggregates Compares Comet with Spark for first/last (with and without IGNORE NULLS and FILTER), any_value, min/max over booleans, bool_and/bool_or/every/some/any and the FILTER shape of Spark's COUNT(DISTINCT) rewrite, over generated data with fixed null patterns per group. Ordered single-split inputs are compared exactly; many-partition and spilling runs check that each value belongs to its group. The default part runs in CI; -Dcomet.test.aggFuzz.full=true runs the full matrix of types, shapes, groupings and a 200k-group dataset. An ignored test records a pre-existing bug: under memory pressure a single COUNT(DISTINCT) overcounts because the PartialMerge aggregate runs as a DataFusion Partial aggregate and emits groups early. Co-Authored-By: Claude Opus 5.5 (1M context) --- .github/workflows/pr_build_linux.yml | 1 + .github/workflows/pr_build_macos.yml | 1 + .../exec/CometFirstLastBoolAggFuzzSuite.scala | 673 ++++++++++++++++++ 3 files changed, 675 insertions(+) create mode 100644 spark/src/test/scala/org/apache/comet/exec/CometFirstLastBoolAggFuzzSuite.scala diff --git a/.github/workflows/pr_build_linux.yml b/.github/workflows/pr_build_linux.yml index 8e837389659..830f8881645 100644 --- a/.github/workflows/pr_build_linux.yml +++ b/.github/workflows/pr_build_linux.yml @@ -526,6 +526,7 @@ jobs: - name: "exec" value: | org.apache.comet.exec.CometAggregateSuite + org.apache.comet.exec.CometFirstLastBoolAggFuzzSuite org.apache.comet.exec.CometExec3_4PlusSuite org.apache.comet.exec.CometExecSuite org.apache.comet.exec.CometTaskBinarySizeSuite diff --git a/.github/workflows/pr_build_macos.yml b/.github/workflows/pr_build_macos.yml index 2fcd137ded6..20f815854e6 100644 --- a/.github/workflows/pr_build_macos.yml +++ b/.github/workflows/pr_build_macos.yml @@ -174,6 +174,7 @@ jobs: - name: "exec" value: | org.apache.comet.exec.CometAggregateSuite + org.apache.comet.exec.CometFirstLastBoolAggFuzzSuite org.apache.comet.exec.CometExec3_4PlusSuite org.apache.comet.exec.CometExecSuite org.apache.comet.exec.CometTaskBinarySizeSuite diff --git a/spark/src/test/scala/org/apache/comet/exec/CometFirstLastBoolAggFuzzSuite.scala b/spark/src/test/scala/org/apache/comet/exec/CometFirstLastBoolAggFuzzSuite.scala new file mode 100644 index 00000000000..1c881335d67 --- /dev/null +++ b/spark/src/test/scala/org/apache/comet/exec/CometFirstLastBoolAggFuzzSuite.scala @@ -0,0 +1,673 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, + * software distributed under the License is distributed on an + * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY + * KIND, either express or implied. See the License for the + * specific language governing permissions and limitations + * under the License. + */ + +package org.apache.comet.exec + +import java.io.File +import java.nio.file.Files +import java.sql.{Date, Timestamp} +import java.time.{Instant, LocalDate, LocalDateTime, ZoneOffset} + +import scala.collection.mutable +import scala.util.Random + +import org.apache.commons.io.FileUtils +import org.apache.spark.sql.{CometTestBase, Row} +import org.apache.spark.sql.comet.CometHashAggregateExec +import org.apache.spark.sql.execution.SparkPlan +import org.apache.spark.sql.execution.adaptive.AdaptiveSparkPlanHelper +import org.apache.spark.sql.execution.aggregate.BaseAggregateExec +import org.apache.spark.sql.internal.SQLConf +import org.apache.spark.sql.types._ + +import org.apache.comet.CometConf + +/** + * Differential tests of first/last (with and without IGNORE NULLS, FILTER, and the FILTER shape + * produced by Spark's rewrite of several COUNT(DISTINCT ...)), any_value, min/max over booleans + * and bool_and/bool_or/every/some/any, comparing Comet against Spark on generated data. + * + * first/last without ordering depend on the input order. On a single input split read in file + * order the answer is deterministic and compared exactly. Over many partitions the Comet (and + * Spark) value must be one of the values of its group that pass the aggregate's FILTER, or NULL + * when the group has a NULL such value (or, with IGNORE NULLS, only when it has no other). + * + * The default run covers every function on the main types and one case of each execution shape. + * The full matrix (all types, shapes, groupings and a 200k-group dataset) runs only with + * `-Dcomet.test.aggFuzz.full=true`; `-Dcomet.test.aggFuzz.seed=` changes the data seed. + */ +class CometFirstLastBoolAggFuzzSuite extends CometTestBase with AdaptiveSparkPlanHelper { + + private val seed: Long = + sys.props.get("comet.test.aggFuzz.seed").map(_.toLong).getOrElse(20261006L) + private val fullMatrix: Boolean = sys.props.get("comet.test.aggFuzz.full").contains("true") + + private case class Col(name: String, dt: DataType, main: Boolean) + + private val structType = + StructType(Seq(StructField("a", IntegerType), StructField("b", StringType))) + + private val allCols: Seq[Col] = Seq( + Col("c_bool", BooleanType, main = true), + Col("c_byte", ByteType, main = false), + Col("c_short", ShortType, main = false), + Col("c_int", IntegerType, main = true), + Col("c_long", LongType, main = true), + Col("c_float", FloatType, main = false), + Col("c_double", DoubleType, main = true), + Col("c_dec10", DecimalType(10, 2), main = false), + Col("c_dec38", DecimalType(38, 10), main = true), + Col("c_str", StringType, main = true), + Col("c_bin", BinaryType, main = true), + Col("c_date", DateType, main = true), + Col("c_ts", TimestampType, main = true), + Col("c_ntz", TimestampNTZType, main = false), + Col("c_arr", ArrayType(IntegerType), main = true), + Col("c_struct", structType, main = false), + Col("c_map", MapType(StringType, IntegerType), main = false)) + + private val largeColNames = + Seq("c_bool", "c_int", "c_long", "c_double", "c_dec38", "c_date", "c_ts") + + private val F1 = "id % 3 <> 1" + private val F2 = "id % 4 = 0" + + // ---------------------------------------------------------------- data generation + + private case class DataSet(name: String, cols: Seq[Col], singleView: String, multiView: String) + + private var tempRoot: File = _ + private val dataSets = mutable.Map[String, DataSet]() + + override protected def afterAll(): Unit = { + try { + if (tempRoot != null) FileUtils.deleteQuietly(tempRoot) + } finally { + super.afterAll() + } + } + + private def withConfs[T](pairs: (String, String)*)(f: => T): T = { + var result: Option[T] = None + withSQLConf(pairs: _*) { result = Some(f) } + result.get + } + + private def pick[T](r: Random, xs: Seq[T]): T = xs(r.nextInt(xs.length)) + + private def randomDigits(r: Random, maxDigits: Int): String = { + val n = 1 + r.nextInt(maxDigits) + (0 until n).map(_ => ('0' + r.nextInt(10)).toChar).mkString + } + + private def genValue(dt: DataType, r: Random): Any = dt match { + case BooleanType => r.nextBoolean() + case ByteType => + if (r.nextInt(5) == 0) pick(r, Seq(Byte.MinValue, Byte.MaxValue, 0.toByte)) + else (r.nextInt(256) - 128).toByte + case ShortType => + if (r.nextInt(5) == 0) pick(r, Seq(Short.MinValue, Short.MaxValue, 0.toShort)) + else (r.nextInt(65536) - 32768).toShort + case IntegerType => + if (r.nextInt(5) == 0) pick(r, Seq(Int.MinValue, Int.MaxValue, 0, -1)) else r.nextInt() + case LongType => + if (r.nextInt(5) == 0) pick(r, Seq(Long.MinValue, Long.MaxValue, 0L, -1L)) else r.nextLong() + case FloatType => + if (r.nextInt(3) == 0) { + pick( + r, + Seq( + Float.NaN, + -0.0f, + 0.0f, + Float.PositiveInfinity, + Float.NegativeInfinity, + Float.MinPositiveValue, + Float.MaxValue)) + } else r.nextFloat() * 2000f - 1000f + case DoubleType => + if (r.nextInt(3) == 0) { + pick( + r, + Seq( + Double.NaN, + -0.0d, + 0.0d, + Double.PositiveInfinity, + Double.NegativeInfinity, + Double.MinPositiveValue, + Double.MaxValue)) + } else r.nextDouble() * 2e6 - 1e6 + case d: DecimalType => + val unscaled = new java.math.BigInteger(randomDigits(r, d.precision)) + val signed = if (r.nextBoolean()) unscaled.negate() else unscaled + new java.math.BigDecimal(signed, d.scale) + case StringType => + r.nextInt(20) match { + case 0 => "" + case 1 => "ж€😀 ünï" + case 2 => r.alphanumeric.take(1500 + r.nextInt(1000)).mkString + case _ => r.alphanumeric.take(1 + r.nextInt(16)).mkString + } + case BinaryType => + if (r.nextInt(10) == 0) Array.emptyByteArray + else Array.fill(1 + r.nextInt(24))(r.nextInt(256).toByte) + case DateType => + Date.valueOf(LocalDate.ofEpochDay(r.nextInt(70000) - 20000L)) + case TimestampType => + Timestamp.from( + Instant + .ofEpochSecond(r.nextInt(2000000000) * 3L - 1800000000L, r.nextInt(1000000) * 1000L)) + case TimestampNTZType => + LocalDateTime.ofEpochSecond( + r.nextInt(2000000000) * 3L - 1800000000L, + r.nextInt(1000000) * 1000, + ZoneOffset.UTC) + case ArrayType(IntegerType, _) => + Seq.fill(r.nextInt(5))(if (r.nextInt(4) == 0) null else r.nextInt(100)) + case s: StructType if s == structType => + Row( + if (r.nextInt(4) == 0) null else r.nextInt(100), + if (r.nextInt(4) == 0) null else r.alphanumeric.take(r.nextInt(6)).mkString) + case MapType(StringType, IntegerType, _) => + (0 until r.nextInt(4)) + .map(i => s"k$i" -> (if (r.nextInt(4) == 0) null else r.nextInt(100))) + .toMap + case other => throw new IllegalArgumentException(s"unsupported type $other") + } + + private def genGroupColumn(c: Col, size: Int, r: Random, pool: IndexedSeq[Any]): Seq[Any] = { + def value(): Any = if (r.nextInt(10) < 7) pool(r.nextInt(pool.length)) else genValue(c.dt, r) + val boolMode = if (c.dt == BooleanType) r.nextInt(3) else -1 + def nonNull(): Any = boolMode match { + case 0 => true + case 1 => false + case _ => value() + } + val single = r.nextInt(size) + r.nextInt(7) match { + case 0 => Seq.fill(size)(nonNull()) + case 1 => Seq.fill(size)(null) + case 2 => (0 until size).map(i => if (i == 0) null else nonNull()) + case 3 => (0 until size).map(i => if (i == size - 1) null else nonNull()) + case 4 => (0 until size).map(i => if (i == single) nonNull() else null) + case 5 => (0 until size).map(i => if (i % 2 == 0) null else nonNull()) + case _ => (0 until size).map(_ => if (r.nextBoolean()) null else nonNull()) + } + } + + private def dataSet(name: String): DataSet = synchronized { + dataSets.getOrElseUpdate( + name, { + val (numGroups, sizeOf, cols) = name match { + case "mixed" => + val sizes: Random => Int = r => + r.nextInt(4) match { + case 0 => 1 + case 1 => 2 + r.nextInt(4) + case _ => 6 + r.nextInt(75) + } + (600, sizes, allCols) + case "large" => + ( + 200000, + (r: Random) => 1 + r.nextInt(4), + allCols.filter(c => largeColNames.contains(c.name))) + } + createDataSet(name, numGroups, sizeOf, cols) + }) + } + + private def createDataSet( + name: String, + numGroups: Int, + sizeOf: Random => Int, + cols: Seq[Col]): DataSet = { + val r = new Random(seed ^ name.hashCode) + val pools = cols.map(c => IndexedSeq.fill(30)(genValue(c.dt, r))) + val sizes = Array.fill(numGroups)(sizeOf(r)) + val groups: Array[Array[Array[Any]]] = Array.tabulate(numGroups) { g => + val perCol = cols.zip(pools).map { case (c, p) => genGroupColumn(c, sizes(g), r, p) } + Array.tabulate(sizes(g))(i => perCol.map(_(i)).toArray) + } + val order = r.shuffle(sizes.indices.flatMap(g => Seq.fill(sizes(g))(g))) + val next = new Array[Int](numGroups) + val rows = order.zipWithIndex.map { case (g, id) => + val values = groups(g)(next(g)) + next(g) += 1 + val k1 = if (g % 41 == 7) null else g / 2 + val k2 = if (g % 3 == 0) null else if (g % 2 == 0) "a" else "b" + Row.fromSeq(Seq[Any](id.toLong, g, k1, k2) ++ values) + } + val schema = StructType( + Seq( + StructField("id", LongType, nullable = false), + StructField("g", IntegerType, nullable = false), + StructField("k1", IntegerType), + StructField("k2", StringType)) ++ cols.map(c => StructField(c.name, c.dt))) + + if (tempRoot == null) tempRoot = Files.createTempDirectory("comet-agg-fuzz").toFile + val singlePath = new File(tempRoot, s"$name-single").getCanonicalPath + val multiPath = new File(tempRoot, s"$name-multi").getCanonicalPath + withSQLConf(CometConf.COMET_ENABLED.key -> "false") { + spark + .createDataFrame(spark.sparkContext.parallelize(rows, 1), schema) + .write + .option("parquet.block.size", 256 * 1024) + .parquet(singlePath) + spark + .createDataFrame(spark.sparkContext.parallelize(rows, 8), schema) + .write + .parquet(multiPath) + } + val ds = DataSet(name, cols, s"fuzz_${name}_single", s"fuzz_${name}_multi") + spark.read.parquet(singlePath).createOrReplaceTempView(ds.singleView) + spark.read.parquet(multiPath).createOrReplaceTempView(ds.multiView) + ds + } + + // ---------------------------------------------------------------- aggregates + + private sealed trait Check + private case object Exact extends Check + private case class Member(value: String, filter: Option[String], ignoreNulls: Boolean) + extends Check + + private case class Agg(sql: String, check: Check) + + private def firstLastAggs(c: String, all: Boolean): Seq[Agg] = { + val base = Seq( + Agg(s"first($c)", Member(c, None, ignoreNulls = false)), + Agg(s"last($c)", Member(c, None, ignoreNulls = false)), + Agg(s"first($c, true)", Member(c, None, ignoreNulls = true)), + Agg(s"last($c, true)", Member(c, None, ignoreNulls = true)), + Agg(s"any_value($c)", Member(c, None, ignoreNulls = false)), + Agg(s"first($c) FILTER (WHERE $F1)", Member(c, Some(F1), ignoreNulls = false)), + Agg(s"last($c, true) FILTER (WHERE $F1)", Member(c, Some(F1), ignoreNulls = true))) + val extra = Seq( + Agg(s"first_value($c) IGNORE NULLS", Member(c, None, ignoreNulls = true)), + Agg(s"last_value($c) IGNORE NULLS", Member(c, None, ignoreNulls = true)), + Agg(s"any_value($c, true)", Member(c, None, ignoreNulls = true)), + Agg(s"last($c) FILTER (WHERE $F2)", Member(c, Some(F2), ignoreNulls = false)), + Agg(s"first($c, true) FILTER (WHERE $F2)", Member(c, Some(F2), ignoreNulls = true))) + if (all) base ++ extra else base + } + + private val boolAggs: Seq[Agg] = Seq( + "max(c_bool)", + "min(c_bool)", + "bool_and(c_bool)", + "bool_or(c_bool)", + "every(c_bool)", + "some(c_bool)", + "any(c_bool)", + s"max(c_bool) FILTER (WHERE $F1)", + s"min(c_bool) FILTER (WHERE $F2)", + s"bool_and(c_bool) FILTER (WHERE $F1)", + s"bool_or(c_bool) FILTER (WHERE $F2)", + "count(c_bool)").map(Agg(_, Exact)) + + private def distinctAggs(native: Boolean): Seq[Agg] = { + val (d2, s, n) = if (native) ("c_date", "c_double", "c_ts") else ("c_str", "c_str", "c_arr") + Seq( + Agg("count(DISTINCT c_int)", Exact), + Agg(s"count(DISTINCT $d2)", Exact), + Agg("first(c_long)", Member("c_long", None, ignoreNulls = false)), + Agg(s"last($s, true)", Member(s, None, ignoreNulls = true)), + Agg("any_value(c_date)", Member("c_date", None, ignoreNulls = false)), + Agg(s"first(c_dec38) FILTER (WHERE $F1)", Member("c_dec38", Some(F1), ignoreNulls = false)), + Agg(s"last($n)", Member(n, None, ignoreNulls = false)), + Agg(s"first($n, true)", Member(n, None, ignoreNulls = true)), + Agg("max(c_bool)", Exact), + Agg("min(c_bool)", Exact), + Agg(s"bool_or(c_bool) FILTER (WHERE $F2)", Exact), + Agg("count(*)", Exact)) + } + + private val singleDistinctAggs: Seq[Agg] = Seq( + Agg("count(DISTINCT c_int)", Exact), + Agg("first(c_long, true)", Member("c_long", None, ignoreNulls = true)), + Agg("last(c_dec38)", Member("c_dec38", None, ignoreNulls = false)), + Agg(s"first(c_ts) FILTER (WHERE $F1)", Member("c_ts", Some(F1), ignoreNulls = false)), + Agg("max(c_bool)", Exact), + Agg("bool_and(c_bool)", Exact)) + + // ---------------------------------------------------------------- shapes + + private case class Shape( + name: String, + multi: Boolean, + spill: Boolean, + confs: Seq[(String, String)]) { + def ordered: Boolean = !multi && !spill + } + + private def aqe(on: Boolean) = SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> on.toString + private def partitions(n: Int) = SQLConf.SHUFFLE_PARTITIONS.key -> n.toString + private def batch(n: Int) = CometConf.COMET_BATCH_SIZE.key -> n.toString + private val tinyPool = CometConf.COMET_OFFHEAP_MEMORY_POOL_FRACTION.key -> "0.00005" + + private val singleAqe = + Shape("single_aqe", multi = false, spill = false, Seq(aqe(true), partitions(1))) + private val singleNoAqe = + Shape("single_noaqe", multi = false, spill = false, Seq(aqe(false), partitions(3))) + private val singleBatch8 = + Shape( + "single_batch8_noaqe", + multi = false, + spill = false, + Seq(aqe(false), partitions(1), batch(8))) + private val singleSpill = + Shape( + "single_spill", + multi = false, + spill = true, + Seq(aqe(true), partitions(1), batch(64), tinyPool)) + private val multiAqe = + Shape("multi_aqe", multi = true, spill = false, Seq(aqe(true), partitions(7))) + private val multiBatch8NoAqe = + Shape( + "multi_batch8_noaqe", + multi = true, + spill = false, + Seq(aqe(false), partitions(5), batch(8))) + private val multiSpill = + Shape( + "multi_spill", + multi = true, + spill = true, + Seq(aqe(true), partitions(3), batch(64), tinyPool)) + + private val allShapes = + Seq(singleAqe, singleNoAqe, singleBatch8, singleSpill, multiAqe, multiBatch8NoAqe, multiSpill) + + private def layoutConfs(multi: Boolean): Seq[(String, String)] = + if (multi) { + Seq( + SQLConf.FILES_MAX_PARTITION_BYTES.key -> "65536", + SQLConf.FILES_OPEN_COST_IN_BYTES.key -> "1") + } else { + Seq( + SQLConf.FILES_MAX_PARTITION_BYTES.key -> (1L << 30).toString, + SQLConf.FILES_OPEN_COST_IN_BYTES.key -> (1L << 30).toString) + } + + // ---------------------------------------------------------------- comparison + + private def norm(v: Any): Any = v match { + case null => null + case b: Array[Byte] => ("bin", b.toList) + case d: Double => ("d", java.lang.Double.doubleToLongBits(d)) + case f: Float => ("f", java.lang.Float.floatToIntBits(f)) + case r: Row => r.toSeq.map(norm).toList + case m: scala.collection.Map[_, _] => m.map { case (k, x) => (norm(k), norm(x)) }.toMap + case s: scala.collection.Seq[_] => s.map(norm).toList + case other => other + } + + private def show(v: Any): String = v match { + case null => "null" + case b: Array[Byte] => b.mkString("bin[", ",", "]") + case s: String if s.length > 60 => s"${s.take(60)}...(${s.length} chars)" + case d: Double if d == 0.0 && 1.0 / d < 0 => "-0.0" + case f: Float if f == 0.0f && 1.0f / f < 0 => "-0.0f" + case other => String.valueOf(other) + } + + private case class CaseResult(native: Boolean, aggSpills: Long, inputPartitions: Int) + + private def aggregatesNative(plan: SparkPlan): Boolean = { + val comet = collect(plan) { case a: CometHashAggregateExec => a } + val spark = collect(plan) { case a: BaseAggregateExec => a } + comet.nonEmpty && spark.isEmpty + } + + private def runCase( + label: String, + ds: DataSet, + shape: Shape, + keys: Seq[String], + aggs: Seq[Agg], + expectNative: Boolean): CaseResult = { + val view = if (shape.multi) ds.multiView else ds.singleView + val select = (keys ++ aggs.zipWithIndex.map { case (a, i) => s"${a.sql} AS a$i" }) + .mkString(", ") + val groupBy = if (keys.isEmpty) "" else keys.mkString(" GROUP BY ", ", ", "") + val query = s"SELECT $select FROM $view$groupBy" + val ctx = s"[seed=$seed case=$label shape=${shape.name} data=${ds.name} " + + s"keys=${keys.mkString(",")}]" + val nk = keys.length + + withConfs(shape.confs ++ layoutConfs(shape.multi): _*) { + val inputPartitions = withConfs(CometConf.COMET_ENABLED.key -> "false") { + spark.table(view).rdd.getNumPartitions + } + assert(shape.multi == (inputPartitions > 1), s"$ctx input partitions: $inputPartitions") + val expected = withConfs(CometConf.COMET_ENABLED.key -> "false") { + sql(query).collect().toSeq + } + val df = sql(query) + val actual = + try df.collect().toSeq + catch { + case e: Throwable => fail(s"$ctx Comet failed for query:\n$query", e) + } + val plan = df.queryExecution.executedPlan + val native = aggregatesNative(plan) + if (expectNative) { + assert(native, s"$ctx expected native aggregates, plan:\n$plan") + } + val aggSpills = collect(plan) { case a: CometHashAggregateExec => a } + .flatMap(_.metrics.get("spill_count")) + .map(_.value) + .sum + + def keyed(rows: Seq[Row], what: String): Map[Any, Row] = { + val m = rows.groupBy(r => norm(Row.fromSeq(r.toSeq.take(nk)))) + val dups = m.filter(_._2.size > 1) + assert(dups.isEmpty, s"$ctx $what has duplicate groups: ${dups.keys.take(5)}") + m.map { case (k, v) => k -> v.head } + } + val exp = keyed(expected, "Spark") + val act = keyed(actual, "Comet") + assert( + exp.keySet == act.keySet, + s"$ctx group sets differ: missing in Comet ${(exp.keySet -- act.keySet).take(5)}, " + + s"extra in Comet ${(act.keySet -- exp.keySet).take(5)}\n$query") + + val members = aggs.map(_.check).collect { case m: Member => m }.distinct + val candidates = + if (members.isEmpty) Map.empty[Member, Map[Any, (Set[Any], Boolean, Long)]] + else candidateSets(view, keys, members) + + val errors = mutable.ArrayBuffer[String]() + for ((k, e) <- exp; a = act(k); (agg, i) <- aggs.zipWithIndex) { + val ev = e.get(nk + i) + val av = a.get(nk + i) + def err(msg: String): Unit = + errors += s"group ${show(k)} ${agg.sql}: $msg (Comet ${show(av)}, Spark ${show(ev)})" + agg.check match { + case Exact => + if (norm(ev) != norm(av)) err("differs") + case m: Member => + if (shape.ordered && native && norm(ev) != norm(av)) { + err("differs on an ordered single split") + } + val (vals, hasNull, n) = candidates(m)(k) + def allowed(v: Any): Boolean = + if (n == 0) v == null + else if (v == null) { if (m.ignoreNulls) vals.isEmpty else hasNull } + else vals.contains(norm(v)) + if (!allowed(av)) err("Comet value is not a value of the group") + if (!allowed(ev)) err("Spark value is not a value of the group") + } + } + if (errors.nonEmpty) { + fail( + s"$ctx ${errors.size} mismatches, first ones:\n" + + errors.take(25).mkString("\n") + s"\nquery: $query\nplan:\n$plan") + } + CaseResult(native, aggSpills, inputPartitions) + } + } + + private def candidateSets( + view: String, + keys: Seq[String], + members: Seq[Member]): Map[Member, Map[Any, (Set[Any], Boolean, Long)]] = { + val nk = keys.length + val cols = members.zipWithIndex.flatMap { case (m, i) => + val f = m.filter.map(c => s" FILTER (WHERE $c)").getOrElse("") + Seq( + s"collect_list(${m.value})$f AS v$i", + s"count(1)$f AS n$i", + s"count(${m.value})$f AS nn$i") + } + val groupBy = if (keys.isEmpty) "" else keys.mkString(" GROUP BY ", ", ", "") + val rows = withConfs(CometConf.COMET_ENABLED.key -> "false") { + sql(s"SELECT ${(keys ++ cols).mkString(", ")} FROM $view$groupBy").collect().toSeq + } + members.zipWithIndex.map { case (m, i) => + m -> rows.map { r => + val k = norm(Row.fromSeq(r.toSeq.take(nk))) + val vals = r.getSeq[Any](nk + 3 * i).map(norm).toSet + val n = r.getLong(nk + 3 * i + 1) + val nn = r.getLong(nk + 3 * i + 2) + k -> ((vals, n > nn, n)) + }.toMap + }.toMap + } + + // ---------------------------------------------------------------- matrix + + private val groupings: Seq[Seq[String]] = Seq(Nil, Seq("k1"), Seq("k1", "k2")) + + private val nativeTypes: Set[String] = Set( + "c_bool", + "c_byte", + "c_short", + "c_int", + "c_long", + "c_float", + "c_double", + "c_dec10", + "c_dec38", + "c_date", + "c_ts", + "c_ntz") + + private def colAggs(cols: Seq[Col], all: Boolean): Seq[Agg] = + cols.flatMap(c => firstLastAggs(c.name, all)) ++ + (if (cols.exists(_.name == "c_bool")) boolAggs else Nil) + + private case class Query(label: String, aggs: Seq[Agg], expectNative: Boolean) + + private def queries(ds: DataSet, onlyMain: Boolean, all: Boolean): Seq[Query] = { + val cols = ds.cols.filter(c => c.main || !onlyMain) + val (native, other) = cols.partition(c => nativeTypes.contains(c.name)) + Query("native_types", colAggs(native, all), expectNative = true) +: + other.map(c => Query(s"type_${c.name}", colAggs(Seq(c), all), expectNative = false)) + } + + private def runAll( + ds: DataSet, + shape: Shape, + keys: Seq[String], + qs: Seq[Query]): Seq[CaseResult] = + qs.map(q => runCase(q.label, ds, shape, keys, q.aggs, q.expectNative)) + + private def assumeFull(): Unit = + assume(fullMatrix, "the full matrix runs only with -Dcomet.test.aggFuzz.full=true") + + // ---------------------------------------------------------------- default tests + + test("ordered single split: exact first/last, bool aggregates and FILTER") { + val ds = dataSet("mixed") + for (keys <- groupings) { + runAll(ds, singleAqe, keys, queries(ds, onlyMain = true, all = keys.nonEmpty)) + } + } + + test("small batches with partial emission, AQE off") { + val ds = dataSet("mixed") + runAll(ds, singleBatch8, Seq("k1", "k2"), queries(ds, onlyMain = true, all = false)) + } + + test("many partitions across a shuffle: values belong to their group") { + val ds = dataSet("mixed") + for (keys <- Seq(Nil, Seq("k1", "k2"))) { + runAll(ds, multiAqe, keys, queries(ds, onlyMain = true, all = false)) + } + } + + test("forced spill of the aggregate state") { + val ds = dataSet("mixed") + val native = ds.cols.filter(c => c.main && nativeTypes.contains(c.name)) + val r = runCase("spill", ds, multiSpill, Seq("k1", "k2"), colAggs(native, all = false), true) + assert(r.aggSpills > 0, s"[seed=$seed] the aggregate did not spill") + } + + test("COUNT(DISTINCT) rewrite: first/last/min/max with FILTER over Expand") { + val ds = dataSet("mixed") + for (shape <- Seq(singleAqe, multiAqe)) { + runCase("distinct", ds, shape, Seq("k1"), distinctAggs(native = true), expectNative = true) + } + runCase("distinct_one", ds, singleAqe, Seq("k1"), singleDistinctAggs, expectNative = true) + } + + // Comet runs Spark's PartialMerge aggregate as a DataFusion Partial aggregate, which emits + // groups early under memory pressure, so the de-duplicating aggregate of a single COUNT(DISTINCT) + // passes duplicates to the count in the same stage. + ignore("known bug: COUNT(DISTINCT) over a spilling PartialMerge aggregate overcounts") { + val ds = dataSet("mixed") + runCase("distinct_one", ds, multiSpill, Seq("k1"), singleDistinctAggs, expectNative = true) + } + + // ---------------------------------------------------------------- full matrix + + for (shape <- allShapes; keys <- groupings) { + test(s"full: ${shape.name}, group by [${keys.mkString(", ")}]") { + assumeFull() + val ds = dataSet("mixed") + val all = queries(ds, onlyMain = false, all = true) ++ Seq( + Query("distinct", distinctAggs(native = true), expectNative = true), + Query("distinct_fallback_types", distinctAggs(native = false), expectNative = false), + Query("distinct_one", singleDistinctAggs, expectNative = true)) + val runnable = + if (shape.spill) all.filter(q => q.expectNative && q.label != "distinct_one") else all + runAll(ds, shape, keys, runnable) + } + } + + for ((shape, keys) <- Seq( + singleAqe -> Seq("g"), + singleSpill -> Seq("g"), + multiAqe -> Seq("k1", "k2"), + multiSpill -> Seq("g"))) { + test(s"full: 200k groups, ${shape.name}, group by [${keys.mkString(", ")}]") { + assumeFull() + val ds = dataSet("large") + val results = runAll(ds, shape, keys, queries(ds, onlyMain = false, all = false)) :+ + runCase("distinct", ds, shape, keys, distinctAggs(native = true), expectNative = true) + if (shape.spill) { + assert(results.exists(_.aggSpills > 0), s"[seed=$seed] the aggregate did not spill") + } + } + } +} From 6c8193e76b2be018f725ff8838c35acd8a2fcce9 Mon Sep 17 00:00:00 2001 From: msaf Date: Tue, 6 Oct 2026 14:38:29 +0100 Subject: [PATCH 06/13] fix: run a PartialMerge aggregate so that it spills instead of emitting a group twice Spark plans a single COUNT(DISTINCT x) (with other aggregates) as Partial(k, x) -> PartialMerge(k, x) -> [PartialMerge, Partial count(x)](k) -> Final. The PartialMerge stage must hand each (k, x) to the next aggregate once: that aggregate counts x without de-duplicating it. Comet ran PartialMerge as a DataFusion Partial aggregate with MergeAsPartial-wrapped accumulators, and DataFusion's PartialHashAggregateStream emits all held groups and starts over when its reservation is refused (hash_stream.rs "Step 4: Larger-than-memory execution (early emit)"), so under memory pressure the same (k, x) left it several times and COUNT(DISTINCT) overcounted (3x in the new test); SUM/MAX and other non-distinct aggregates stayed correct. PartialMerge now runs as DataFusion PartialReduce (states in, states out). DataFusion 55.1 routes PartialReduce to PartialReduceHashAggregateStream, which cannot spill, unless the pool reports a finite limit, and then to the legacy GroupedHashAggregateStream. Comet's pools report an unknown limit, so the vendored crate now sends PartialReduce under any non-infinite pool to FinalHashAggregateStream / OrderedFinalAggregateStream, which spill and replay sorted runs, and emit merged states instead of final values for it. The mixed {PartialMerge, Partial} stage keeps Partial + MergeAsPartial: a Final always merges its output. Skip-partial stays disabled for plans with a PartialMerge. Tests: a native test feeds a PartialMerge three copies of 24000 (k, x) states through a 256 KiB fair pool (23679 groups came out more than once before); CometAggregateSuite checks a grouped and a global single COUNT(DISTINCT) whose PartialMerge spills against Spark (25716 vs 8572 per group before); the fuzz suite's known-bug case is a test again and the spill shapes of the full matrix include distinct_one. Co-Authored-By: Claude Opus 5.5 (1M context) --- native/core/src/execution/jni_api.rs | 167 ++++++++++++++++++ native/core/src/execution/planner.rs | 15 +- .../aggregate_hash_table/final_table.rs | 8 + .../ordered_final_table.rs | 7 + .../src/aggregates/hash_stream.rs | 15 +- .../src/aggregates/mod.rs | 22 ++- .../src/aggregates/ordered_final_stream.rs | 23 ++- .../comet/exec/CometAggregateSuite.scala | 36 ++++ .../exec/CometFirstLastBoolAggFuzzSuite.scala | 12 +- 9 files changed, 279 insertions(+), 26 deletions(-) diff --git a/native/core/src/execution/jni_api.rs b/native/core/src/execution/jni_api.rs index 9002e83194b..a7f5f89aa2c 100644 --- a/native/core/src/execution/jni_api.rs +++ b/native/core/src/execution/jni_api.rs @@ -4225,3 +4225,170 @@ mod aggregate_offset_overflow_tests { ); } } + +#[cfg(test)] +mod partial_merge_spill_tests { + use super::*; + use crate::execution::memory_pools::fair_unified_pool_with_fake_spark; + use crate::execution::operators::InputBatch; + use arrow::array::{Int32Array, Int64Array}; + use datafusion_comet_proto::spark_expression::{self, Expr}; + use datafusion_comet_proto::spark_operator::{self, operator::OpStruct}; + use std::collections::HashMap as StdHashMap; + + const KEYS: i32 = 4; + const VALUES: i32 = 6000; + const COPIES: usize = 3; + const BATCH_ROWS: usize = 1024; + + fn data_type(type_id: i32) -> spark_expression::DataType { + spark_expression::DataType { + type_id, + type_info: None, + } + } + + fn bound(index: i32, type_id: i32) -> Expr { + Expr { + expr_struct: Some(spark_expression::expr::ExprStruct::Bound( + spark_expression::BoundReference { + index, + datatype: Some(data_type(type_id)), + }, + )), + query_context: None, + expr_id: None, + } + } + + /// `HashAggregate(keys = [k, x], functions = [merge count])`: the de-duplicating stage + /// Spark plans under a single COUNT(DISTINCT x) GROUP BY k, fed `(k, x, count)` states. + fn partial_merge_count() -> Operator { + let scan = Operator { + op_struct: Some(OpStruct::Scan(spark_operator::Scan { + fields: vec![data_type(3), data_type(3), data_type(4)], + source: "states".to_string(), + })), + ..Default::default() + }; + let count = spark_expression::AggExpr { + expr_struct: Some(AggExprStruct::Count(spark_expression::Count { + children: vec![bound(1, 3)], + })), + ..Default::default() + }; + Operator { + children: vec![scan], + op_struct: Some(OpStruct::HashAgg(spark_operator::HashAggregate { + grouping_exprs: vec![bound(0, 3), bound(1, 3)], + agg_exprs: vec![count], + mode: AggregateMode::PartialMerge as i32, + expr_modes: vec![AggregateMode::PartialMerge as i32], + initial_input_buffer_offset: 2, + })), + ..Default::default() + } + } + + fn input_batches() -> Vec { + let mut rows: Vec<(i32, i32)> = (0..COPIES) + .flat_map(|_| (0..KEYS).flat_map(|k| (0..VALUES).map(move |x| (k, x)))) + .collect(); + let mut seed = 0x9e37_79b9_7f4a_7c15u64; + for i in (1..rows.len()).rev() { + seed ^= seed << 13; + seed ^= seed >> 7; + seed ^= seed << 17; + rows.swap(i, (seed % (i as u64 + 1)) as usize); + } + rows.chunks(BATCH_ROWS) + .map(|chunk| { + InputBatch::Batch( + vec![ + Arc::new(Int32Array::from_iter_values(chunk.iter().map(|r| r.0))), + Arc::new(Int32Array::from_iter_values(chunk.iter().map(|r| r.1))), + Arc::new(Int64Array::from(vec![1i64; chunk.len()])), + ], + chunk.len(), + ) + }) + .chain(std::iter::once(InputBatch::EOF)) + .collect() + } + + /// Under memory pressure a PartialMerge must spill rather than emit a group twice, or the + /// COUNT(DISTINCT) above it counts the repeated value again. + #[tokio::test] + async fn partial_merge_spills_instead_of_emitting_a_group_twice() { + let share = 256 * 1024; + let (pool, _spark) = fair_unified_pool_with_fake_spark(share, share); + let spill_dir = tempfile::tempdir().unwrap(); + let operator = partial_merge_count(); + let session = Arc::new( + prepare_datafusion_session_context( + BATCH_ROWS, + PlanCancellation::new(), + pool, + vec![spill_dir.path().to_string_lossy().into_owned()], + u64::MAX, + 1, + &StdHashMap::new(), + &operator, + Some(share), + ) + .unwrap(), + ); + let planner = PhysicalPlanner::new(Arc::clone(&session), 0); + let (mut scans, _, plan) = planner.create_plan(&operator, &mut vec![], 1).unwrap(); + let mut stream = plan.native_plan.execute(0, session.task_ctx()).unwrap(); + let mut input = input_batches().into_iter(); + + let mut counts: StdHashMap<(i32, i32), Vec> = StdHashMap::new(); + while let Some(batch) = futures::future::poll_fn(|cx| { + let result = stream.poll_next_unpin(cx); + if result.is_pending() && scans[0].batch.try_lock().unwrap().is_none() { + if let Some(batch) = input.next() { + scans[0].set_input_batch(batch); + cx.waker().wake_by_ref(); + } + } + result + }) + .await + { + let batch = batch.unwrap(); + let k = batch + .column(0) + .as_any() + .downcast_ref::() + .unwrap(); + let x = batch + .column(1) + .as_any() + .downcast_ref::() + .unwrap(); + let c = batch + .column(2) + .as_any() + .downcast_ref::() + .unwrap(); + for row in 0..batch.num_rows() { + counts + .entry((k.value(row), x.value(row))) + .or_default() + .push(c.value(row)); + } + } + + let spills = plan + .native_plan + .metrics() + .and_then(|m| m.spill_count()) + .unwrap_or(0); + let repeated = counts.values().filter(|c| c.len() > 1).count(); + assert_eq!(repeated, 0, "{repeated} groups emitted more than once"); + assert_eq!(counts.len(), (KEYS * VALUES) as usize); + assert!(counts.values().all(|c| c == &[COPIES as i64])); + assert!(spills > 0, "the PartialMerge aggregate did not spill"); + } +} diff --git a/native/core/src/execution/planner.rs b/native/core/src/execution/planner.rs index fb691cdc904..620fb99b8f2 100644 --- a/native/core/src/execution/planner.rs +++ b/native/core/src/execution/planner.rs @@ -1544,19 +1544,20 @@ impl PhysicalPlanner { agg.mode )) })?; + // A PartialMerge feeds groups it must not repeat (the distinct values of a + // single COUNT(DISTINCT)), so it runs as PartialReduce, which spills under + // memory pressure where Partial emits groups early. let mode = match proto_mode { ProtoAggregateMode::Partial => DFAggregateMode::Partial, ProtoAggregateMode::Final => DFAggregateMode::Final, - // PartialMerge: Partial + MergeAsPartial - ProtoAggregateMode::PartialMerge => DFAggregateMode::Partial, + ProtoAggregateMode::PartialMerge => DFAggregateMode::PartialReduce, }; - // Check if any expression uses PartialMerge mode. When present, - // those expressions are wrapped with MergeAsPartial to get merge - // semantics inside a Partial-mode AggregateExec. + // A mixed {Partial, PartialMerge} aggregate runs as Partial and wraps its + // PartialMerge expressions with MergeAsPartial to get merge semantics. let partial_merge_value = ProtoAggregateMode::PartialMerge as i32; - let has_partial_merge = proto_mode == ProtoAggregateMode::PartialMerge - || agg.expr_modes.contains(&partial_merge_value); + let has_partial_merge = proto_mode == ProtoAggregateMode::Partial + && agg.expr_modes.contains(&partial_merge_value); let agg_exprs: PhyAggResult = agg .agg_exprs diff --git a/native/vendor/datafusion-physical-plan/src/aggregates/aggregate_hash_table/final_table.rs b/native/vendor/datafusion-physical-plan/src/aggregates/aggregate_hash_table/final_table.rs index 47758ecab16..d61214c4a64 100644 --- a/native/vendor/datafusion-physical-plan/src/aggregates/aggregate_hash_table/final_table.rs +++ b/native/vendor/datafusion-physical-plan/src/aggregates/aggregate_hash_table/final_table.rs @@ -61,6 +61,14 @@ impl AggregateHashTable { self.next_output_batch_inner(HashAggregateAccumulator::evaluate_to_columns) } + /// COMET PATCH: emits group keys with merged states, for a + /// [`crate::aggregates::AggregateMode::PartialReduce`] run by the final stream. + pub(in crate::aggregates) fn next_state_output_batch( + &mut self, + ) -> Result> { + self.next_output_batch_inner(HashAggregateAccumulator::state) + } + /// Final aggregation consumes partial aggregate states and merges them into /// the table's partial-state accumulators. pub(in crate::aggregates) fn aggregate_batch( diff --git a/native/vendor/datafusion-physical-plan/src/aggregates/aggregate_hash_table/ordered_final_table.rs b/native/vendor/datafusion-physical-plan/src/aggregates/aggregate_hash_table/ordered_final_table.rs index fd064ebffec..d13139a03f5 100644 --- a/native/vendor/datafusion-physical-plan/src/aggregates/aggregate_hash_table/ordered_final_table.rs +++ b/native/vendor/datafusion-physical-plan/src/aggregates/aggregate_hash_table/ordered_final_table.rs @@ -82,4 +82,11 @@ impl OrderedAggregateTable { ) -> Result> { self.next_output_batch_for_mode(true) } + + /// COMET PATCH: see `AggregateHashTable::::next_state_output_batch`. + pub(in crate::aggregates) fn next_state_output_batch( + &mut self, + ) -> Result> { + self.next_output_batch_for_mode(false) + } } diff --git a/native/vendor/datafusion-physical-plan/src/aggregates/hash_stream.rs b/native/vendor/datafusion-physical-plan/src/aggregates/hash_stream.rs index d3b63c7083f..1834259b70c 100644 --- a/native/vendor/datafusion-physical-plan/src/aggregates/hash_stream.rs +++ b/native/vendor/datafusion-physical-plan/src/aggregates/hash_stream.rs @@ -251,6 +251,10 @@ pub(crate) struct FinalHashAggregateStream { /// See comments for the same variable in [`PartialHashAggregateStream`]. group_values_soft_limit: Option, + /// COMET PATCH: a `PartialReduce` aggregate runs here so that it spills under memory + /// pressure, and emits merged states instead of final values. + emit_states: bool, + /// Tracks the high-level stream lifecycle. The hash table owns the lower-level /// state for emitting output batches. state: Option, @@ -1031,7 +1035,9 @@ impl FinalHashAggregateStream { ) -> Result { debug_assert!(matches!( agg.mode, - super::AggregateMode::Final | super::AggregateMode::FinalPartitioned + super::AggregateMode::Final + | super::AggregateMode::FinalPartitioned + | super::AggregateMode::PartialReduce )); debug_assert_eq!(agg.input_order_mode, InputOrderMode::Linear); @@ -1074,6 +1080,7 @@ impl FinalHashAggregateStream { baseline_metrics, reservation, group_values_soft_limit: agg.limit_options().map(|config| config.limit()), + emit_states: agg.mode == super::AggregateMode::PartialReduce, state: Some(FinalHashAggregateState::ReadingInput { hash_table, spill_context, @@ -1435,7 +1442,11 @@ impl FinalHashAggregateStream { let elapsed_compute = self.baseline_metrics.elapsed_compute().clone(); let timer = elapsed_compute.timer(); - let result = hash_table.next_output_batch(); + let result = if self.emit_states { + hash_table.next_state_output_batch() + } else { + hash_table.next_output_batch() + }; timer.done(); match result { diff --git a/native/vendor/datafusion-physical-plan/src/aggregates/mod.rs b/native/vendor/datafusion-physical-plan/src/aggregates/mod.rs index ddcf46a69c2..e69d1bad404 100644 --- a/native/vendor/datafusion-physical-plan/src/aggregates/mod.rs +++ b/native/vendor/datafusion-physical-plan/src/aggregates/mod.rs @@ -1254,7 +1254,9 @@ impl AggregateExec { fn should_use_final_hash_stream(&self, _context: &TaskContext) -> bool { matches!( self.mode, - AggregateMode::Final | AggregateMode::FinalPartitioned + AggregateMode::Final + | AggregateMode::FinalPartitioned + | AggregateMode::PartialReduce ) && self.limit_options_supported_by_hash_stream() && self.input_order_mode == InputOrderMode::Linear && !self.group_by.is_true_no_grouping() @@ -1263,7 +1265,9 @@ impl AggregateExec { fn should_use_partial_reduce_hash_stream(&self, context: &TaskContext) -> bool { // TODO: implement memory-limited path and remove this limitation - if matches!(context.memory_pool().memory_limit(), MemoryLimit::Finite(_)) { + // COMET PATCH: a pool of unknown size can refuse memory too, and this stream fails + // then. Leave it to the final hash stream, which spills. + if !matches!(context.memory_pool().memory_limit(), MemoryLimit::Infinite) { return false; } @@ -1287,7 +1291,9 @@ impl AggregateExec { fn should_use_ordered_final_aggregate_stream(&self, _context: &TaskContext) -> bool { matches!( self.mode, - AggregateMode::Final | AggregateMode::FinalPartitioned + AggregateMode::Final + | AggregateMode::FinalPartitioned + | AggregateMode::PartialReduce ) && self.limit_options_supported_by_hash_stream() && self.input_order_mode != InputOrderMode::Linear && !self.group_by.is_true_no_grouping() @@ -4336,8 +4342,8 @@ mod tests { Ok(()) } - /// Spilling behavior is not implemented for partial-reduce stream yet, so fall - /// back to the existing `GroupedHashAggregateStream` + /// Spilling behavior is not implemented for partial-reduce stream yet. + /// COMET PATCH: fall back to `FinalHashAggregateStream`, which spills and emits states. #[tokio::test] async fn partial_reduce_aggregate_with_memory_limit_planning() -> Result<()> { let partial_reduce = partial_reduce_test_aggregate()?; @@ -4355,7 +4361,11 @@ mod tests { ); let stream = partial_reduce.execute_typed(0, &task_ctx)?; - assert!(matches!(stream, StreamType::GroupedHash(_))); + assert!(matches!(stream, StreamType::FinalHash(_))); + let stream: SendableRecordBatchStream = stream.into(); + let output = collect(stream).await?; + assert_eq!(output.iter().map(RecordBatch::num_rows).sum::(), 3); + assert_eq!(output[0].schema(), partial_reduce.schema()); Ok(()) } diff --git a/native/vendor/datafusion-physical-plan/src/aggregates/ordered_final_stream.rs b/native/vendor/datafusion-physical-plan/src/aggregates/ordered_final_stream.rs index 5696a6fdc10..35ccdfda49c 100644 --- a/native/vendor/datafusion-physical-plan/src/aggregates/ordered_final_stream.rs +++ b/native/vendor/datafusion-physical-plan/src/aggregates/ordered_final_stream.rs @@ -70,6 +70,8 @@ pub(crate) struct OrderedFinalAggregateStream { /// COMET PATCH: the group keys are fully ordered, so the table holds only the groups /// the last batch may continue. fully_ordered: bool, + /// COMET PATCH: a `PartialReduce` aggregate emits merged states instead of final values. + emit_states: bool, } /// Spill configuration and accumulated runs for partially ordered final @@ -269,7 +271,9 @@ impl OrderedFinalAggregateStream { ) -> Result { debug_assert!(matches!( agg.mode, - AggregateMode::Final | AggregateMode::FinalPartitioned + AggregateMode::Final + | AggregateMode::FinalPartitioned + | AggregateMode::PartialReduce )); debug_assert_ne!(agg.input_order_mode, InputOrderMode::Linear); @@ -329,7 +333,9 @@ impl OrderedFinalAggregateStream { ) -> Result { debug_assert!(matches!( agg.mode, - AggregateMode::Final | AggregateMode::FinalPartitioned + AggregateMode::Final + | AggregateMode::FinalPartitioned + | AggregateMode::PartialReduce )); debug_assert_ne!(*input_order_mode, InputOrderMode::Linear); @@ -374,6 +380,7 @@ impl OrderedFinalAggregateStream { spill_context, }), fully_ordered: *input_order_mode == InputOrderMode::Sorted, + emit_states: agg.mode == AggregateMode::PartialReduce, }) } @@ -504,7 +511,11 @@ impl OrderedFinalAggregateStream { Ok(None) } else { let timer = elapsed_compute.timer(); - let result = table.next_output_batch(); + let result = if self.emit_states { + table.next_state_output_batch() + } else { + table.next_output_batch() + }; timer.done(); result }; @@ -760,7 +771,11 @@ impl OrderedFinalAggregateStream { let mut table = table; let elapsed_compute = self.baseline_metrics.elapsed_compute().clone(); let timer = elapsed_compute.timer(); - let result = table.next_output_batch(); + let result = if self.emit_states { + table.next_state_output_batch() + } else { + table.next_output_batch() + }; timer.done(); match result { 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 dca5f3f05f5..e82831ec426 100644 --- a/spark/src/test/scala/org/apache/comet/exec/CometAggregateSuite.scala +++ b/spark/src/test/scala/org/apache/comet/exec/CometAggregateSuite.scala @@ -1359,6 +1359,42 @@ class CometAggregateSuite extends CometTestBase with AdaptiveSparkPlanHelper { }) } + test("partialMerge - single count distinct is exact when the PartialMerge aggregate spills") { + withTempPath { dir => + withSQLConf(CometConf.COMET_ENABLED.key -> "false") { + spark + .range(0, 180000, 1, 6) + .selectExpr("(id * 7919) % 60000 AS x", "id AS v") + .selectExpr("x % 7 AS k", "x", "v") + .write + .parquet(dir.getAbsolutePath) + } + spark.read.parquet(dir.getAbsolutePath).createOrReplaceTempView("pm_spill") + for (partitions <- Seq(1, 3)) { + withSQLConf( + CometConf.COMET_OFFHEAP_MEMORY_POOL_FRACTION.key -> "0.00005", + CometConf.COMET_BATCH_SIZE.key -> "64", + SQLConf.SHUFFLE_PARTITIONS.key -> partitions.toString) { + for (q <- Seq( + "SELECT k, count(DISTINCT x), sum(v), max(v) FROM pm_spill GROUP BY k", + "SELECT count(DISTINCT x), sum(v), min(v) FROM pm_spill")) { + val (_, plan) = checkSparkAnswerAndOperator(sql(q)) + val partialMerges = collect(plan) { + case a: CometHashAggregateExec + if a.aggregateExpressions.nonEmpty && + a.aggregateExpressions.forall(_.mode == PartialMerge) => + a + } + assert(partialMerges.nonEmpty, s"no PartialMerge aggregate in:\n$plan") + val spills = + partialMerges.map(_.metrics.get("spill_count").map(_.value).getOrElse(0L)).sum + assert(spills > 0, s"the PartialMerge aggregate did not spill: $q") + } + } + } + } + } + test("partialMerge - distinct + non-distinct aggregates (Expand pattern)") { withParquetTable((1 to 100).map(i => (i, i.toString)), "tbl", false) { checkSparkAnswerAndOperator("SELECT avg(_1), sum(_1), count(distinct _1) FROM tbl") diff --git a/spark/src/test/scala/org/apache/comet/exec/CometFirstLastBoolAggFuzzSuite.scala b/spark/src/test/scala/org/apache/comet/exec/CometFirstLastBoolAggFuzzSuite.scala index 1c881335d67..1e082a3df20 100644 --- a/spark/src/test/scala/org/apache/comet/exec/CometFirstLastBoolAggFuzzSuite.scala +++ b/spark/src/test/scala/org/apache/comet/exec/CometFirstLastBoolAggFuzzSuite.scala @@ -631,12 +631,11 @@ class CometFirstLastBoolAggFuzzSuite extends CometTestBase with AdaptiveSparkPla runCase("distinct_one", ds, singleAqe, Seq("k1"), singleDistinctAggs, expectNative = true) } - // Comet runs Spark's PartialMerge aggregate as a DataFusion Partial aggregate, which emits - // groups early under memory pressure, so the de-duplicating aggregate of a single COUNT(DISTINCT) - // passes duplicates to the count in the same stage. - ignore("known bug: COUNT(DISTINCT) over a spilling PartialMerge aggregate overcounts") { + test("COUNT(DISTINCT) over a spilling PartialMerge aggregate") { val ds = dataSet("mixed") - runCase("distinct_one", ds, multiSpill, Seq("k1"), singleDistinctAggs, expectNative = true) + val r = + runCase("distinct_one", ds, multiSpill, Seq("k1"), singleDistinctAggs, expectNative = true) + assert(r.aggSpills > 0, s"[seed=$seed] the aggregate did not spill") } // ---------------------------------------------------------------- full matrix @@ -649,8 +648,7 @@ class CometFirstLastBoolAggFuzzSuite extends CometTestBase with AdaptiveSparkPla Query("distinct", distinctAggs(native = true), expectNative = true), Query("distinct_fallback_types", distinctAggs(native = false), expectNative = false), Query("distinct_one", singleDistinctAggs, expectNative = true)) - val runnable = - if (shape.spill) all.filter(q => q.expectNative && q.label != "distinct_one") else all + val runnable = if (shape.spill) all.filter(_.expectNative) else all runAll(ds, shape, keys, runnable) } } From 88c3edfb28fce8c912d1b51d913d4fd5467b076d Mon Sep 17 00:00:00 2001 From: msaf Date: Tue, 6 Oct 2026 18:54:05 +0100 Subject: [PATCH 07/13] fix: exempt a validity interval only when both bounds read one row of one input fbj_order_type's left join took its lower bound from replenishments and its upper bound from orders joined below the sort-merge join, a band join that the validity-interval shape let run natively. Trace each bound to the operator that produces it and charge smjCondition when a join or union lies between them; a LEAD-built upper bound counts as its own source. Co-Authored-By: Claude Opus 5.5 (1M context) --- .../comet/rules/CostBasedEngineChoice.scala | 11 ++- .../apache/comet/rules/EngineCostTable.scala | 5 +- .../comet/rules/JoinConditionShape.scala | 91 ++++++++++++++++--- .../rules/CostBasedEngineChoiceSuite.scala | 30 +++++- 4 files changed, 118 insertions(+), 19 deletions(-) diff --git a/spark/src/main/scala/org/apache/comet/rules/CostBasedEngineChoice.scala b/spark/src/main/scala/org/apache/comet/rules/CostBasedEngineChoice.scala index 5292c914a05..05b47ff3fb6 100644 --- a/spark/src/main/scala/org/apache/comet/rules/CostBasedEngineChoice.scala +++ b/spark/src/main/scala/org/apache/comet/rules/CostBasedEngineChoice.scala @@ -125,12 +125,17 @@ class EngineCostModel( if (agg.aggregateExpressions.exists(_.mode == Complete)) 1.0 else 0.5 /** Whether `condition` of `join`, run as `plan`, is one validity interval. */ - def validityInterval(condition: Expression, join: SortMergeJoinExec, plan: SparkPlan): Boolean = + def validityInterval( + condition: Expression, + join: SortMergeJoinExec, + plan: SparkPlan): Boolean = { + val sides = if (plan.children.size == 2) plan.children else join.children JoinConditionShape.isValidityInterval( condition, - join.left.outputSet, - join.right.outputSet, + sides.head, + sides(1), JoinConditionShape.aliases(plan.children ++ join.children)) + } private def classTerms(costClass: CostClass, plan: SparkPlan, engine: Engine): Seq[Term] = { val op = sparkOperator(plan) diff --git a/spark/src/main/scala/org/apache/comet/rules/EngineCostTable.scala b/spark/src/main/scala/org/apache/comet/rules/EngineCostTable.scala index 7d1c5eaae7a..f807bcfc18b 100644 --- a/spark/src/main/scala/org/apache/comet/rules/EngineCostTable.scala +++ b/spark/src/main/scala/org/apache/comet/rules/EngineCostTable.scala @@ -313,8 +313,9 @@ object EngineCostTable { * left band joins took 23 to 28 us per output row natively against 6 to 8 in Spark. The * pairs per output row are not estimated, so the line keeps the ratio of 3.5 at a price * between the two, high enough to outweigh `smj` and the conversions around the join at any - * width. A condition that is one validity interval ([[JoinConditionShape]]), where Comet - * costs about what Spark does, adds nothing. + * width. A condition that is one validity interval ([[JoinConditionShape]]) with both + * bounds from one row of one input, where Comet costs about what Spark does, adds nothing; + * bounds from two inputs joined below make a band again. * - `predicate`: a filter over the leaves its predicate reads, Spark noisy (maxrel 0.5 to * 0.8); passing rows costs the filter scalars. * - `projectPassThrough`: a project over every output leaf. Spark copies the row only when diff --git a/spark/src/main/scala/org/apache/comet/rules/JoinConditionShape.scala b/spark/src/main/scala/org/apache/comet/rules/JoinConditionShape.scala index d046b280931..3128473981c 100644 --- a/spark/src/main/scala/org/apache/comet/rules/JoinConditionShape.scala +++ b/spark/src/main/scala/org/apache/comet/rules/JoinConditionShape.scala @@ -21,14 +21,16 @@ package org.apache.comet.rules import scala.collection.mutable -import org.apache.spark.sql.catalyst.expressions.{Add, AddMonths, Alias, And, Attribute, AttributeSet, Cast, Coalesce, DateAdd, DateAddInterval, DateAddYMInterval, DateSub, Expression, ExprId, GreaterThan, GreaterThanOrEqual, IsNull, LessThan, LessThanOrEqual, Or, Subtract, TimeAdd, TimestampAddYMInterval, TruncDate, TruncTimestamp} +import org.apache.spark.sql.catalyst.expressions.{Add, AddMonths, Alias, And, Attribute, AttributeSet, Cast, Coalesce, DateAdd, DateAddInterval, DateAddYMInterval, DateSub, Expression, ExprId, GreaterThan, GreaterThanOrEqual, IsNull, LessThan, LessThanOrEqual, Or, Subtract, TimeAdd, TimestampAddYMInterval, TruncDate, TruncTimestamp, WindowExpression} import org.apache.spark.sql.comet.CometExec -import org.apache.spark.sql.execution.SparkPlan -import org.apache.spark.sql.execution.adaptive.QueryStageExec +import org.apache.spark.sql.execution.{ProjectExec, SparkPlan, UnionExec} +import org.apache.spark.sql.execution.adaptive.{AdaptiveSparkPlanExec, QueryStageExec} +import org.apache.spark.sql.execution.exchange.ReusedExchangeExec /** * Recognizes a join condition that is a single validity interval: `L <= V < U` with `V` from one - * side, `L` and `U` from the other, and `U` not derived from `L` by a constant offset. + * side, `L` and `U` from the other, `U` not derived from `L` by a constant offset, and `L` and + * `U` read from one row of one input of the other side, with no join or union between them. */ object JoinConditionShape { @@ -110,20 +112,80 @@ object JoinConditionShape { strip(e) match { case a: Attribute => aliases.get(a.exprId) match { + case Some(_: WindowExpression) => Set(a.exprId) case Some(child) if depth < 64 => sources(child, aliases, depth + 1) case _ => Set(a.exprId) } case other => other.references.map(_.exprId).toSet } + private def inputs(node: SparkPlan): Seq[SparkPlan] = node match { + case s: QueryStageExec => Seq(s.plan) + case a: AdaptiveSparkPlanExec => Seq(a.executedPlan) + case other => other.children + } + + private def original(node: SparkPlan): SparkPlan = node match { + case c: CometExec => c.originalPlan + case other => other + } + + private def outputs(node: SparkPlan, id: ExprId): Boolean = node.output.exists(_.exprId == id) + + private def producers(node: SparkPlan, id: ExprId, depth: Int): Option[Seq[List[SparkPlan]]] = + if (depth > 256 || !outputs(node, id)) { + None + } else { + node match { + case r: ReusedExchangeExec => + val i = r.output.indexWhere(_.exprId == id) + producers(r.child, r.child.output(i).exprId, depth + 1).map(_.map(r :: _)) + case _ => + inputs(node).find(outputs(_, id)) match { + case Some(child) => producers(child, id, depth + 1).map(_.map(node :: _)) + case None => + val projected = original(node) match { + case p: ProjectExec => + p.projectList.collectFirst { case a: Alias if a.exprId == id => a.child } + case _ => None + } + (projected, inputs(node)) match { + case (Some(e), Seq(child)) if e.references.nonEmpty => + val traced = e.references.toSeq.map(a => producers(child, a.exprId, depth + 1)) + if (traced.forall(_.isDefined)) { + Some(traced.flatMap(_.get).map(node :: _)) + } else { + None + } + case _ => Some(Seq(List(node))) + } + } + } + } + + private def oneRow(side: SparkPlan, refs: AttributeSet): Boolean = { + val traced = refs.toSeq.map(a => producers(side, a.exprId, 0)) + traced.nonEmpty && traced.forall(_.isDefined) && { + val paths = traced.flatMap(_.get) + val common = (0 until paths.map(_.size).min) + .takeWhile(i => paths.forall(_(i) eq paths.head(i))) + .size + val below = paths.map(_.drop(common)) + val heads = below.flatMap(_.headOption) + heads.forall(_ eq heads.head) && + paths.forall(_.forall(n => !original(n).isInstanceOf[UnionExec])) && + below.forall(_.forall(n => inputs(n).size <= 1)) + } + } + /** - * Whether `condition` of a join between `left` and `right` outputs is one validity interval - * whose bounds `aliases` does not derive from one another. + * Whether `condition` of a join between `left` and `right` is one validity interval whose + * bounds `aliases` does not derive from one another and that read one row of one input. */ def isValidityInterval( condition: Expression, - left: AttributeSet, - right: AttributeSet, + left: SparkPlan, + right: SparkPlan, aliases: Map[ExprId, Expression]): Boolean = { conjuncts(condition).map(bound) match { case Seq(Some(first), Some(second)) => @@ -133,12 +195,17 @@ object JoinConditionShape { } shapes.exists { case (v, l, u, nullCheck) => val (vRefs, lRefs, uRefs) = (v.references, l.references, u.references) - def onOtherSides(side: AttributeSet, other: AttributeSet): Boolean = - vRefs.subsetOf(side) && lRefs.subsetOf(other) && uRefs.subsetOf(other) + def onOtherSides(side: SparkPlan, other: SparkPlan): Boolean = + vRefs.subsetOf(side.outputSet) && lRefs.subsetOf(other.outputSet) && + uRefs.subsetOf(other.outputSet) + val boundsSide = + if (onOtherSides(left, right)) Some(right) + else if (onOtherSides(right, left)) Some(left) + else None vRefs.nonEmpty && lRefs.nonEmpty && uRefs.nonEmpty && - (onOtherSides(left, right) || onOtherSides(right, left)) && nullCheck.forall(same(_, u)) && - sources(l, aliases, 0).intersect(sources(u, aliases, 0)).isEmpty + sources(l, aliases, 0).intersect(sources(u, aliases, 0)).isEmpty && + boundsSide.exists(oneRow(_, lRefs ++ uRefs)) } case _ => false } diff --git a/spark/src/test/scala/org/apache/comet/rules/CostBasedEngineChoiceSuite.scala b/spark/src/test/scala/org/apache/comet/rules/CostBasedEngineChoiceSuite.scala index 92e350904b0..7bb0319534b 100644 --- a/spark/src/test/scala/org/apache/comet/rules/CostBasedEngineChoiceSuite.scala +++ b/spark/src/test/scala/org/apache/comet/rules/CostBasedEngineChoiceSuite.scala @@ -968,7 +968,11 @@ class CostBasedEngineChoiceSuite extends CometTestBase { .parquet(s"${dir.getCanonicalPath}/dim") spark.read.parquet(s"${dir.getCanonicalPath}/ev").createOrReplaceTempView("ev") spark.read.parquet(s"${dir.getCanonicalPath}/dim").createOrReplaceTempView("dim") - withTempView("ev", "dim")(f) + sql( + "CREATE OR REPLACE TEMP VIEW rates AS SELECT s AS currency, " + + "CAST(to_date(l) AS timestamp) AS effective_date, " + + "CAST(to_date(u) AS timestamp) AS next_effective_date FROM dim") + withTempView("ev", "dim", "rates")(f) } } @@ -980,7 +984,21 @@ class CostBasedEngineChoiceSuite extends CometTestBase { "to_date(e.v) >= to_date(d.l) AND CAST(e.v AS date) < CAST(d.u AS date)", "date_trunc('day', e.v) >= d.l AND e.v < date_trunc('day', d.u)") + private val intervalQueries = Seq( + "SELECT e.k, e.v, r.effective_date FROM ev e LEFT JOIN rates r ON e.s = r.currency " + + "AND e.v > r.effective_date AND e.v <= r.next_effective_date", + "SELECT e.k, e.v, d.l FROM ev e LEFT JOIN dim d ON e.k = d.k AND e.v >= d.l " + + "AND e.v < CAST(COALESCE(CAST(d.u AS string), '9999-12-31') AS timestamp)", + "WITH p AS (SELECT k, l, LEAD(l) OVER (PARTITION BY k ORDER BY l) AS nl FROM dim) " + + "SELECT e.k, e.v, p.l FROM ev e LEFT JOIN p ON e.k = p.k AND e.v >= p.l " + + "AND e.v < COALESCE(p.nl, timestamp '9999-12-31')") + + private val crossInputBounds = + "SELECT o.k, o.v, g.v AS gv FROM ev o JOIN dim p ON o.k = p.k " + + "LEFT JOIN ev g ON o.k = g.k AND p.completed_dt < g.v AND o.v > g.v" + private val chargedConditions = Seq( + crossInputBounds, "SELECT e.k, e.v, d.l FROM ev e JOIN dim d ON e.k = d.k " + "AND e.vs BETWEEN d.ls - 2592000 AND d.ls", "SELECT e.k, e.v, d.l FROM ev e JOIN dim d ON e.k = d.k " + @@ -1021,6 +1039,9 @@ class CostBasedEngineChoiceSuite extends CometTestBase { s"SELECT e.k, e.v, d.l FROM ev e $joinType dim d ON e.k = d.k AND $condition" assert(conditionTerms(query).forall(_.isEmpty), query) } + for (query <- intervalQueries) { + assert(conditionTerms(query).forall(_.isEmpty), query) + } for (query <- chargedConditions) { assert(conditionTerms(query).exists(_.nonEmpty), query) } @@ -1043,9 +1064,14 @@ class CostBasedEngineChoiceSuite extends CometTestBase { val (intervalOff, intervalOn) = offAndOn(run(interval)) assert(joins(intervalOff) == (1, 0), s"plan:\n$intervalOff") assert(joins(intervalOn) == (1, 0), s"plan:\n$intervalOn") - val (bandOff, bandOn) = offAndOn(run(chargedConditions.head)) + val (bandOff, bandOn) = offAndOn(run(chargedConditions(1))) assert(joins(bandOff) == (1, 0), s"plan:\n$bandOff") assert(joins(bandOn) == (0, 1), s"plan:\n$bandOn") + val (crossOff, crossOn) = offAndOn(run(crossInputBounds)) + assert(joins(crossOff) == (2, 0), s"plan:\n$crossOff") + assert( + count(crossOn) { case j: SortMergeJoinExec if j.condition.isDefined => j } == 1, + s"plan:\n$crossOn") } } } From b5fdf8668859d452c5183e66a8cc937b965231e6 Mon Sep 17 00:00:00 2001 From: msaf Date: Tue, 6 Oct 2026 19:11:08 +0100 Subject: [PATCH 08/13] fix: let a Spark producer feed a root Comet shuffle through the columnar format A Comet shuffle with no consumer in the plan, such as the repartition at the root of a dbt model's write (DISTRIBUTE BY), could only keep its native format, which a Spark producer cannot feed. Every operator of its stage was pinned native, so fbj_order_type's band sort-merge join stayed native whatever its smjCondition price. Such a shuffle may now also take the columnar format: its output is Arrow either way. Co-Authored-By: Claude Opus 5.5 (1M context) --- .../apache/comet/rules/BoundaryFormats.scala | 11 +- .../comet/rules/CostBasedEngineChoice.scala | 3 +- .../rules/CostBasedEngineChoiceSuite.scala | 354 ++++++++++++++++++ 3 files changed, 362 insertions(+), 6 deletions(-) diff --git a/spark/src/main/scala/org/apache/comet/rules/BoundaryFormats.scala b/spark/src/main/scala/org/apache/comet/rules/BoundaryFormats.scala index 5f6b1b2313b..f2725906ddf 100644 --- a/spark/src/main/scala/org/apache/comet/rules/BoundaryFormats.scala +++ b/spark/src/main/scala/org/apache/comet/rules/BoundaryFormats.scala @@ -56,9 +56,11 @@ import org.apache.comet.shims.CometTypeShim * Spark when their keys may hash differently (see [[modes]]). * * Boundaries whose format is fixed: materialized or reused query stages (`QueryStageExec`, - * `ReusedExchangeExec`, `AQEShuffleReadExec`), any other exchange implementation, and a boundary - * with no consumer in the plan (the root of a subquery or query stage, or an exchange directly - * over another exchange). Their consumers still pay the conversion their fixed format implies. + * `ReusedExchangeExec`, `AQEShuffleReadExec`) and any other exchange implementation. Their + * consumers still pay the conversion their fixed format implies. A boundary with no consumer in + * the plan (the root of a subquery or query stage, or an exchange directly over another exchange) + * keeps its output: a broadcast keeps its format, and a Comet shuffle stays a Comet shuffle, + * native or columnar, so that a Spark producer can still feed it. * * Range and round-robin shuffles, and single-partition ones, are never co-partitioned with * another input, so only their conversions count. `shuffleOrigin` and the advisory partition size @@ -363,8 +365,7 @@ object BoundaryFormats extends Logging with CometTypeShim { conversionsInto(input.consumer, NativeShuffle), CometHash)) } - if ((!keep || current == ColumnarShuffle) && current != SparkShuffle && - columnarAvailable(input)) { + if (current != SparkShuffle && columnarAvailable(input)) { val write = if (producerIsComet) 2 else 1 candidates += (( ColumnarShuffle, diff --git a/spark/src/main/scala/org/apache/comet/rules/CostBasedEngineChoice.scala b/spark/src/main/scala/org/apache/comet/rules/CostBasedEngineChoice.scala index 05b47ff3fb6..0b66d06b94d 100644 --- a/spark/src/main/scala/org/apache/comet/rules/CostBasedEngineChoice.scala +++ b/spark/src/main/scala/org/apache/comet/rules/CostBasedEngineChoice.scala @@ -349,7 +349,8 @@ object EngineCostModel { * keep the engine they were converted to (the aggregate test is the one of * `COMET_UNSAFE_PARTIAL` and [[RevertNativeForTransitionHeavyStages]]). * - Materialized and reused stages are leaves of fixed format, and a boundary with no consumer - * in the plan (a subquery or stage root) keeps its format. + * in the plan (a subquery or stage root) keeps its output, Arrow or rows: a Comet shuffle + * there may switch between native and columnar, so its producer may run in Spark. * - The plan's own output is rows, so a native root pays one conversion. The root of a subquery * keeps its engine: an operator outside the plan, such as the broadcast that dynamic * partition pruning builds around it, may rely on it. diff --git a/spark/src/test/scala/org/apache/comet/rules/CostBasedEngineChoiceSuite.scala b/spark/src/test/scala/org/apache/comet/rules/CostBasedEngineChoiceSuite.scala index 7bb0319534b..d9b7f67363f 100644 --- a/spark/src/test/scala/org/apache/comet/rules/CostBasedEngineChoiceSuite.scala +++ b/spark/src/test/scala/org/apache/comet/rules/CostBasedEngineChoiceSuite.scala @@ -1077,6 +1077,360 @@ class CostBasedEngineChoiceSuite extends CometTestBase { } } + private def withFbj(f: => Unit): Unit = { + withTempPath { dir => + spark + .range(4000) + .selectExpr( + "cast(id AS string) AS order_id", + "timestamp_seconds(1735689600 + id * 3600) AS order_datetime_utc", + "concat('m', cast(id % 13 AS string)) AS merchant_id", + "concat('p', cast(id % 41 AS string)) AS product_id", + "concat('v', cast(id % 61 AS string)) AS product_variant_id", + "IF(id % 10 = 0, 'Other', 'Chinese') AS origin_name", + "id % 3 + 1 AS product_quantity", + "id % 4 = 0 AS is_fbj", + "cast(date_trunc('month', timestamp_seconds(1735689600 + id * 3600)) AS date) " + + "AS order_month_msk") + .write + .parquet(s"${dir.getCanonicalPath}/orders") + spark + .range(600) + .selectExpr( + "concat('v', cast(id % 61 AS string)) AS variant_id", + "concat('g', cast(id DIV 2 AS string)) AS replenishment_group_id", + "concat('g', cast(id DIV 2 + id % 2 * 100000 AS string)) AS replenishment_id", + "IF(id % 7 = 0, 'Merchant', 'Joom') AS source", + "cast(timestamp_seconds(1735689600 + id * 20000) AS date) AS partition_date", + "timestamp_seconds(1735689600 + id * 20000) AS created_at", + "'completed' AS current_status", + "timestamp_seconds(1735689600 + id * 20000 + 3600) AS 2_pending_inbound_dt", + "timestamp_seconds(1735689600 + id * 20000 + 7200) AS 3_pending_shipping_dt", + "timestamp_seconds(1735689600 + id * 20000 + 10800) AS 4_shipped_dt", + "timestamp_seconds(1735689600 + id * 20000 + 14400) AS 5_action_required_dt", + "timestamp_seconds(1735689600 + id * 20000 + 18000) AS 6_on_review_dt", + "IF(id % 5 = 0, NULL, timestamp_seconds(1735689600 + id * 20000 + 604800)) " + + "AS completed_dt", + "id % 20 + 5 AS min_count", + "id % 20 + 50 AS max_count", + "id % 30 AS requested_count", + "id % 25 AS accepted_count") + .write + .parquet(s"${dir.getCanonicalPath}/repl") + spark + .range(200) + .selectExpr( + "cast(date_trunc('quarter', timestamp_seconds(1735689600 + (id % 4) * 7776000)) " + + "AS date) AS quarter", + "concat('m', cast(id % 13 AS string)) AS merchant_id", + "concat('name', cast(id AS string)) AS merchant_name", + "concat('main', cast(id AS string)) AS main_merchant_name", + "concat('kam', cast(id % 3 AS string)) AS kam_email") + .write + .parquet(s"${dir.getCanonicalPath}/kam") + spark.read.parquet(s"${dir.getCanonicalPath}/orders").createOrReplaceTempView("gold_orders") + spark.read.parquet(s"${dir.getCanonicalPath}/repl").createOrReplaceTempView("fbj_repl") + spark.read.parquet(s"${dir.getCanonicalPath}/kam").createOrReplaceTempView("fbj_kam") + withTempView("gold_orders", "fbj_repl", "fbj_kam")(f) + } + } + + private val fbjOrderType = + """ + |WITH orders AS ( + | SELECT + | order_datetime_utc, + | CAST(order_datetime_utc AS DATE) AS dt, + | order_id, + | product_variant_id, + | product_quantity, + | is_fbj, + | product_id, + | merchant_id + | FROM gold_orders + | WHERE + | origin_name = 'Chinese' + | AND order_datetime_utc >= '2025-01-01' + | AND order_month_msk >= '2024-12-01' + |), + | + |variant_daily_demand AS ( + | SELECT + | product_variant_id, + | dt, + | SUM(product_quantity) AS daily_qty + | FROM orders + | GROUP BY product_variant_id, dt + |), + | + |forward_demand AS ( + | SELECT + | a.product_variant_id, + | a.dt, + | SUM(f.daily_qty) AS forward_14d_qty + | FROM variant_daily_demand AS a + | INNER JOIN variant_daily_demand AS f + | ON + | a.product_variant_id = f.product_variant_id + | AND a.dt <= f.dt + | AND f.dt <= DATE_ADD(a.dt, 13) + | GROUP BY a.product_variant_id, a.dt + |), + | + |parent_repl AS ( + | SELECT + | variant_id, + | replenishment_group_id, + | partition_date, + | created_at, + | current_status, + | 2_pending_inbound_dt, + | 3_pending_shipping_dt, + | 4_shipped_dt, + | 5_action_required_dt, + | 6_on_review_dt, + | completed_dt, + | min_count, + | max_count, + | requested_count + | FROM fbj_repl + | WHERE + | source IN ('Joom', 'Warehouse') + | AND replenishment_id = replenishment_group_id + |), + | + |parent_repl_amount AS ( + | SELECT + | parent_repl.replenishment_group_id, + | parent_repl.min_count, + | SUM(r.accepted_count) AS acc_count + | FROM parent_repl + | INNER JOIN fbj_repl AS r + | ON parent_repl.replenishment_group_id = r.replenishment_group_id + | GROUP BY + | parent_repl.replenishment_group_id, + | parent_repl.min_count + |), + | + |current_active_repl AS ( + | SELECT + | o.order_datetime_utc, + | o.dt, + | o.order_id, + | o.product_variant_id, + | pr.replenishment_group_id, + | pr.created_at, + | COALESCE(pr.completed_dt, '5999-12-31 23:59:59') AS completed_dt, + | ROW_NUMBER() OVER (PARTITION BY o.order_id ORDER BY pr.created_at DESC) AS rn + | FROM orders AS o + | LEFT JOIN parent_repl AS pr + | ON + | o.product_variant_id = pr.variant_id + | AND o.order_datetime_utc > pr.created_at + | AND COALESCE(pr.completed_dt, '5999-12-31 23:59:59') > o.order_datetime_utc + |), + | + |previous_completed_repl AS ( + | SELECT + | o.order_datetime_utc, + | o.dt, + | o.order_id, + | o.product_variant_id, + | pr.replenishment_group_id, + | pr.created_at, + | COALESCE(pr.completed_dt, '5999-12-31 23:59:59') AS completed_dt, + | pra.min_count, + | pra.acc_count, + | ROW_NUMBER() OVER ( + | PARTITION BY o.order_id + | ORDER BY COALESCE(pr.completed_dt, '5999-12-31 23:59:59') DESC + | ) AS rn + | FROM orders AS o + | LEFT JOIN parent_repl AS pr + | ON + | o.product_variant_id = pr.variant_id + | AND o.order_datetime_utc > pr.created_at + | AND COALESCE(pr.completed_dt, '5999-12-31 23:59:59') <= o.order_datetime_utc + | LEFT JOIN parent_repl_amount AS pra + | ON pr.replenishment_group_id = pra.replenishment_group_id + |), + | + |var_kam AS ( + | SELECT + | mkm.quarter, + | mkm.merchant_id, + | MAX(mkm.merchant_name) AS merchant_name, + | MAX(mkm.main_merchant_name) AS main_merchant_name, + | MAX(mkm.kam_email) AS kam_email + | FROM fbj_kam AS mkm + | GROUP BY + | mkm.quarter, + | mkm.merchant_id + |), + | + |final AS ( + | SELECT /*+ BROADCAST(m_k) */ + | o.order_datetime_utc, + | o.dt, + | o.order_id, + | o.product_variant_id, + | o.product_quantity, + | o.is_fbj, + | o.product_id, + | o.merchant_id, + | m_k.kam_email AS kam, + | fd.forward_14d_qty, + | car.replenishment_group_id AS c_replenishment_group_id, + | car.created_at AS c_created_at, + | car.completed_dt AS c_completed_dt, + | pcr.replenishment_group_id AS p_replenishment_group_id, + | pcr.created_at AS p_created_at, + | pcr.completed_dt AS p_completed_dt, + | pcr.min_count AS p_min_count, + | COALESCE(pcr.acc_count, 0) AS p_acc_count, + | COALESCE(SUM(go.product_quantity), 0) AS p_product_qnt + | FROM orders AS o + | LEFT JOIN forward_demand AS fd + | ON + | o.product_variant_id = fd.product_variant_id + | AND o.dt = fd.dt + | LEFT JOIN current_active_repl AS car + | ON + | o.order_id = car.order_id + | AND car.rn = 1 + | LEFT JOIN previous_completed_repl AS pcr + | ON + | o.order_id = pcr.order_id + | AND pcr.rn = 1 + | LEFT JOIN var_kam AS m_k + | ON + | o.merchant_id = m_k.merchant_id + | AND m_k.quarter = DATE_TRUNC('quarter', o.dt) + | LEFT JOIN gold_orders AS go + | ON + | o.product_variant_id = go.product_variant_id + | AND pcr.completed_dt < go.order_datetime_utc + | AND o.order_datetime_utc > go.order_datetime_utc + | GROUP BY + | o.order_datetime_utc, + | o.dt, + | o.order_id, + | o.product_variant_id, + | o.product_quantity, + | o.is_fbj, + | o.product_id, + | o.merchant_id, + | m_k.kam_email, + | fd.forward_14d_qty, + | car.replenishment_group_id, + | car.created_at, + | car.completed_dt, + | pcr.replenishment_group_id, + | pcr.created_at, + | pcr.completed_dt, + | pcr.min_count, + | COALESCE(pcr.acc_count, 0) + |) + | + |SELECT + | dt AS partition_date, + | order_datetime_utc, + | order_id, + | product_variant_id, + | product_quantity, + | is_fbj, + | product_id, + | merchant_id, + | kam, + | forward_14d_qty, + | forward_14d_qty > 4 AS is_eligible, + | CASE + | WHEN is_fbj AND forward_14d_qty > 4 + | THEN '1.1 FBJ | Eligible' + | WHEN is_fbj AND forward_14d_qty <= 4 + | THEN '1.2 FBJ | Not Eligible' + | WHEN NOT is_fbj AND forward_14d_qty <= 4 + | THEN '2.1 FBM | Not Eligible' + | WHEN + | NOT is_fbj AND forward_14d_qty > 4 + | AND p_replenishment_group_id IS NULL + | THEN '2.2.1 FBM Eligible | No previous replenishment' + | WHEN + | NOT is_fbj AND forward_14d_qty > 4 + | AND p_replenishment_group_id IS NOT NULL + | AND c_replenishment_group_id IS NULL + | THEN '2.2.2 FBM Eligible | Prev repl completed, no active repl' + | WHEN + | NOT is_fbj AND forward_14d_qty > 4 + | AND p_replenishment_group_id IS NOT NULL + | AND c_replenishment_group_id IS NOT NULL + | AND p_acc_count >= p_min_count + | THEN '2.2.3 FBM Eligible | Fully delivered, late new repl' + | WHEN + | NOT is_fbj AND forward_14d_qty > 4 + | AND p_replenishment_group_id IS NOT NULL + | AND c_replenishment_group_id IS NOT NULL + | AND p_acc_count < p_min_count + | AND p_product_qnt >= p_min_count - p_acc_count + | THEN '2.2.4 FBM Eligible | Partially delivered, demand covered' + | WHEN + | NOT is_fbj AND forward_14d_qty > 4 + | AND p_replenishment_group_id IS NOT NULL + | AND c_replenishment_group_id IS NOT NULL + | AND p_acc_count < p_min_count + | AND p_product_qnt < p_min_count - p_acc_count + | THEN '3.0 FBM Eligible | Partially delivered, demand not covered' + | END AS bucket, + | c_replenishment_group_id, + | c_created_at, + | c_completed_dt, + | p_replenishment_group_id, + | p_created_at, + | p_completed_dt, + | p_min_count, + | p_acc_count, + | p_product_qnt + |FROM final + |DISTRIBUTE BY partition_date + |""".stripMargin + + for (aqe <- Seq("false", "true")) { + test(s"fbj_order_type's band join under a root repartition runs in Spark (AQE=$aqe)") { + withFbj { + withAqe( + aqe, + SQLConf.AUTO_BROADCASTJOIN_THRESHOLD.key -> "-1", + SQLConf.ADAPTIVE_AUTO_BROADCASTJOIN_THRESHOLD.key -> "-1") { + def conditions(plan: SparkPlan): (Seq[String], Seq[String]) = { + val all = nodes(plan).collect { + case j: CometSortMergeJoinExec => + (true, j.originalPlan.asInstanceOf[SortMergeJoinExec].condition) + case j: SortMergeJoinExec => (false, j.condition) + } + def named(native: Boolean) = + all.collect { case (`native`, Some(c)) => c.references.map(_.name).toSeq.sorted } + def kind(refs: Seq[String]): String = + if (refs.contains("completed_dt") && refs.count(_ == "order_datetime_utc") == 2) { + "band" + } else if (refs.contains("created_at")) { + "replenishment" + } else { + "forward" + } + (named(true).map(kind).sorted, named(false).map(kind).sorted) + } + val (off, on) = offAndOn(run(fbjOrderType)) + assert( + conditions(off) == (Seq("band", "forward", "replenishment", "replenishment"), Nil), + s"plan:\n$off") + val (native, spark) = conditions(on) + assert(spark.contains("band"), s"plan:\n$on") + assert(native.contains("replenishment"), s"plan:\n$on") + } + } + } + } + test("a window costs its line and the classes of its functions") { withTables { withSQLConf(SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "false", flag -> "false") { From c7df21932c1ab16eda70e4ad81cf5b73f958c86d Mon Sep 17 00:00:00 2001 From: msaf Date: Wed, 7 Oct 2026 10:30:47 +0100 Subject: [PATCH 09/13] fix: recognise a constant time shift without referencing TimeAdd, which Spark 4.1 renamed Co-Authored-By: Claude Opus 5.5 (1M context) --- .../org/apache/comet/rules/JoinConditionShape.scala | 9 +++++++-- 1 file changed, 7 insertions(+), 2 deletions(-) diff --git a/spark/src/main/scala/org/apache/comet/rules/JoinConditionShape.scala b/spark/src/main/scala/org/apache/comet/rules/JoinConditionShape.scala index 3128473981c..de0d07eff9d 100644 --- a/spark/src/main/scala/org/apache/comet/rules/JoinConditionShape.scala +++ b/spark/src/main/scala/org/apache/comet/rules/JoinConditionShape.scala @@ -21,7 +21,7 @@ package org.apache.comet.rules import scala.collection.mutable -import org.apache.spark.sql.catalyst.expressions.{Add, AddMonths, Alias, And, Attribute, AttributeSet, Cast, Coalesce, DateAdd, DateAddInterval, DateAddYMInterval, DateSub, Expression, ExprId, GreaterThan, GreaterThanOrEqual, IsNull, LessThan, LessThanOrEqual, Or, Subtract, TimeAdd, TimestampAddYMInterval, TruncDate, TruncTimestamp, WindowExpression} +import org.apache.spark.sql.catalyst.expressions.{Add, AddMonths, Alias, And, Attribute, AttributeSet, Cast, Coalesce, DateAdd, DateAddInterval, DateAddYMInterval, DateSub, Expression, ExprId, GreaterThan, GreaterThanOrEqual, IsNull, LessThan, LessThanOrEqual, Or, Subtract, TimestampAddYMInterval, TruncDate, TruncTimestamp, WindowExpression} import org.apache.spark.sql.comet.CometExec import org.apache.spark.sql.execution.{ProjectExec, SparkPlan, UnionExec} import org.apache.spark.sql.execution.adaptive.{AdaptiveSparkPlanExec, QueryStageExec} @@ -45,13 +45,18 @@ object JoinConditionShape { case other => other } + private val timeAddClasses = Set("TimeAdd", "TimestampAddInterval") + private def offsetBase(e: Expression): Option[Expression] = e match { case a: Add if a.right.foldable => Some(a.left) case a: Add if a.left.foldable => Some(a.right) case s: Subtract if s.right.foldable => Some(s.left) case d: DateAdd if d.days.foldable => Some(d.startDate) case d: DateSub if d.days.foldable => Some(d.startDate) - case t: TimeAdd if t.interval.foldable => Some(t.start) + case t + if timeAddClasses(t.getClass.getSimpleName) && t.children.size >= 2 && + t.children(1).foldable => + Some(t.children.head) case d: DateAddInterval if d.interval.foldable => Some(d.start) case d: DateAddYMInterval if d.interval.foldable => Some(d.date) case t: TimestampAddYMInterval if t.interval.foldable => Some(t.timestamp) From 6975ebae0cfbd6e9879f1a58ab177de62bacb907 Mon Sep 17 00:00:00 2001 From: msaf Date: Wed, 7 Oct 2026 13:36:39 +0100 Subject: [PATCH 10/13] perf: gather a wide native sort's payload from concatenated columns The late-materialized sort interleaved every output batch from all the batches it buffered, ~1600 batches of 8192 rows at 13M rows per task, column by column. It now concatenates each payload column of the buffered batches once, releasing the source column as soon as it is copied, and takes every output batch from it. The memory held grows by at most one column while it is copied. A column is concatenated whole when its reservation grows by the column's size, in up to 16 chunks otherwise, and is left in its batches when not even a chunk fits; an output batch then interleaves the few pieces its rows come from. A spill concatenates in chunks of at most 1/16 of what it buffered, since its reservation cannot refuse. Columns whose concatenation or take is unsupported, or whose 32-bit list offsets would overflow, stay in their batches. An order that reads at most 16 batches per output batch on average (input already nearly sorted) keeps interleaving and releases each batch once its last row is output. SortSpillBenchSuite, 22 leaf columns (string, date, long, double and boolean, mostly null), 4 keys, one task, in memory, end to end ns/row, alternating runs, before -> after: 6M rows 505/465 -> 431/406; 13M rows 1166/556 -> 484/498 (the old gather ran at either speed per JVM). Native sort time at 13M: 12.3/6.9 -> 5.2/5.6 s. Co-Authored-By: Claude Opus 5.5 (1M context) --- .../src/sorts/sort.rs | 1 + .../src/sorts/sort/late_materialize.rs | 960 ++++++++++++++++-- 2 files changed, 879 insertions(+), 82 deletions(-) diff --git a/native/vendor/datafusion-physical-plan/src/sorts/sort.rs b/native/vendor/datafusion-physical-plan/src/sorts/sort.rs index 82d55594e5c..51514d63b94 100644 --- a/native/vendor/datafusion-physical-plan/src/sorts/sort.rs +++ b/native/vendor/datafusion-physical-plan/src/sorts/sort.rs @@ -844,6 +844,7 @@ impl ExternalSorter { std::mem::take(&mut self.in_mem_batches), self.expr.clone(), rows_per_batch, + !is_output_stream, self.reservation.take(), elapsed_compute, ); diff --git a/native/vendor/datafusion-physical-plan/src/sorts/sort/late_materialize.rs b/native/vendor/datafusion-physical-plan/src/sorts/sort/late_materialize.rs index 04ecb7420f6..a8d5681c7cb 100644 --- a/native/vendor/datafusion-physical-plan/src/sorts/sort/late_materialize.rs +++ b/native/vendor/datafusion-physical-plan/src/sorts/sort/late_materialize.rs @@ -17,15 +17,15 @@ //! COMET PATCH: sort the keys of the buffered batches and gather each payload once. +use std::ops::Range; use std::sync::Arc; use arrow::array::{ - Array, ArrayData, ArrayRef, RecordBatch, RecordBatchOptions, UInt32Array, + Array, ArrayData, ArrayRef, AsArray, RecordBatch, RecordBatchOptions, UInt32Array, + new_empty_array, }; -use arrow::compute::{ - SortColumn, concat, interleave, lexsort_to_indices, take_record_batch, -}; -use arrow::datatypes::SchemaRef; +use arrow::compute::{SortColumn, concat, interleave, lexsort_to_indices, take}; +use arrow::datatypes::{DataType, SchemaRef}; use arrow::row::{RowConverter, Rows, SortField}; use datafusion_common::HashMap; use datafusion_common::Result; @@ -45,6 +45,10 @@ const ORDER_BYTES_PER_ROW: usize = 48; const OUTPUT_BATCH_BYTES: usize = 4 << 20; const MIN_SPILL_BATCH_BYTES: usize = 16 << 10; const SPILL_BATCHES_PER_RUN: usize = 64; +const LOCAL_SOURCES: usize = 16; +const SAMPLED_CHUNKS: usize = 16; +const SPILL_COLUMN_SHARE: usize = 16; +const CHUNKS_PER_COLUMN: usize = 16; #[derive(Debug)] pub(super) struct LateMaterialization { @@ -130,6 +134,7 @@ impl LateMaterialization { batches: Vec, ordering: LexOrdering, rows_per_batch: usize, + spilling: bool, reservation: MemoryReservation, elapsed_compute: Time, ) -> SendableRecordBatchStream { @@ -143,6 +148,7 @@ impl LateMaterialization { batches, &ordering, rows_per_batch, + spilling, reservation, elapsed_compute.clone(), )? @@ -229,13 +235,14 @@ fn row_order(rows: &Rows) -> Vec { struct Gather { schema: SchemaRef, - batches: Vec, - starts: Vec, + columns: Vec, + layouts: Vec, + holders: Vec>>, + gathered_bytes: usize, + local: bool, remaining: Vec, - buffers: Vec>, owners: HashMap, live_bytes: usize, - slots: Vec, order: UInt32Array, cursor: usize, rows_per_batch: usize, @@ -243,6 +250,79 @@ struct Gather { elapsed_compute: Time, } +struct Pieces { + arrays: Vec, + starts: Vec, + batches: Vec>, + layout: usize, +} + +struct Layout { + starts: Vec, + slots: Vec, + used: Vec, + indices: Vec<(usize, usize)>, + local: Option, +} + +impl Layout { + fn new(starts: Vec) -> Self { + Self { + slots: vec![usize::MAX; starts.len()], + starts, + used: vec![], + indices: vec![], + local: None, + } + } + + fn piece_of(starts: &[usize], row: usize) -> usize { + starts.partition_point(|&start| start <= row) - 1 + } + + fn map(&mut self, order: &UInt32Array) { + self.used.clear(); + self.indices.clear(); + self.local = None; + if self.starts.len() == 1 { + return; + } + for &row in order.values() { + let row = row as usize; + let piece = Self::piece_of(&self.starts, row); + if self.slots[piece] == usize::MAX { + self.slots[piece] = self.used.len(); + self.used.push(piece); + } + self.indices + .push((self.slots[piece], row - self.starts[piece])); + } + for &piece in &self.used { + self.slots[piece] = usize::MAX; + } + if self.used.len() == 1 { + self.local = Some(UInt32Array::from_iter_values( + self.indices.iter().map(|&(_, row)| row as u32), + )); + } + } + + fn gather(&self, arrays: &[ArrayRef], order: &UInt32Array) -> Result { + if self.starts.len() == 1 { + return Ok(take(&arrays[0], order, None)?); + } + if let Some(local) = &self.local { + return Ok(take(&arrays[self.used[0]], local, None)?); + } + let used: Vec<&dyn Array> = self + .used + .iter() + .map(|&piece| arrays[piece].as_ref()) + .collect(); + Ok(interleave(&used, &self.indices)?) + } +} + fn collect_buffers(data: &ArrayData, buffers: &mut Vec<(usize, usize)>) { for buffer in data.buffers() { buffers.push((buffer.data_ptr().as_ptr() as usize, buffer.capacity())); @@ -256,46 +336,137 @@ fn collect_buffers(data: &ArrayData, buffers: &mut Vec<(usize, usize)>) { } } +fn sliced_bytes(array: &dyn Array) -> Result { + let mut bytes = array.to_data().get_slice_memory_size()?; + match array.data_type() { + DataType::Utf8View => { + bytes += array + .as_string_view() + .data_buffers() + .iter() + .map(|b| b.len()) + .sum::() + } + DataType::BinaryView => { + bytes += array + .as_binary_view() + .data_buffers() + .iter() + .map(|b| b.len()) + .sum::() + } + _ => {} + } + Ok(bytes) +} + +fn offsets_fit(arrays: &[ArrayData]) -> bool { + match arrays[0].data_type() { + DataType::List(_) | DataType::Map(_, _) => { + let mut total = 0usize; + let mut children = Vec::with_capacity(arrays.len()); + for data in arrays { + let offsets = + &data.buffer::(0)[data.offset()..=data.offset() + data.len()]; + let start = offsets[0] as usize; + let end = offsets[data.len()] as usize; + total += end - start; + children.push(data.child_data()[0].slice(start, end - start)); + } + total <= i32::MAX as usize && offsets_fit(&children) + } + DataType::LargeList(_) => { + let mut children = Vec::with_capacity(arrays.len()); + for data in arrays { + let offsets = + &data.buffer::(0)[data.offset()..=data.offset() + data.len()]; + let start = offsets[0] as usize; + let end = offsets[data.len()] as usize; + children.push(data.child_data()[0].slice(start, end - start)); + } + offsets_fit(&children) + } + DataType::FixedSizeList(_, size) => { + let size = *size as usize; + let children: Vec = arrays + .iter() + .map(|data| { + data.child_data()[0].slice(data.offset() * size, data.len() * size) + }) + .collect(); + offsets_fit(&children) + } + DataType::Struct(fields) => (0..fields.len()).all(|field| { + let children: Vec = arrays + .iter() + .map(|data| data.child_data()[field].slice(data.offset(), data.len())) + .collect(); + offsets_fit(&children) + }), + DataType::ListView(_) + | DataType::LargeListView(_) + | DataType::Union(_, _) + | DataType::RunEndEncoded(_, _) => false, + _ => true, + } +} + impl Gather { fn try_new( schema: SchemaRef, batches: Vec, ordering: &LexOrdering, rows_per_batch: usize, + spilling: bool, reservation: MemoryReservation, elapsed_compute: Time, ) -> Result { let order = sort_order(&batches, ordering)?; + let width = schema.fields().len(); let mut starts = Vec::with_capacity(batches.len()); - let mut buffers = Vec::with_capacity(batches.len()); + let mut holders = Vec::with_capacity(batches.len()); let mut owners: HashMap = HashMap::new(); let mut live_bytes = 0; let mut rows = 0; for batch in &batches { starts.push(rows); rows += batch.num_rows(); - let mut found = vec![]; + let mut held = Vec::with_capacity(width); for column in batch.columns() { + let mut found = vec![]; collect_buffers(&column.to_data(), &mut found); + found.sort_unstable(); + found.dedup_by_key(|(ptr, _)| *ptr); + for &(ptr, capacity) in &found { + let owner = owners.entry(ptr).or_insert_with(|| { + live_bytes += capacity; + (0, capacity) + }); + owner.0 += 1; + } + held.push(found.into_iter().map(|(ptr, _)| ptr).collect()); } - found.sort_unstable(); - found.dedup_by_key(|(ptr, _)| *ptr); - for &(ptr, capacity) in &found { - let owner = owners.entry(ptr).or_insert_with(|| { - live_bytes += capacity; - (0, capacity) - }); - owner.0 += 1; - } - buffers.push(found.into_iter().map(|(ptr, _)| ptr).collect()); + holders.push(held); } + let columns = (0..width) + .map(|column| Pieces { + arrays: batches + .iter() + .map(|batch| Arc::clone(batch.column(column))) + .collect(), + starts: starts.clone(), + batches: (0..batches.len()).map(|batch| batch..batch + 1).collect(), + layout: 0, + }) + .collect(); let mut gather = Self { schema, + columns, + layouts: vec![], + holders, + gathered_bytes: 0, + local: true, remaining: batches.iter().map(RecordBatch::num_rows).collect(), - slots: vec![usize::MAX; batches.len()], - batches, - starts, - buffers, owners, live_bytes, order, @@ -304,20 +475,176 @@ impl Gather { reservation, elapsed_compute, }; - gather.shrink(); + drop(batches); + gather.fit(); + if starts.len() > 1 && !gather.scattered_sources_fit_locally(&starts) { + gather.local = false; + let buffered = gather.live_bytes; + for column in 0..width { + if spilling { + gather.concatenate(column, buffered / SPILL_COLUMN_SHARE)?; + } else { + gather.concatenate(column, usize::MAX)?; + let total = gather.column_bytes(column)?; + gather.concatenate(column, total / CHUNKS_PER_COLUMN)?; + } + } + } + let mut layouts: Vec = vec![]; + for column in gather.columns.iter_mut() { + column.layout = match layouts + .iter() + .position(|layout| layout.starts == column.starts) + { + Some(layout) => layout, + None => { + layouts.push(Layout::new(column.starts.clone())); + layouts.len() - 1 + } + }; + } + gather.layouts = layouts; Ok(gather) } - fn shrink(&mut self) { - let needed = self.live_bytes + self.order.get_array_memory_size(); - if self.reservation.size() > needed { - self.reservation.shrink(self.reservation.size() - needed); + fn needed(&self) -> usize { + self.live_bytes + self.gathered_bytes + self.order.get_array_memory_size() + } + + fn fit(&mut self) { + let needed = self.needed(); + let size = self.reservation.size(); + if size > needed { + self.reservation.shrink(size - needed); + } else if size < needed { + self.reservation.grow(needed - size); } } - fn finish(&mut self, batch: usize) { - self.batches[batch] = RecordBatch::new_empty(Arc::clone(&self.schema)); - for ptr in std::mem::take(&mut self.buffers[batch]) { + fn scattered_sources_fit_locally(&self, starts: &[usize]) -> bool { + let chunks = self.order.len().div_ceil(self.rows_per_batch); + let sampled = chunks.min(SAMPLED_CHUNKS); + let mut slots = vec![false; starts.len()]; + let mut used = vec![]; + let mut sources = 0; + for sample in 0..sampled { + let start = sample * chunks / sampled * self.rows_per_batch; + let end = (start + self.rows_per_batch).min(self.order.len()); + for &row in &self.order.values()[start..end] { + let batch = Layout::piece_of(starts, row as usize); + if !slots[batch] { + slots[batch] = true; + used.push(batch); + } + } + sources += used.len(); + for batch in used.drain(..) { + slots[batch] = false; + } + } + sources <= LOCAL_SOURCES * sampled + } + + fn column_bytes(&self, column: usize) -> Result { + let mut bytes = 0; + for array in &self.columns[column].arrays { + bytes += sliced_bytes(array.as_ref())?; + } + Ok(bytes) + } + + fn concatenate(&mut self, column: usize, budget: usize) -> Result<()> { + let mut sizes = Vec::with_capacity(self.columns[column].arrays.len()); + for array in &self.columns[column].arrays { + sizes.push(sliced_bytes(array.as_ref())?); + } + let mut groups = vec![]; + let mut start = 0; + while start < sizes.len() { + let mut end = start + 1; + let mut bytes = sizes[start]; + while end < sizes.len() && bytes + sizes[end] <= budget { + bytes += sizes[end]; + end += 1; + } + groups.push((start..end, bytes)); + start = end; + } + if groups.len() == sizes.len() { + return Ok(()); + } + let old = std::mem::replace( + &mut self.columns[column], + Pieces { + arrays: vec![], + starts: vec![], + batches: vec![], + layout: 0, + }, + ); + let mut pieces = Pieces { + arrays: Vec::with_capacity(groups.len()), + starts: Vec::with_capacity(groups.len()), + batches: Vec::with_capacity(groups.len()), + layout: 0, + }; + for (group, bytes) in groups { + let merged = if group.len() > 1 { + self.merge(&old.arrays[group.clone()], bytes) + } else { + None + }; + match merged { + Some(array) => { + let batches = + old.batches[group.start].start..old.batches[group.end - 1].end; + self.gathered_bytes += array.get_array_memory_size(); + for batch in batches.clone() { + self.release(batch, column); + } + self.fit(); + pieces.arrays.push(array); + pieces.starts.push(old.starts[group.start]); + pieces.batches.push(batches); + } + None => { + for piece in group { + pieces.arrays.push(Arc::clone(&old.arrays[piece])); + pieces.starts.push(old.starts[piece]); + pieces.batches.push(old.batches[piece].clone()); + } + } + } + } + drop(old); + self.columns[column] = pieces; + self.fit(); + Ok(()) + } + + fn merge(&mut self, arrays: &[ArrayRef], bytes: usize) -> Option { + let data: Vec = arrays.iter().map(|array| array.to_data()).collect(); + if !offsets_fit(&data) { + return None; + } + drop(data); + let needed = self.needed() + bytes; + let size = self.reservation.size(); + if size < needed && self.reservation.try_grow(needed - size).is_err() { + return None; + } + let arrays: Vec<&dyn Array> = arrays.iter().map(|array| array.as_ref()).collect(); + let merged = concat(&arrays) + .ok() + .filter(|array| take(array, &UInt32Array::from(vec![0u32]), None).is_ok()); + if merged.is_none() { + self.fit(); + } + merged + } + + fn release(&mut self, batch: usize, column: usize) { + for ptr in std::mem::take(&mut self.holders[batch][column]) { if let Some(owner) = self.owners.get_mut(&ptr) { owner.0 -= 1; if owner.0 == 0 { @@ -328,67 +655,69 @@ impl Gather { } } + fn finish(&mut self, batch: usize) { + for column in 0..self.columns.len() { + let pieces = &mut self.columns[column]; + pieces.arrays[batch] = new_empty_array(pieces.arrays[batch].data_type()); + self.release(batch, column); + } + } + fn next_batch(&mut self) -> Result { let elapsed_compute = self.elapsed_compute.clone(); let _timer = elapsed_compute.timer(); let end = (self.cursor + self.rows_per_batch).min(self.order.len()); let order = self.order.slice(self.cursor, end - self.cursor); self.cursor = end; + for layout in self.layouts.iter_mut() { + layout.map(&order); + } + let columns = self + .columns + .iter() + .map(|pieces| self.layouts[pieces.layout].gather(&pieces.arrays, &order)) + .collect::>>()?; let mut finished = vec![]; - let batch = if self.batches.len() == 1 { - let batch = take_record_batch(&self.batches[0], &order)?; - self.remaining[0] -= order.len(); - if self.remaining[0] == 0 { - finished.push(0); - } - batch - } else { - let mut used = vec![]; - let indices: Vec<(usize, usize)> = order - .values() - .iter() - .map(|&row| { - let row = row as usize; - let batch = self.starts.partition_point(|&start| start <= row) - 1; - if self.slots[batch] == usize::MAX { - self.slots[batch] = used.len(); - used.push(batch); + if self.local { + let layout = &self.layouts[0]; + if layout.starts.len() == 1 { + self.remaining[0] -= order.len(); + if self.remaining[0] == 0 { + finished.push(0); + } + } else { + for &(slot, _) in &layout.indices { + let batch = layout.used[slot]; + self.remaining[batch] -= 1; + if self.remaining[batch] == 0 { + finished.push(batch); } - (self.slots[batch], row - self.starts[batch]) - }) - .collect(); - let columns = (0..self.schema.fields().len()) - .map(|column| { - let arrays: Vec<&dyn Array> = used - .iter() - .map(|&batch| self.batches[batch].column(column).as_ref()) - .collect(); - interleave(&arrays, &indices) - }) - .collect::, _>>()?; - for &batch in &used { - self.slots[batch] = usize::MAX; - } - for &(slot, _) in &indices { - let batch = used[slot]; - self.remaining[batch] -= 1; - if self.remaining[batch] == 0 { - finished.push(batch); } } - RecordBatch::try_new_with_options( - Arc::clone(&self.schema), - columns, - &RecordBatchOptions::new().with_row_count(Some(indices.len())), - )? - }; - if !finished.is_empty() { - for batch in finished { - self.finish(batch); + } + for batch in &finished { + self.finish(*batch); + } + let done = self.cursor == self.order.len(); + if done { + for pieces in self.columns.iter_mut() { + pieces.arrays.clear(); + } + for batch in 0..self.holders.len() { + for column in 0..self.columns.len() { + self.release(batch, column); + } } - self.shrink(); + self.gathered_bytes = 0; + } + if !finished.is_empty() || done { + self.fit(); } - Ok(batch) + Ok(RecordBatch::try_new_with_options( + Arc::clone(&self.schema), + columns, + &RecordBatchOptions::new().with_row_count(Some(order.len())), + )?) } } @@ -411,6 +740,7 @@ mod tests { use arrow::array::{ BinaryArray, DictionaryArray, Int32Array, ListArray, StringArray, StringViewArray, }; + use arrow::compute::take_record_batch; use arrow::compute::{SortOptions, concat_batches}; use arrow::datatypes::{DataType, Field, Int32Type, Schema}; use datafusion_common::config::SpillCompression; @@ -870,4 +1200,470 @@ mod tests { assert_eq!(pool.reserved(), 0); Ok(()) } + + fn mixed_schema() -> SchemaRef { + let item = Arc::new(Field::new_list_field(DataType::Int32, true)); + Arc::new(Schema::new(vec![ + Field::new("k", DataType::Int32, true), + Field::new("s", DataType::Utf8, true), + Field::new("flag", DataType::Boolean, true), + Field::new("day", DataType::Date32, true), + Field::new("amount", DataType::Float64, true), + Field::new("count", DataType::Int64, true), + Field::new("name", DataType::Utf8, true), + Field::new("blob", DataType::Binary, true), + Field::new( + "dict", + DataType::Dictionary(Box::new(DataType::Int8), Box::new(DataType::Utf8)), + true, + ), + Field::new("list", DataType::List(Arc::clone(&item)), true), + Field::new( + "pair", + DataType::Struct( + vec![ + Field::new("a", DataType::Int64, true), + Field::new("b", DataType::Utf8, true), + ] + .into(), + ), + true, + ), + Field::new("view", DataType::Utf8View, true), + ])) + } + + fn mixed_batches(count: usize, rows: usize, seed: u64) -> Vec { + use arrow::array::{ + BooleanArray, Date32Array, Float64Array, Int64Array, StructArray, + }; + let schema = mixed_schema(); + let mut state = seed | 1; + let mut next = move || { + state ^= state << 13; + state ^= state >> 7; + state ^= state << 17; + state + }; + (0..count) + .map(|_| { + let values: Vec = (0..rows).map(|_| next()).collect(); + let k = Int32Array::from_iter( + values + .iter() + .map(|v| (v % 9 != 0).then_some((v % 7) as i32 - 3)), + ); + let s = StringArray::from_iter(values.iter().map(|v| { + (v % 5 != 0).then(|| { + format!("{:0>w$}", (v >> 20) % 300, w = (v % 12) as usize) + }) + })); + let flag = BooleanArray::from_iter( + values.iter().map(|v| (v % 3 != 0).then_some(v & 64 != 0)), + ); + let day = Date32Array::from_iter( + values + .iter() + .map(|v| (v % 4 == 0).then_some((v % 400) as i32)), + ); + let amount = Float64Array::from_iter( + values + .iter() + .map(|v| (v % 6 == 0).then_some((v % 1000) as f64 / 7.0)), + ); + let count = Int64Array::from_iter( + values + .iter() + .map(|v| (v % 2 == 0).then_some((v >> 3) as i64)), + ); + let name = StringArray::from_iter(values.iter().map(|v| { + (v % 8 == 0).then(|| { + format!("name-{}-{}", v % 977, "x".repeat((v % 40) as usize)) + }) + })); + let blob = BinaryArray::from_iter( + values + .iter() + .map(|v| (v % 10 != 0).then(|| v.to_le_bytes().repeat(20))), + ); + let dict: DictionaryArray = values + .iter() + .map(|v| (v % 7 != 0).then_some(["p", "qq", "rrr"][(v % 3) as usize])) + .collect(); + let list = ListArray::from_iter_primitive::( + values.iter().map(|v| { + (v % 5 != 1).then(|| (0..(v % 3) as i32).map(|i| Some(i + 1))) + }), + ); + let pair = StructArray::from(vec![ + ( + Arc::new(Field::new("a", DataType::Int64, true)), + Arc::new(Int64Array::from_iter( + values.iter().map(|v| (v % 3 == 1).then_some(*v as i64)), + )) as ArrayRef, + ), + ( + Arc::new(Field::new("b", DataType::Utf8, true)), + Arc::new(StringArray::from_iter( + values + .iter() + .map(|v| (v % 4 == 1).then(|| format!("b{}", v % 13))), + )) as ArrayRef, + ), + ]); + let view = StringViewArray::from_iter(values.iter().map(|v| { + (v % 11 != 0) + .then(|| format!("a view of more than twelve bytes {}", v % 51)) + })); + RecordBatch::try_new( + Arc::clone(&schema), + vec![ + Arc::new(k), + Arc::new(s), + Arc::new(flag), + Arc::new(day), + Arc::new(amount), + Arc::new(count), + Arc::new(name), + Arc::new(blob), + Arc::new(dict), + Arc::new(list), + Arc::new(pair), + Arc::new(view), + ], + ) + .unwrap() + }) + .collect() + } + + fn mixed_ordering(options: &[(&str, bool, bool)]) -> LexOrdering { + let schema = mixed_schema(); + LexOrdering::new(options.iter().map(|(name, descending, nulls_first)| { + PhysicalSortExpr::new( + col(name, &schema).unwrap(), + SortOptions { + descending: *descending, + nulls_first: *nulls_first, + }, + ) + })) + .unwrap() + } + + fn mixed_orderings() -> Vec { + vec![ + mixed_ordering(&[("k", false, false)]), + mixed_ordering(&[("k", true, true)]), + mixed_ordering(&[("k", true, false), ("s", false, true)]), + mixed_ordering(&[ + ("s", true, true), + ("day", false, false), + ("k", false, true), + ]), + mixed_ordering(&[("flag", false, true), ("s", false, false)]), + ] + } + + fn reference(input: &[RecordBatch], ordering: &LexOrdering) -> RecordBatch { + let batch = concat_batches(&input[0].schema(), input).unwrap(); + let columns: Vec = ordering + .iter() + .map(|sort| sort.evaluate_to_sort_column(&batch).unwrap()) + .collect(); + let fields: Vec = columns + .iter() + .map(|c| { + SortField::new_with_options( + c.values.data_type().clone(), + c.options.unwrap(), + ) + }) + .collect(); + let values: Vec = columns.into_iter().map(|c| c.values).collect(); + let rows = RowConverter::new(fields) + .unwrap() + .convert_columns(&values) + .unwrap(); + let mut order: Vec = (0..batch.num_rows() as u32).collect(); + order.sort_by(|a, b| { + rows.row(*a as usize) + .cmp(&rows.row(*b as usize)) + .then(a.cmp(b)) + }); + take_record_batch(&batch, &UInt32Array::from(order)).unwrap() + } + + fn assert_matches_reference( + input: &[RecordBatch], + output: &[RecordBatch], + ordering: &LexOrdering, + ) { + let expected = reference(input, ordering); + let actual = concat_batches(&input[0].schema(), output).unwrap(); + assert_eq!(expected.num_rows(), actual.num_rows()); + let key_columns: Vec = ordering + .iter() + .flat_map(|sort| collect_columns(&sort.expr)) + .map(|column| column.index()) + .collect(); + for column in key_columns { + assert_eq!(expected.column(column), actual.column(column)); + } + let all = |batch: &RecordBatch| { + let fields = batch + .schema() + .fields() + .iter() + .map(|field| SortField::new(field.data_type().clone())) + .collect(); + let mut rows = encoded_rows(batch, batch.columns(), fields); + rows.sort(); + rows + }; + assert!(all(&expected) == all(&actual)); + let fields = actual + .schema() + .fields() + .iter() + .map(|field| SortField::new(field.data_type().clone())) + .collect::>(); + let expected_rows = encoded_rows(&expected, expected.columns(), fields.clone()); + let actual_rows = encoded_rows(&actual, actual.columns(), fields); + let keys = |batch: &RecordBatch| { + let columns: Vec = ordering + .iter() + .map(|sort| sort.evaluate_to_sort_column(batch).unwrap()) + .collect(); + let fields = columns + .iter() + .map(|c| { + SortField::new_with_options( + c.values.data_type().clone(), + c.options.unwrap(), + ) + }) + .collect(); + let values: Vec = columns.into_iter().map(|c| c.values).collect(); + encoded_rows(batch, &values, fields) + }; + let (expected_keys, actual_keys) = (keys(&expected), keys(&actual)); + let mut start = 0; + while start < expected_keys.len() { + let mut end = start + 1; + while end < expected_keys.len() && expected_keys[end] == expected_keys[start] + { + end += 1; + } + assert!( + actual_keys[start..end] + .iter() + .all(|k| k == &expected_keys[start]) + ); + let mut want = expected_rows[start..end].to_vec(); + let mut got = actual_rows[start..end].to_vec(); + want.sort(); + got.sort(); + assert!(want == got); + start = end; + } + } + + fn mixed_sort_exec(input: &[RecordBatch], ordering: LexOrdering) -> Arc { + let source = + TestMemoryExec::try_new_exec(&[input.to_vec()], mixed_schema(), None) + .unwrap(); + Arc::new(SortExec::new(ordering, source)) + } + + fn gather( + input: Vec, + ordering: &LexOrdering, + rows_per_batch: usize, + spilling: bool, + pool: &Arc, + ) -> Result { + let reservation = + datafusion_execution::memory_pool::MemoryConsumer::new("gather") + .register(pool); + let mut counter = RecordBatchMemoryCounter::new(); + let late = LateMaterialization::select(&input[0], ordering)?.unwrap(); + let mut size = 0; + for batch in &input { + size += late.reserved_bytes(batch, &mut counter)?; + } + reservation.try_grow(size)?; + Gather::try_new( + mixed_schema(), + input, + ordering, + rows_per_batch, + spilling, + reservation, + Time::new(), + ) + } + + #[tokio::test] + async fn wide_mixed_rows_in_many_batches_match_the_reference() -> Result<()> { + for (count, rows, batch_size) in [(1, 3000, 512), (40, 97, 256), (300, 7, 8192)] { + let input = mixed_batches(count, rows, count as u64 * 7919 + rows as u64); + assert!( + LateMaterialization::select(&input[0], &mixed_orderings()[0])?.is_some() + ); + for ordering in mixed_orderings() { + let sort = mixed_sort_exec(&input, ordering.clone()); + let output = collect(sort, context(None, batch_size, 1 << 20)).await?; + assert_matches_reference(&input, &output, &ordering); + assert!(output.iter().all(|batch| batch.num_rows() <= batch_size)); + } + } + Ok(()) + } + + #[tokio::test] + async fn spilled_wide_mixed_rows_match_the_reference() -> Result<()> { + let input = mixed_batches(64, 150, 42); + let bytes: usize = input.iter().map(RecordBatch::get_array_memory_size).sum(); + for ordering in mixed_orderings() { + let pool = PeakPool::new(bytes / 3); + let sort = mixed_sort_exec(&input, ordering.clone()); + let output = collect( + Arc::clone(&sort) as Arc, + context(Some(Arc::clone(&pool) as _), 1024, 256 << 10), + ) + .await?; + assert!(sort.metrics().unwrap().spill_count().unwrap() > 0); + assert_matches_reference(&input, &output, &ordering); + assert_eq!(pool.reserved(), 0); + } + Ok(()) + } + + #[test] + fn concatenates_every_column_of_a_scattered_order() -> Result<()> { + let input = mixed_batches(50, 100, 7); + let ordering = mixed_ordering(&[("k", false, true), ("s", true, false)]); + let pool: Arc = Arc::new(GreedyMemoryPool::new(1 << 30)); + let mut gather = gather(input.clone(), &ordering, 256, false, &pool)?; + assert!(!gather.local); + assert!(gather.columns.iter().all(|pieces| pieces.arrays.len() == 1)); + assert_eq!(gather.layouts.len(), 1); + assert_eq!(gather.live_bytes, 0); + assert_eq!(pool.reserved(), gather.needed()); + let mut output = vec![]; + while gather.cursor < gather.order.len() { + output.push(gather.next_batch()?); + assert_eq!(pool.reserved(), gather.needed()); + } + assert_eq!(gather.gathered_bytes, 0); + assert_eq!(pool.reserved(), gather.order.get_array_memory_size()); + drop(gather); + assert_eq!(pool.reserved(), 0); + assert_matches_reference(&input, &output, &ordering); + Ok(()) + } + + #[test] + fn interleaves_an_order_that_reads_few_batches_at_a_time() -> Result<()> { + let input = mixed_batches(50, 100, 7); + let mut keyed = vec![]; + for (i, batch) in input.iter().enumerate() { + let mut columns = batch.columns().to_vec(); + columns[0] = Arc::new(Int32Array::from(vec![i as i32; batch.num_rows()])); + keyed.push(RecordBatch::try_new(mixed_schema(), columns)?); + } + let ordering = mixed_ordering(&[("k", false, true)]); + let pool: Arc = Arc::new(GreedyMemoryPool::new(1 << 30)); + let mut gather = gather(keyed.clone(), &ordering, 64, false, &pool)?; + assert!(gather.local); + assert!( + gather + .columns + .iter() + .all(|pieces| pieces.arrays.len() == 50) + ); + let held = pool.reserved(); + let mut output = vec![gather.next_batch()?]; + let mut lowest = held; + while gather.cursor < gather.order.len() { + output.push(gather.next_batch()?); + lowest = lowest.min(pool.reserved()); + } + assert!(lowest < held / 4); + drop(gather); + assert_matches_reference(&keyed, &output, &ordering); + assert_eq!(pool.reserved(), 0); + Ok(()) + } + + #[test] + fn concatenates_in_chunks_a_column_the_pool_cannot_hold_twice() -> Result<()> { + let input = mixed_batches(50, 100, 11); + let ordering = mixed_ordering(&[("s", false, true), ("k", true, true)]); + let late = LateMaterialization::select(&input[0], &ordering)?.unwrap(); + let mut counter = RecordBatchMemoryCounter::new(); + let mut held = 0; + for batch in &input { + held += late.reserved_bytes(batch, &mut counter)?; + } + let pool: Arc = Arc::new(GreedyMemoryPool::new(held)); + let gather = gather(input.clone(), &ordering, 256, false, &pool)?; + let blob = gather.columns[7].arrays.len(); + assert!(blob > 1 && blob < 50, "{blob} pieces"); + assert_eq!(gather.columns[2].arrays.len(), 1); + assert_eq!(gather.layouts.len(), 2); + assert_eq!(pool.reserved(), gather.needed()); + let output = gather.collect::>>()?; + assert_matches_reference(&input, &output, &ordering); + assert_eq!(pool.reserved(), 0); + Ok(()) + } + + #[test] + fn a_spill_concatenates_columns_in_chunks_of_a_share_of_its_input() -> Result<()> { + let input = mixed_batches(50, 100, 13); + let ordering = mixed_ordering(&[("k", false, false), ("s", false, false)]); + let pool: Arc = Arc::new(GreedyMemoryPool::new(1 << 30)); + let gather = gather(input.clone(), &ordering, 256, true, &pool)?; + let buffered: usize = input.iter().map(|b| b.get_sliced_size().unwrap()).sum(); + for pieces in &gather.columns { + assert!(pieces.arrays.len() < 50); + for (array, batches) in pieces.arrays.iter().zip(&pieces.batches) { + if batches.len() > 1 { + let bytes = sliced_bytes(array.as_ref())?; + assert!( + bytes <= 2 * buffered / SPILL_COLUMN_SHARE, + "{bytes} of {buffered}" + ); + } + } + } + assert!(gather.columns[7].arrays.len() > 1); + assert_eq!(gather.columns[2].arrays.len(), 1); + let output = gather.collect::>>()?; + assert_matches_reference(&input, &output, &ordering); + assert_eq!(pool.reserved(), 0); + Ok(()) + } + + #[test] + fn list_offsets_that_would_overflow_are_not_concatenated() -> Result<()> { + let list_type = + DataType::List(Arc::new(Field::new_list_field(DataType::Null, true))); + let list = |len: i32| { + ArrayData::builder(list_type.clone()) + .len(1) + .add_buffer(arrow::buffer::Buffer::from_slice_ref([0i32, len])) + .add_child_data(ArrayData::new_null(&DataType::Null, len as usize)) + .build() + }; + let small = list(1)?; + let big = list(i32::MAX - 1)?; + assert!(offsets_fit(&[small.clone(), small.clone()])); + assert!(offsets_fit(&[big.clone(), small.clone()])); + assert!(!offsets_fit(&[big.clone(), small.clone(), small])); + assert!(!offsets_fit(&[big.clone(), big])); + Ok(()) + } } From 5e4ef4aeb73df0c7f5b0a5725c68b8c53132c42a Mon Sep 17 00:00:00 2001 From: msaf Date: Wed, 7 Oct 2026 13:51:01 +0100 Subject: [PATCH 11/13] perf: sort a wide native sort's row keys by an 8-byte prefix first The keys of a late-materialized sort over several columns were sorted as (row bytes, index) pairs, every comparison a memcmp through a pointer into the rows. They are now sorted as (first 8 bytes big-endian, index) pairs, and only runs of equal prefixes are sorted again by the whole row. The row format already encodes descending order, null order and every type as bytes, and a row shorter than 8 bytes is zero-padded, which can only tie with the rows it is a prefix of, so the order is the same. Keys of the SortSpillBenchSuite rows (a 24-character hex id, two dates, an int; every row ties on its prefix with ~17 others) alone, ns/row, before -> after: 1M 97 -> 81, 6M 132 -> 123, 13M 160 -> 123. The whole sort, end to end, alternating runs: 6M 428/395 -> 380/386, 13M 539/513 -> 518/459. Co-Authored-By: Claude Opus 5.5 (1M context) --- .../src/sorts/sort/late_materialize.rs | 91 ++++++++++++++++++- 1 file changed, 89 insertions(+), 2 deletions(-) diff --git a/native/vendor/datafusion-physical-plan/src/sorts/sort/late_materialize.rs b/native/vendor/datafusion-physical-plan/src/sorts/sort/late_materialize.rs index a8d5681c7cb..c73d914f4db 100644 --- a/native/vendor/datafusion-physical-plan/src/sorts/sort/late_materialize.rs +++ b/native/vendor/datafusion-physical-plan/src/sorts/sort/late_materialize.rs @@ -224,15 +224,37 @@ fn row_order(rows: &Rows) -> Vec { keys.sort_unstable(); return keys.into_iter().map(|(_, index)| index).collect(); } - let mut keys: Vec<(&[u8], u32)> = rows + let mut keys: Vec<(u64, u32)> = rows .iter() .enumerate() - .map(|(index, row)| (row.data(), index as u32)) + .map(|(index, row)| (prefix(row.data()), index as u32)) .collect(); keys.sort_unstable(); + let mut start = 0; + while start < keys.len() { + let mut end = start + 1; + while end < keys.len() && keys[end].0 == keys[start].0 { + end += 1; + } + if end - start > 1 { + keys[start..end].sort_unstable_by(|a, b| { + let left = rows.row(a.1 as usize); + let right = rows.row(b.1 as usize); + left.data().cmp(right.data()).then(a.1.cmp(&b.1)) + }); + } + start = end; + } keys.into_iter().map(|(_, index)| index).collect() } +fn prefix(data: &[u8]) -> u64 { + let mut bytes = [0u8; 8]; + let len = data.len().min(8); + bytes[..len].copy_from_slice(&data[..len]); + u64::from_be_bytes(bytes) +} + struct Gather { schema: SchemaRef, columns: Vec, @@ -1666,4 +1688,69 @@ mod tests { assert!(!offsets_fit(&[big.clone(), big])); Ok(()) } + + #[test] + fn row_order_matches_a_full_comparison_of_the_rows() -> Result<()> { + let mut state = 0x9E37_79B9_7F4A_7C15_u64; + let mut next = move || { + state ^= state << 13; + state ^= state >> 7; + state ^= state << 17; + state + }; + for options in [ + SortOptions::default(), + SortOptions { + descending: true, + nulls_first: true, + }, + SortOptions { + descending: false, + nulls_first: true, + }, + SortOptions { + descending: true, + nulls_first: false, + }, + ] { + let values: Vec = (0..20000).map(|_| next()).collect(); + let strings = StringArray::from_iter(values.iter().map(|v| { + (v % 9 != 0).then(|| { + let len = (v >> 8) % 14; + let base = ["", "a", "ab", "abcdefg", "abcdefgh", "abcdefghij", "b"] + [(v % 7) as usize]; + format!("{base}{}", "z".repeat(len as usize % 3)) + .repeat((len / 5) as usize + 1) + }) + })); + let ints = Int32Array::from_iter( + values + .iter() + .map(|v| (v % 5 != 0).then_some(((v >> 16) % 4) as i32 - 2)), + ); + let bytes = BinaryArray::from_iter( + values + .iter() + .map(|v| (v % 6 != 0).then(|| vec![0u8; ((v >> 24) % 10) as usize])), + ); + let columns: Vec = + vec![Arc::new(strings), Arc::new(ints), Arc::new(bytes)]; + for width in 1..=3 { + let fields = columns[..width] + .iter() + .map(|c| SortField::new_with_options(c.data_type().clone(), options)) + .collect(); + let rows = + RowConverter::new(fields)?.convert_columns(&columns[..width])?; + let mut expected: Vec = (0..rows.num_rows() as u32).collect(); + expected.sort_by(|a, b| { + rows.row(*a as usize) + .cmp(&rows.row(*b as usize)) + .then(a.cmp(b)) + }); + assert_eq!(row_order(&rows), expected); + } + } + Ok(()) + } } From 3023e0b024cd98b62fe498020732b9e8441e91bc Mon Sep 17 00:00:00 2001 From: msaf Date: Wed, 7 Oct 2026 15:57:45 +0100 Subject: [PATCH 12/13] feat: recalibrate the sort lines of the engine cost table on prod-sized tasks The sort lines came from calib29's small tasks, where Comet's sort cost 224 + 0.023 * L^2 against Spark's flat 646 ns per row. A local single-core benchmark (sortWithinPartitions minus the same scan and conversion, so Spark's inserts and copies count, not only its sort time) after the gather and key-sort changes measured 3 to 30 flat leaves and 8 and 22 leaves in structs and arrays at 1M, 6M and 13M rows per task. Both engines cost more per row the more rows a task sorts, Comet more: Comet/Spark is 0.4 at 1M rows, 0.75 flat and about 1 nested at 13M, where Comet's per-leaf gather from separate columns outgrows the caches while Spark copies whole rows. The lines are fitted at 13M rows and scaled by 1.6, the ratio of the pr29 prices of small tasks to the benchmark at 1M rows: sort flat Comet 321 + 21.2 L (was 224 + 0.023 L^2), Spark 432 + 24.1 L (646) sort nested Comet 1110 + 3.3 L (244 + 2.69 L + 0.016 L^2), Spark 621 + 25.4 L (770 + 2.34 L) sortSpill flat Comet 22.2 L (265 + 49 L), Spark 93 + 39.3 L (495 + 38.5 L) sortSpill nested Comet 34.6 L (324 + 45.3 L), Spark 66.7 L (502 + 33.5 L) sortSpillFraction stays 0. A narrow native sort now gains about 120 ns per row over Spark instead of 420, less than a columnar shuffle written for it by a Spark producer, so the reduce-side sort of a Spark sort aggregate moves to Spark with its shuffle; the suite asserts that, and that both sorts stay native when the native sort is cheap. Co-Authored-By: Claude Opus 5.5 (1M context) --- .../apache/comet/rules/EngineCostTable.scala | 37 +++++++++++++------ .../rules/CostBasedEngineChoiceSuite.scala | 23 ++++++++---- 2 files changed, 42 insertions(+), 18 deletions(-) diff --git a/spark/src/main/scala/org/apache/comet/rules/EngineCostTable.scala b/spark/src/main/scala/org/apache/comet/rules/EngineCostTable.scala index f807bcfc18b..1bc084adbcb 100644 --- a/spark/src/main/scala/org/apache/comet/rules/EngineCostTable.scala +++ b/spark/src/main/scala/org/apache/comet/rules/EngineCostTable.scala @@ -290,17 +290,32 @@ object EngineCostTable { Seq((costClass, Flat) -> line, (costClass, Nested) -> line) /** - * The default prices, measured on pr29 (calib29, calib29c and calib29d) with L counting every - * leaf of the row. Lines are `(class, form) -> Line(Comet c0, Comet k0, Comet k1, Spark c0, - * Spark k)`; a class priced alike in both forms has one line for both. + * The default prices, measured on pr29 (calib29, calib29c and calib29d) except `sort` and + * `sortSpill`, with L counting every leaf of the row. Lines are `(class, form) -> Line(Comet + * c0, Comet k0, Comet k1, Spark c0, Spark k)`; a class priced alike in both forms has one line + * for both. * * - `shuffleWrite` (at `shuffleWritePartitionBase` partitions) and `shuffleRead`: every leaf * of the shuffled rows, Spark's fetch wait and the read conversion excluded. Spark is * averaged over 250 to 4800 partitions, on which it does not depend. - * - `sort`: every leaf of the sorted rows. Spark sorts pointers, so its price barely depends - * on the width; Comet flat is noisy (maxrel 0.6). `sortSpill` is what spilling adds, on a - * fraction `sortSpillFraction` of rows, none by default: rows are not estimated, so a spill - * cannot be predicted. + * - `sort`: every leaf of the sorted rows, measured apart from the other lines in a local + * benchmark on one core: the time of a `sortWithinPartitions` minus that of the same scan + * and conversion without it, so Spark's inserts, copies and output count, not only its + * `sort time` (about 0.7 of it). Rows of 3, 8, 15, 22 and 30 flat leaves (strings, dates, + * longs, doubles, booleans, ints and timestamps, mostly null) sorted by a 24-character id + * and an int, and of 8 and 22 leaves with all but the keys in structs and an array of + * structs, at 1M, 6M and 13M rows per task (6M and 13M twice). The lines are fitted at 13M + * rows, as wide production sorts run, and scaled by 1.6 to the cluster: the pr29 prices of + * small tasks are 1.5 (Comet) and 1.7 (Spark) times this harness at 1M rows. Both engines + * cost more per row the more rows a task sorts, Comet more, and the model does not know the + * rows: at 1M rows Comet costs about 150 ns locally whatever the width, Spark 300 + 5.5 * + * L, a ratio of 0.4; at 13M rows the ratio is 0.75 flat and about 1 nested. Comet gathers + * every leaf of a row from its own column, which outgrows the caches; Spark copies the row + * whole. Comet is noisy (maxrel 0.34), Spark less (0.13), and the nested lines rest on two + * widths. `sortSpill` is what spilling every row once adds, at 6M rows (some in-memory runs + * at 13M spilled on their own), through the origin where the intercept came out negative; + * it is priced on a fraction `sortSpillFraction` of rows, none by default: rows are not + * estimated, so a spill cannot be predicted. * - `smj`: the join over its sorted inputs, every output leaf, noisy. `bhj`: the probe side, * every output leaf; the nested `bhj` is noisy (Comet 0 to 350, Spark 50 to 4200 ns) and * takes the flat line. `smjCondition`: what a join condition adds to `smj`, every output @@ -366,10 +381,10 @@ object EngineCostTable { (ShuffleWrite, Nested) -> Line(46, 39.59, 0.031, 686, 47.77), (ShuffleRead, Flat) -> Line(0, 14.73, 0.032, 67, 26.34), (ShuffleRead, Nested) -> Line(0, 10.79, 0.019, 167, 19.02), - (Sort, Flat) -> Line(224, 0, 0.023, 646, 0), - (Sort, Nested) -> Line(244, 2.69, 0.016, 770, 2.34), - (SortSpill, Flat) -> Line(265, 49.0, 0, 495, 38.5), - (SortSpill, Nested) -> Line(324, 45.3, 0, 502, 33.5), + (Sort, Flat) -> Line(321, 21.2, 0, 432, 24.1), + (Sort, Nested) -> Line(1110, 3.3, 0, 621, 25.4), + (SortSpill, Flat) -> Line(0, 22.2, 0, 93, 39.3), + (SortSpill, Nested) -> Line(0, 34.6, 0, 0, 66.7), (Smj, Flat) -> Line(0, 4, 0, 0, 35), (Smj, Nested) -> Line(0, 0.45, 0.046, 0, 32), (Bhj, Flat) -> Line(72, 2.3, 0.002, 0, 20.5), diff --git a/spark/src/test/scala/org/apache/comet/rules/CostBasedEngineChoiceSuite.scala b/spark/src/test/scala/org/apache/comet/rules/CostBasedEngineChoiceSuite.scala index d9b7f67363f..4c2d2e461ed 100644 --- a/spark/src/test/scala/org/apache/comet/rules/CostBasedEngineChoiceSuite.scala +++ b/spark/src/test/scala/org/apache/comet/rules/CostBasedEngineChoiceSuite.scala @@ -46,6 +46,7 @@ class CostBasedEngineChoiceSuite extends CometTestBase { private val flag = CometConf.COMET_EXEC_COST_BASED_ENGINES_ENABLED.key private val costTable = CometConf.COMET_EXEC_COST_BASED_ENGINES_COST_TABLE.key private val expensiveNativeSort = costTable -> "sort.flat.comet=1000,0,0" + private val cheapNativeSort = costTable -> "sort.flat.comet=0,0,0" private def same(actual: Double, expected: Double): Unit = assert(math.abs(actual - expected) < 1e-9, s"$actual != $expected") @@ -157,7 +158,14 @@ class CostBasedEngineChoiceSuite extends CometTestBase { assert(count(off) { case s: SortAggregateExec => s } == 2, s"plan:\n$off") assert(count(off) { case s: CometSortExec => s } == 2, s"plan:\n$off") assert(count(on) { case s: SortAggregateExec => s } == 2, s"plan:\n$on") - assert(count(on) { case s: CometSortExec => s } == 2, s"plan:\n$on") + assert(count(on) { case s: CometSortExec => s } == 1, s"plan:\n$on") + assert(count(on) { case s: SortExec => s } == 1, s"plan:\n$on") + val mapSide = nodes(on).collectFirst { case s: CometSortExec => s }.get + assert(nodes(mapSide).exists(_.isInstanceOf[CometNativeScanExec]), s"plan:\n$on") + withSQLConf(flag -> "true", cheapNativeSort) { + val kept = run(query) + assert(count(kept) { case s: CometSortExec => s } == 2, s"plan:\n$kept") + } withSQLConf(flag -> "true", expensiveNativeSort) { val moved = run(query) assert(count(moved) { case s: CometSortExec => s } == 0, s"plan:\n$moved") @@ -500,7 +508,7 @@ class CostBasedEngineChoiceSuite extends CometTestBase { "keepFiltersOverNativeScans=false;keepPartialAggregatesOverNativeInputs=false;" + "shuffleReadPerByte.spark=3;cometShuffleBytesRatio=0.75;quadraticLeafCap=10;" + "sortSpillFraction=0.5") - assert(table.line(Sort, Form.Flat) == Line(224, 1, 2, 646, 0)) + assert(table.line(Sort, Form.Flat) == Line(321, 1, 2, 432, 24.1)) assert(table.line(ShuffleRead, Form.Nested) == Line(0, 10.79, 0.019, 3, 4)) assert(table.line(ShuffleWrite, Form.Flat) == Line(5, 6, 7, 69, 67.21)) assert(table.line(C2R, Form.Nested).cometK0 == 7) @@ -515,7 +523,7 @@ class CostBasedEngineChoiceSuite extends CometTestBase { assert(table.filterPassThroughPerLeafSpark == 0) same(table.cometShuffleBytes(100), (0.5 + 2) * 0.75 * 100) same(table.sparkShuffleBytes(100), (3.6 + 3) * 100) - same(table.comet(Sort, Width(20, 0)), 224 + 1 * 20 + 2 * 20 * 10) + same(table.comet(Sort, Width(20, 0)), 321 + 1 * 20 + 2 * 20 * 10) assert(table.sortSpillFraction == 0.5) assert(!table.keepFiltersOverNativeScans) assert(EngineCostTable.default.keepFiltersOverNativeScans) @@ -570,10 +578,11 @@ class CostBasedEngineChoiceSuite extends CometTestBase { same(t.spark(ShuffleWrite, flat10), 69 + 67.21 * 10) same(t.comet(ShuffleRead, Width(4, 4)), 10.79 * 4 + 0.019 * 16) same(t.spark(ShuffleRead, Width(4, 4)), 167 + 19.02 * 4) - same(t.spark(Sort, Width(200, 0)), 646) - same(t.comet(Sort, Width(200, 0)), 224 + 0.023 * 40000) - same(t.comet(Sort, Width(200, 200)), 244 + 2.69 * 200 + 0.016 * 40000) - same(t.comet(Sort, Width(1000, 0)), 224 + 0.023 * 1000 * 600) + same(t.spark(Sort, Width(200, 0)), 432 + 24.1 * 200) + same(t.comet(Sort, Width(200, 0)), 321 + 21.2 * 200) + same(t.comet(Sort, Width(200, 200)), 1110 + 3.3 * 200) + same(t.spark(Sort, Width(200, 200)), 621 + 25.4 * 200) + same(t.comet(ShuffleWrite, Width(1000, 0)), 1000 * (48.95 + 0.037 * 600)) same(t.comet(Smj, flat10), 40) same(t.spark(Smj, flat10), 350) same(t.comet(Bhj, flat10), 72 + 23 + 0.2) From c90d450adfd68efac07ab34a9ba27e57c8cec1a8 Mon Sep 17 00:00:00 2001 From: msaf Date: Wed, 7 Oct 2026 17:44:21 +0100 Subject: [PATCH 13/13] fix: keep the calib29 prices of nested sorts The local fit on 8 and 22 nested leaves extrapolated to 509 nested leaves sent mongo finance's sort before its window group limit native, where it cost 13.7 us per row and spilled 955 GB against 0.28 us and no spill in Spark (the whole job 39.96 task-hours against 35.80). The nested sort and sortSpill lines go back to calib29, measured on the cluster at 8 to 512 leaves, which keep that sort in Spark; the flat lines stay recalibrated. Co-Authored-By: Claude Opus 5.5 (1M context) --- .../apache/comet/rules/EngineCostTable.scala | 20 +++++++++++-------- .../rules/CostBasedEngineChoiceSuite.scala | 5 +++-- 2 files changed, 15 insertions(+), 10 deletions(-) diff --git a/spark/src/main/scala/org/apache/comet/rules/EngineCostTable.scala b/spark/src/main/scala/org/apache/comet/rules/EngineCostTable.scala index 1bc084adbcb..d900df08d2d 100644 --- a/spark/src/main/scala/org/apache/comet/rules/EngineCostTable.scala +++ b/spark/src/main/scala/org/apache/comet/rules/EngineCostTable.scala @@ -298,7 +298,11 @@ object EngineCostTable { * - `shuffleWrite` (at `shuffleWritePartitionBase` partitions) and `shuffleRead`: every leaf * of the shuffled rows, Spark's fetch wait and the read conversion excluded. Spark is * averaged over 250 to 4800 partitions, on which it does not depend. - * - `sort`: every leaf of the sorted rows, measured apart from the other lines in a local + * - `sort`: every leaf of the sorted rows. The nested `sort` and `sortSpill` lines keep the + * calib29 prices, measured on the cluster at 8 to 512 leaves: a local fit on 8 and 22 + * nested leaves extrapolated to 509 nested leaves sent mongo finance's sort before its + * window group limit native, where it cost 13.7 us per row and spilled 955 GB against 0.28 + * us and no spill in Spark. The flat lines are measured apart from the others in a local * benchmark on one core: the time of a `sortWithinPartitions` minus that of the same scan * and conversion without it, so Spark's inserts, copies and output count, not only its * `sort time` (about 0.7 of it). Rows of 3, 8, 15, 22 and 30 flat leaves (strings, dates, @@ -311,11 +315,11 @@ object EngineCostTable { * rows: at 1M rows Comet costs about 150 ns locally whatever the width, Spark 300 + 5.5 * * L, a ratio of 0.4; at 13M rows the ratio is 0.75 flat and about 1 nested. Comet gathers * every leaf of a row from its own column, which outgrows the caches; Spark copies the row - * whole. Comet is noisy (maxrel 0.34), Spark less (0.13), and the nested lines rest on two - * widths. `sortSpill` is what spilling every row once adds, at 6M rows (some in-memory runs - * at 13M spilled on their own), through the origin where the intercept came out negative; - * it is priced on a fraction `sortSpillFraction` of rows, none by default: rows are not - * estimated, so a spill cannot be predicted. + * whole. Comet is noisy (maxrel 0.34), Spark less (0.13). The flat `sortSpill` is what + * spilling every row once adds, at 6M rows (some in-memory runs at 13M spilled on their + * own), through the origin where the intercept came out negative; it is priced on a + * fraction `sortSpillFraction` of rows, none by default: rows are not estimated, so a spill + * cannot be predicted. * - `smj`: the join over its sorted inputs, every output leaf, noisy. `bhj`: the probe side, * every output leaf; the nested `bhj` is noisy (Comet 0 to 350, Spark 50 to 4200 ns) and * takes the flat line. `smjCondition`: what a join condition adds to `smj`, every output @@ -382,9 +386,9 @@ object EngineCostTable { (ShuffleRead, Flat) -> Line(0, 14.73, 0.032, 67, 26.34), (ShuffleRead, Nested) -> Line(0, 10.79, 0.019, 167, 19.02), (Sort, Flat) -> Line(321, 21.2, 0, 432, 24.1), - (Sort, Nested) -> Line(1110, 3.3, 0, 621, 25.4), + (Sort, Nested) -> Line(244, 2.69, 0.016, 770, 2.34), (SortSpill, Flat) -> Line(0, 22.2, 0, 93, 39.3), - (SortSpill, Nested) -> Line(0, 34.6, 0, 0, 66.7), + (SortSpill, Nested) -> Line(324, 45.3, 0, 502, 33.5), (Smj, Flat) -> Line(0, 4, 0, 0, 35), (Smj, Nested) -> Line(0, 0.45, 0.046, 0, 32), (Bhj, Flat) -> Line(72, 2.3, 0.002, 0, 20.5), diff --git a/spark/src/test/scala/org/apache/comet/rules/CostBasedEngineChoiceSuite.scala b/spark/src/test/scala/org/apache/comet/rules/CostBasedEngineChoiceSuite.scala index 4c2d2e461ed..2ccc7f6f19e 100644 --- a/spark/src/test/scala/org/apache/comet/rules/CostBasedEngineChoiceSuite.scala +++ b/spark/src/test/scala/org/apache/comet/rules/CostBasedEngineChoiceSuite.scala @@ -580,8 +580,9 @@ class CostBasedEngineChoiceSuite extends CometTestBase { same(t.spark(ShuffleRead, Width(4, 4)), 167 + 19.02 * 4) same(t.spark(Sort, Width(200, 0)), 432 + 24.1 * 200) same(t.comet(Sort, Width(200, 0)), 321 + 21.2 * 200) - same(t.comet(Sort, Width(200, 200)), 1110 + 3.3 * 200) - same(t.spark(Sort, Width(200, 200)), 621 + 25.4 * 200) + same(t.comet(Sort, Width(200, 200)), 244 + 200 * (2.69 + 0.016 * 200)) + same(t.spark(Sort, Width(200, 200)), 770 + 2.34 * 200) + assert(t.spark(Sort, Width(509, 487)) < t.comet(Sort, Width(509, 487))) same(t.comet(ShuffleWrite, Width(1000, 0)), 1000 * (48.95 + 0.037 * 600)) same(t.comet(Smj, flat10), 40) same(t.spark(Smj, flat10), 350)