This is an automated email from the ASF dual-hosted git repository.

zhztheplayer pushed a commit to branch main
in repository https://gitbox.apache.org/repos/asf/gluten.git


The following commit(s) were added to refs/heads/main by this push:
     new fcda79c1ea [VL] Option to insert partial-merge aggregations for eager 
aggregation (#13086)
fcda79c1ea is described below

commit fcda79c1eacaae9b9256e598a7c33385af5c9f8d
Author: Hongze Zhang <[email protected]>
AuthorDate: Thu Sep 24 04:22:34 2026 +0200

    [VL] Option to insert partial-merge aggregations for eager aggregation 
(#13086)
---
 .../execution/StarSchemaJoinAggregateSuite.scala   |   8 +
 docs/Configuration.md                              |   1 +
 .../org/apache/gluten/config/GlutenConfig.scala    |  10 +
 .../extension/joinagg/ImplementJoinAggregate.scala | 156 +++++++----
 .../joinagg/JoinAggregateFunctionWrapper.scala     |  38 ++-
 .../joinagg/PushAggregateThroughJoin.scala         | 114 +++++---
 .../execution/PushAggregateThroughJoinSuite.scala  | 289 +++++++++++++++++----
 7 files changed, 463 insertions(+), 153 deletions(-)

diff --git 
a/backends-velox/src/test/scala/org/apache/gluten/execution/StarSchemaJoinAggregateSuite.scala
 
b/backends-velox/src/test/scala/org/apache/gluten/execution/StarSchemaJoinAggregateSuite.scala
index 54aa3770d7..9e848847ae 100644
--- 
a/backends-velox/src/test/scala/org/apache/gluten/execution/StarSchemaJoinAggregateSuite.scala
+++ 
b/backends-velox/src/test/scala/org/apache/gluten/execution/StarSchemaJoinAggregateSuite.scala
@@ -70,6 +70,13 @@ class StarSchemaJoinAggregateSingleDepthSuite extends 
StarSchemaJoinAggregateSui
   }
 }
 
+class StarSchemaJoinAggregatePartialMergeEnabledSuite extends 
StarSchemaJoinAggregateSuite {
+  override protected def sparkConf: SparkConf = {
+    super.sparkConf
+      .set(GlutenConfig.PUSH_AGGREGATE_THROUGH_JOIN_PARTIAL_MERGE_ENABLED.key, 
"true")
+  }
+}
+
 class StarSchemaJoinAggregateSuite extends VeloxTPCHTableSupport with 
AdaptiveSparkPlanHelper {
   private val factMeasureColumnNames = Set(
     "sales_price",
@@ -95,6 +102,7 @@ class StarSchemaJoinAggregateSuite extends 
VeloxTPCHTableSupport with AdaptiveSp
       .set(GlutenConfig.COLUMNAR_FORCE_SHUFFLED_HASH_JOIN_ENABLED.key, "true")
       .set(GlutenConfig.PUSH_AGGREGATE_THROUGH_JOIN_ENABLED.key, "true")
       .set(GlutenConfig.PUSH_AGGREGATE_THROUGH_JOIN_MAX_DEPTH.key, 
s"${Int.MaxValue}")
+      .set(GlutenConfig.PUSH_AGGREGATE_THROUGH_JOIN_PARTIAL_MERGE_ENABLED.key, 
"false")
       .set("spark.sql.adaptive.enabled", "false")
   }
 
diff --git a/docs/Configuration.md b/docs/Configuration.md
index a9b5c669e1..83e3432343 100644
--- a/docs/Configuration.md
+++ b/docs/Configuration.md
@@ -138,6 +138,7 @@ nav_order: 15
 | spark.gluten.sql.orc.charType.scan.fallback.enabled                 | 🔄 
Dynamic    | true              | Force fallback for orc char type scan.         
                                                                                
                                                                                
                                                                                
                                                                                
                   [...]
 | spark.gluten.sql.pushAggregateThroughJoin.enabled                   | 🔄 
Dynamic    | false             | Enables the push-aggregate-through-join 
optimization in Gluten. When enabled, aggregate operators may be pushed below 
joins during logical optimization and corresponding physical plans may be 
rewritten to execute the aggregation earlier.                                   
                                                                                
                                  [...]
 | spark.gluten.sql.pushAggregateThroughJoin.maxDepth                  | 🔄 
Dynamic    | 2147483647        | Maximum join traversal depth when applying the 
push-aggregate-through-join optimization. A value of 1 allows pushing an 
aggregate through a single join; larger values allow the rule to traverse and 
push through multiple consecutive joins.                                        
                                                                                
                            [...]
+| spark.gluten.sql.pushAggregateThroughJoin.partialMerge.enabled      | 🔄 
Dynamic    | false             | Enables a PartialMerge aggregate above each 
aggregate pushed through a join.                                                
                                                                                
                                                                                
                                                                                
                      [...]
 | spark.gluten.sql.removeNativeWriteFilesSortAndProject               | 🔄 
Dynamic    | true              | When true, Gluten will remove the vanilla 
Spark V1Writes added sort and project for velox backend.                        
                                                                                
                                                                                
                                                                                
                        [...]
 | spark.gluten.sql.rewrite.dateTimestampComparison                    | 🔄 
Dynamic    | true              | Rewrite the comparision between date and 
timestamp to timestamp comparison.For example `from_unixtime(ts) > date` will 
be rewritten to `ts > to_unixtime(date)`                                        
                                                                                
                                                                                
                           [...]
 | spark.gluten.sql.scan.detailedMetrics.enabled                       | 🔄 
Dynamic    | true              | When true (default), Velox backend scan 
operators register all detailed SQL metrics. When false, only essential scan 
metrics are registered to reduce driver memory usage. Also enabled 
automatically when spark.gluten.sql.debug is true. Does not affect the 
ClickHouse backend.                                                             
                                                   [...]
diff --git 
a/gluten-substrait/src/main/scala/org/apache/gluten/config/GlutenConfig.scala 
b/gluten-substrait/src/main/scala/org/apache/gluten/config/GlutenConfig.scala
index 0f2805c8d5..753bc865ab 100644
--- 
a/gluten-substrait/src/main/scala/org/apache/gluten/config/GlutenConfig.scala
+++ 
b/gluten-substrait/src/main/scala/org/apache/gluten/config/GlutenConfig.scala
@@ -158,6 +158,9 @@ class GlutenConfig(conf: SQLConf) extends 
GlutenCoreConfig(conf) {
   def pushAggregateThroughJoinMaxDepth: Int =
     getConf(PUSH_AGGREGATE_THROUGH_JOIN_MAX_DEPTH)
 
+  def pushAggregateThroughJoinPartialMergeEnabled: Boolean =
+    getConf(PUSH_AGGREGATE_THROUGH_JOIN_PARTIAL_MERGE_ENABLED)
+
   def forceOrcCharTypeScanFallbackEnabled: Boolean =
     getConf(VELOX_FORCE_ORC_CHAR_TYPE_SCAN_FALLBACK)
 
@@ -784,6 +787,13 @@ object GlutenConfig extends ConfigRegistry {
       .checkValue(_ >= 1, "must be greater than or equal to 1.")
       .createWithDefault(Int.MaxValue)
 
+  val PUSH_AGGREGATE_THROUGH_JOIN_PARTIAL_MERGE_ENABLED =
+    buildConf("spark.gluten.sql.pushAggregateThroughJoin.partialMerge.enabled")
+      .doc(
+        "Enables a PartialMerge aggregate above each aggregate pushed through 
a join.")
+      .booleanConf
+      .createWithDefault(false)
+
   val GLUTEN_SOFT_AFFINITY_ENABLED =
     buildConf("spark.gluten.soft-affinity.enabled")
       .doc("Whether to enable Soft Affinity scheduling.")
diff --git 
a/gluten-substrait/src/main/scala/org/apache/gluten/extension/joinagg/ImplementJoinAggregate.scala
 
b/gluten-substrait/src/main/scala/org/apache/gluten/extension/joinagg/ImplementJoinAggregate.scala
index bc01889eed..5a9434e473 100644
--- 
a/gluten-substrait/src/main/scala/org/apache/gluten/extension/joinagg/ImplementJoinAggregate.scala
+++ 
b/gluten-substrait/src/main/scala/org/apache/gluten/extension/joinagg/ImplementJoinAggregate.scala
@@ -97,6 +97,8 @@ case class ImplementJoinAggregate(spark: SparkSession) 
extends SparkStrategy {
       phase match {
         case JoinAggregateFunctionWrapper.PartialPhase =>
           planPartialPhase(grouping, aggExpressions, resultExpressions, 
childPlan)
+        case JoinAggregateFunctionWrapper.PartialMergePhase =>
+          planPartialMergePhase(grouping, aggExpressions, resultExpressions, 
childPlan)
         case JoinAggregateFunctionWrapper.FinalPhase =>
           planFinalPhase(grouping, aggExpressions, resultExpressions, 
childPlan)
       }
@@ -109,20 +111,100 @@ case class ImplementJoinAggregate(spark: SparkSession) 
extends SparkStrategy {
       aggregateExpressions: Seq[AggregateExpression],
       resultExpressions: Seq[NamedExpression],
       childPlan: SparkPlan): Option[SparkPlan] = {
-    // The pushed logical aggregate exposes one wrapper-typed output per 
pushed aggregate. Spark
-    // physically computes ordinary aggregate buffers, so this phase runs a 
normal HashAggregateExec
-    // first and then repacks those buffers into the struct-valued wrapper 
outputs expected by the
-    // logical plan above.
-    val rewrittenAggExprs = aggregateExpressions.map {
+    planPhase(
+      grouping,
+      aggregateExpressions,
+      resultExpressions,
+      childPlan,
+      unpackInputBuffers = false,
+      packOutputBuffers = true)
+  }
+
+  private def planPartialMergePhase(
+      grouping: Seq[NamedExpression],
+      aggregateExpressions: Seq[AggregateExpression],
+      resultExpressions: Seq[NamedExpression],
+      childPlan: SparkPlan): Option[SparkPlan] = {
+    planPhase(
+      grouping,
+      aggregateExpressions,
+      resultExpressions,
+      childPlan,
+      unpackInputBuffers = true,
+      packOutputBuffers = true)
+  }
+
+  private def planPhase(
+      grouping: Seq[NamedExpression],
+      aggregateExpressions: Seq[AggregateExpression],
+      resultExpressions: Seq[NamedExpression],
+      childPlan: SparkPlan,
+      unpackInputBuffers: Boolean,
+      packOutputBuffers: Boolean): Option[SparkPlan] = {
+    val rewrittenAggExprs = rewriteWrapperAggregates(aggregateExpressions)
+    if (rewrittenAggExprs.isEmpty) {
+      return None
+    }
+    val preparedChild = if (unpackInputBuffers) {
+      unpackInputBufferFields(childPlan, aggregateExpressions, 
rewrittenAggExprs)
+    } else {
+      childPlan
+    }
+    if (packOutputBuffers) {
+      planBufferPhase(
+        grouping,
+        aggregateExpressions,
+        resultExpressions,
+        preparedChild,
+        rewrittenAggExprs)
+    } else {
+      planFinalOutput(grouping, resultExpressions, preparedChild, 
rewrittenAggExprs)
+    }
+  }
+
+  private def unpackInputBufferFields(
+      childPlan: SparkPlan,
+      aggregateExpressions: Seq[AggregateExpression],
+      rewrittenAggExprs: Seq[AggregateExpression]): SparkPlan = {
+    // Recreate the wrapped aggregate's physical input-buffer attributes from 
the struct payload.
+    val unpackAliases = ArrayBuffer.empty[Alias]
+    val seenExprIds = scala.collection.mutable.HashSet.empty[Long]
+    rewrittenAggExprs.zip(aggregateExpressions).foreach {
+      case (rewrittenAe, AggregateExpression(wrapper: 
JoinAggregateFunctionWrapper, _, _, _, _)) =>
+        val bufferExpr = wrapper.children.head
+        
rewrittenAe.aggregateFunction.inputAggBufferAttributes.zipWithIndex.foreach {
+          case (bufferAttr, index) if seenExprIds.add(bufferAttr.exprId.id) =>
+            unpackAliases += Alias(
+              GetStructField(bufferExpr, index, Some(bufferAttr.name)),
+              s"_joinagg_buf_${bufferAttr.exprId.id}_$index"
+            )(exprId = bufferAttr.exprId, qualifier = bufferAttr.qualifier)
+          case _ =>
+        }
+      case _ =>
+    }
+    if (unpackAliases.nonEmpty) {
+      ProjectExec(childPlan.output ++ unpackAliases, childPlan)
+    } else {
+      childPlan
+    }
+  }
+
+  private def rewriteWrapperAggregates(
+      aggregateExpressions: Seq[AggregateExpression]): 
Seq[AggregateExpression] = {
+    aggregateExpressions.map {
       case ae @ AggregateExpression(_: JoinAggregateFunctionWrapper, _, _, _, 
_) =>
         rewriteSingleAggregateExpression(ae)
       case ae =>
         ae
     }
-    if (rewrittenAggExprs.isEmpty) {
-      return None
-    }
+  }
 
+  private def planBufferPhase(
+      grouping: Seq[NamedExpression],
+      aggregateExpressions: Seq[AggregateExpression],
+      resultExpressions: Seq[NamedExpression],
+      childPlan: SparkPlan,
+      rewrittenAggExprs: Seq[AggregateExpression]): Option[SparkPlan] = {
     val hashAgg = HashAggregateExec(
       requiredChildDistributionExpressions = None,
       isStreaming = false,
@@ -199,50 +281,20 @@ case class ImplementJoinAggregate(spark: SparkSession) 
extends SparkStrategy {
       aggregateExpressions: Seq[AggregateExpression],
       resultExpressions: Seq[NamedExpression],
       childPlan: SparkPlan): Option[SparkPlan] = {
-    // Lower the final wrapper phase by first unpacking the wrapper struct 
into the wrapped
-    // aggregate's input buffer attributes, then running a normal Spark final 
/ merge aggregate.
-    val wrapperWithRewritten: Seq[(JoinAggregateFunctionWrapper, 
AggregateExpression)] =
-      aggregateExpressions.flatMap {
-        case originalAe @ AggregateExpression(wrapper: 
JoinAggregateFunctionWrapper, _, _, _, _) =>
-          Some((wrapper, rewriteSingleAggregateExpression(originalAe)))
-        case _ =>
-          None
-      }
-
-    val unpackAliases = ArrayBuffer.empty[Alias]
-    val seenExprIds = scala.collection.mutable.HashSet.empty[Long]
-    wrapperWithRewritten.foreach {
-      case (wrapper, rewrittenAe) =>
-        val bufferExpr = wrapper.children.head
-        
rewrittenAe.aggregateFunction.inputAggBufferAttributes.zipWithIndex.foreach {
-          case (bufferAttr, idx) if seenExprIds.add(bufferAttr.exprId.id) =>
-            // Keep exprId for binding correctness, but avoid dotted names 
(e.g. a.b) in the
-            // temporary unpack projection. This projection only recreates the 
physical buffer attrs
-            // that Spark's final / merge aggregate expects to read from the 
wrapper struct.
-            val safeName = s"_joinagg_buf_${bufferAttr.exprId.id}_$idx"
-            unpackAliases += Alias(
-              GetStructField(bufferExpr, idx, Some(bufferAttr.name)),
-              safeName
-            )(exprId = bufferAttr.exprId, qualifier = bufferAttr.qualifier)
-          case _ =>
-        }
-    }
-
-    val childWithUnpacked = if (unpackAliases.nonEmpty) {
-      ProjectExec(childPlan.output ++ unpackAliases, childPlan)
-    } else {
-      childPlan
-    }
+    planPhase(
+      grouping,
+      aggregateExpressions,
+      resultExpressions,
+      childPlan,
+      unpackInputBuffers = true,
+      packOutputBuffers = false)
+  }
 
-    val rewrittenAggExprs = aggregateExpressions.map {
-      case ae @ AggregateExpression(_: JoinAggregateFunctionWrapper, _, _, _, 
_) =>
-        rewriteSingleAggregateExpression(ae)
-      case ae =>
-        ae
-    }
-    if (rewrittenAggExprs.isEmpty) {
-      return None
-    }
+  private def planFinalOutput(
+      grouping: Seq[NamedExpression],
+      resultExpressions: Seq[NamedExpression],
+      childPlan: SparkPlan,
+      rewrittenAggExprs: Seq[AggregateExpression]): Option[SparkPlan] = {
     val aggregateAttrs = rewrittenAggExprs.map(_.resultAttribute)
     val rewrittenResultExpressions =
       rewriteResultAsAggregateAttributes(resultExpressions, rewrittenAggExprs)
@@ -257,7 +309,7 @@ case class ImplementJoinAggregate(spark: SparkSession) 
extends SparkStrategy {
         aggregateAttributes = aggregateAttrs,
         initialInputBufferOffset = 0,
         resultExpressions = rewrittenResultExpressions,
-        child = childWithUnpacked
+        child = childPlan
       ))
   }
 
diff --git 
a/gluten-substrait/src/main/scala/org/apache/gluten/extension/joinagg/JoinAggregateFunctionWrapper.scala
 
b/gluten-substrait/src/main/scala/org/apache/gluten/extension/joinagg/JoinAggregateFunctionWrapper.scala
index 1474bdf2ba..01a40bc989 100644
--- 
a/gluten-substrait/src/main/scala/org/apache/gluten/extension/joinagg/JoinAggregateFunctionWrapper.scala
+++ 
b/gluten-substrait/src/main/scala/org/apache/gluten/extension/joinagg/JoinAggregateFunctionWrapper.scala
@@ -25,9 +25,10 @@ import java.util.Locale
 import scala.collection.mutable
 
 object JoinAggregateFunctionWrapper {
-  // The wrapper is used in exactly two logical phases:
-  //   - PartialPhase: a pushed aggregate below / through joins
-  //   - FinalPhase: the aggregate above the join that restores the original 
query semantics
+  // The wrapper is used in three logical phases:
+  //   - PartialPhase: aggregate raw input below a join
+  //   - PartialMergePhase: merge buffers after each pushed join edge
+  //   - FinalPhase: restore the original aggregate result above the joins
   sealed trait TargetPhase {
     def sqlName: String
   }
@@ -36,6 +37,10 @@ object JoinAggregateFunctionWrapper {
     override val sqlName: String = "PARTIAL"
   }
 
+  case object PartialMergePhase extends TargetPhase {
+    override val sqlName: String = "PARTIAL_MERGE"
+  }
+
   case object FinalPhase extends TargetPhase {
     override val sqlName: String = "FINAL"
   }
@@ -50,6 +55,17 @@ object JoinAggregateFunctionWrapper {
       wrapperKey = wrapperKey)
   }
 
+  def wrapperPartialMerge(
+      innerAgg: DeclarativeAggregate,
+      inputBuffer: Expression,
+      wrapperKey: String = "0"): JoinAggregateFunctionWrapper = {
+    JoinAggregateFunctionWrapper(
+      innerAgg = innerAgg,
+      targetPhase = PartialMergePhase,
+      inputBuffer = Some(inputBuffer),
+      wrapperKey = wrapperKey)
+  }
+
   def wrapperFinal(
       innerAgg: DeclarativeAggregate,
       inputBuffer: Expression,
@@ -74,6 +90,7 @@ object JoinAggregateFunctionWrapper {
       case (PartialMerge, PartialPhase) => PartialMerge
       case (Final, PartialPhase) => PartialMerge
       case (Complete, PartialPhase) => Partial
+      case (_, PartialMergePhase) => PartialMerge
       case (Partial, FinalPhase) => PartialMerge
       case (PartialMerge, FinalPhase) => PartialMerge
       case (Final, FinalPhase) => Final
@@ -103,6 +120,7 @@ case class JoinAggregateFunctionWrapper(
    *
    * The wrapper therefore changes only the *logical contract* across the join:
    *   - PartialPhase exposes the wrapped aggregate buffer as a single 
struct-valued output.
+   *   - PartialMergePhase merges one or more struct-valued buffers into 
another buffer.
    *   - FinalPhase consumes that struct-valued buffer and delegates merge / 
evaluate semantics
    *     back to the wrapped Spark aggregate.
    *
@@ -123,7 +141,7 @@ case class JoinAggregateFunctionWrapper(
   override lazy val nullable: Boolean = true
 
   override lazy val dataType: DataType = targetPhase match {
-    case PartialPhase =>
+    case PartialPhase | PartialMergePhase =>
       // The pushed phase carries the aggregate buffer through the plan as a 
single struct-valued
       // payload so the join sees one logical column per pushed aggregate.
       CreateStruct(wrappedBufferAttrs).dataType
@@ -133,7 +151,7 @@ case class JoinAggregateFunctionWrapper(
 
   override def children: Seq[Expression] = targetPhase match {
     case PartialPhase => innerAgg.children
-    case FinalPhase => Seq(outputBufferExpr)
+    case PartialMergePhase | FinalPhase => Seq(outputBufferExpr)
   }
 
   override lazy val aggBufferAttributes: Seq[AttributeReference] = 
wrappedBufferAttrs
@@ -150,7 +168,7 @@ case class JoinAggregateFunctionWrapper(
         childReplacements = innerAgg.children.zip(children).toMap,
         useInputBufferField = false
       )
-    case FinalPhase =>
+    case PartialMergePhase | FinalPhase =>
       // Merge expressions read from the struct-valued input buffer produced 
by the pushed phase.
       rewrite(innerAgg.mergeExpressions, childReplacements = Map.empty, 
useInputBufferField = true)
   }
@@ -160,7 +178,7 @@ case class JoinAggregateFunctionWrapper(
   }
 
   override lazy val evaluateExpression: Expression = targetPhase match {
-    case PartialPhase =>
+    case PartialPhase | PartialMergePhase =>
       // The pushed phase returns the entire aggregate buffer, not the final 
aggregate value.
       CreateStruct(aggBufferAttributes)
     case FinalPhase =>
@@ -182,7 +200,7 @@ case class JoinAggregateFunctionWrapper(
   override lazy val deterministic: Boolean = innerAgg.deterministic
 
   override lazy val defaultResult: Option[Literal] = targetPhase match {
-    case PartialPhase => None
+    case PartialPhase | PartialMergePhase => None
     case FinalPhase => innerAgg.defaultResult
   }
 
@@ -192,10 +210,10 @@ case class JoinAggregateFunctionWrapper(
       case PartialPhase =>
         val newInner = 
innerAgg.withNewChildren(newChildren).asInstanceOf[DeclarativeAggregate]
         copy(innerAgg = newInner, inputBuffer = None)
-      case FinalPhase =>
+      case PartialMergePhase | FinalPhase =>
         if (newChildren.size != 1) {
           throw new IllegalArgumentException(
-            s"Final JoinAggregateWrapper expects exactly one child, got 
${newChildren.size}")
+            s"$targetPhase JoinAggregateWrapper expects exactly one child, got 
${newChildren.size}")
         }
         copy(inputBuffer = Some(newChildren.head))
     }
diff --git 
a/gluten-substrait/src/main/scala/org/apache/gluten/extension/joinagg/PushAggregateThroughJoin.scala
 
b/gluten-substrait/src/main/scala/org/apache/gluten/extension/joinagg/PushAggregateThroughJoin.scala
index bb115de796..14d652fddc 100644
--- 
a/gluten-substrait/src/main/scala/org/apache/gluten/extension/joinagg/PushAggregateThroughJoin.scala
+++ 
b/gluten-substrait/src/main/scala/org/apache/gluten/extension/joinagg/PushAggregateThroughJoin.scala
@@ -91,12 +91,12 @@ case class PushAggregateThroughJoin(spark: SparkSession)
             !hasDistinctAggExpr(agg.aggregateExpressions) =>
         // 1) Aggregate+Join => FinalWrapperAgg(PartialWrapperAgg(...Join...))
         splitAggregate(agg) match {
-          case Some(newAgg) =>
+          case Some((finalAgg, lowerPartialAgg)) =>
             splitCount += 1
             // 2) Exhaustively push PartialWrapperAgg through join edges.
-            val pushed = pushPartialWrapperAggregate(newAgg)
+            val pushed = pushPartialWrapperAggregate(lowerPartialAgg)
             // 3) Return rewritten plan with pushed partial wrapper aggregates.
-            pushed
+            finalAgg.copy(child = pushed)
           case None => agg
         }
     }
@@ -113,7 +113,10 @@ case class PushAggregateThroughJoin(spark: SparkSession)
 
   private def maxDepth: Int = GlutenConfig.get.pushAggregateThroughJoinMaxDepth
 
-  private def splitAggregate(agg: Aggregate): Option[Aggregate] = {
+  private def partialMergeEnabled: Boolean =
+    GlutenConfig.get.pushAggregateThroughJoinPartialMergeEnabled
+
+  private def splitAggregate(agg: Aggregate): Option[(Aggregate, Aggregate)] = 
{
     // Split is intentionally child-agnostic. It only rewrites:
     //
     //   Aggregate(resultExprs, child)
@@ -166,29 +169,30 @@ case class PushAggregateThroughJoin(spark: SparkSession)
 
       rewriteAggregateExpressions(agg.aggregateExpressions, partialRefs).map {
         rewrittenAggExprs =>
-          agg.copy(
-            aggregateExpressions = rewrittenAggExprs,
-            child = partialAgg
-          )
+          (agg.copy(aggregateExpressions = rewrittenAggExprs, child = 
partialAgg), partialAgg)
       }
     }
   }
 
-  private def pushPartialWrapperAggregate(agg: Aggregate): LogicalPlan = {
+  private def pushPartialWrapperAggregate(lowerPartialAgg: Aggregate): 
LogicalPlan = {
     // Push one join edge per iteration. This keeps the rewrite local and lets 
`maxDepth` bound
     // how far a pushed aggregate is allowed to travel through a multi-join 
subtree.
-    var current: LogicalPlan = agg
+    var current: LogicalPlan = lowerPartialAgg
     var changed = true
     var pushCount = 0
     while (changed && pushCount < maxDepth) {
       changed = false
       current = current.transformUp {
         case partialAgg: Aggregate if 
isPurePartialWrapperAggregate(partialAgg) =>
-          pushOnce(
-            partialAgg,
-            partialAgg.groupingExpressions,
-            partialAgg.aggregateExpressions,
-            partialAgg.child) match {
+          val pushed = if (partialMergeEnabled) {
+            val (partialMergeAgg, freshLowerPartialAgg) = 
splitPartialAggregateForMerge(partialAgg)
+            pushOnce(freshLowerPartialAgg).map {
+              pushedPlan => partialMergeAgg.copy(child = pushedPlan)
+            }
+          } else {
+            pushOnce(partialAgg)
+          }
+          pushed match {
             case Some(newPlan) =>
               pushCount += 1
               changed = true
@@ -202,22 +206,17 @@ case class PushAggregateThroughJoin(spark: SparkSession)
     current
   }
 
-  private def pushOnce(
-      partialAgg: Aggregate,
-      groupingExprs: Seq[Expression],
-      aggExprs: Seq[NamedExpression],
-      child: LogicalPlan): Option[LogicalPlan] = {
-    extractJoin(child).flatMap {
+  private def pushOnce(partialAgg: Aggregate): Option[LogicalPlan] = {
+    extractJoin(partialAgg.child).flatMap {
       case (join, wrapperRequiredAttrs, rebuild) =>
         Seq(JoinLeft, JoinRight).iterator
           .flatMap {
             side =>
               val maybePushedJoin =
-                pushPartialAggToJoinSide(join, groupingExprs, aggExprs, 
wrapperRequiredAttrs, side)
+                pushPartialAggToJoinSide(partialAgg, join, 
wrapperRequiredAttrs, side)
               maybePushedJoin match {
                 case Some(pushedJoin) =>
-                  val requiredAttrs = partialAgg.output.collect { case a: 
Attribute => a }
-                  Some(rebuild(pushedJoin, requiredAttrs))
+                  Some(rebuild(pushedJoin, partialAgg.output.collect { case a: 
Attribute => a }))
                 case None => None
               }
           }
@@ -227,9 +226,8 @@ case class PushAggregateThroughJoin(spark: SparkSession)
   }
 
   private def pushPartialAggToJoinSide(
+      partialAgg: Aggregate,
       join: Join,
-      groupingExprs: Seq[Expression],
-      aggExprs: Seq[NamedExpression],
       wrapperRequiredAttrs: Seq[Attribute],
       side: JoinSide): Option[Join] = {
     // A pushed wrapper aggregate may move to a join side only when all of its 
aggregate inputs
@@ -242,7 +240,7 @@ case class PushAggregateThroughJoin(spark: SparkSession)
     //
     // The pushed grouping must *not* include pure measure inputs of the 
pushed aggregate such as
     // `ss_net_profit`, otherwise the pre-aggregation becomes over-constrained 
and ineffective.
-    val wrapperAliases = collectPartialWrapperAliases(aggExprs)
+    val wrapperAliases = 
collectPartialWrapperAliases(partialAgg.aggregateExpressions)
     if (wrapperAliases.isEmpty) {
       return None
     }
@@ -255,7 +253,7 @@ case class PushAggregateThroughJoin(spark: SparkSession)
       return None
     }
 
-    val sideGroupingAttrs = groupingExprs
+    val sideGroupingAttrs = partialAgg.groupingExpressions
       .flatMap(referencedAttrsInOrder)
       .collect { case a: Attribute if sideOutputSet.contains(a) => a }
     val sideJoinKeys = 
join.condition.toSeq.flatMap(splitConjunctivePredicates).collect {
@@ -273,7 +271,7 @@ case class PushAggregateThroughJoin(spark: SparkSession)
     // These attrs belong to aggregate subexpressions that stay above the 
pushed aggregate. They
     // must survive subtree rebuild, but they are not themselves proof that 
the pushed aggregate
     // needs to group by those measures.
-    val sideNonPushableAggAttrs = dedupeAttrs(aggExprs.flatMap {
+    val sideNonPushableAggAttrs = 
dedupeAttrs(partialAgg.aggregateExpressions.flatMap {
       case Alias(expr, _) if !isPushableExpr(expr) && 
!containsWrapperAggregateExpr(expr) =>
         referencedAttrsInOrder(expr).collect {
           case a: Attribute if sideOutputSet.contains(a) => a
@@ -312,18 +310,8 @@ case class PushAggregateThroughJoin(spark: SparkSession)
       return None
     }
 
-    val pushedWrapperAliases = wrapperAliases.map {
-      case (alias, wrapper) =>
-        val wrapped = JoinAggregateFunctionWrapper
-          .wrapperPartial(wrapper.innerAgg, wrapper.wrapperKey)
-          .toAggregateExpression()
-        Alias(wrapped, alias.name)(
-          exprId = alias.exprId,
-          qualifier = alias.qualifier,
-          explicitMetadata = alias.explicitMetadata,
-          nonInheritableMetadataKeys = alias.nonInheritableMetadataKeys
-        )
-    }
+    // Move the existing Partial expressions unchanged, including their 
ExprIds.
+    val pushedWrapperAliases = wrapperAliases.map(_._1)
 
     val pushedAgg = Aggregate(
       groupingExpressions = pushedGrouping,
@@ -335,6 +323,50 @@ case class PushAggregateThroughJoin(spark: SparkSession)
     Some(pushedJoin)
   }
 
+  // Retain the current buffer IDs above a fresh lower Partial. The lower 
Partial is then the
+  // only aggregate pushed through the next join edge.
+  private def splitPartialAggregateForMerge(partialAgg: Aggregate): 
(Aggregate, Aggregate) = {
+    val wrapperAliases = 
collectPartialWrapperAliases(partialAgg.aggregateExpressions)
+    val lowerAliases = wrapperAliases.map {
+      case (alias, wrapper) =>
+        Alias(
+          JoinAggregateFunctionWrapper
+            .wrapperPartial(wrapper.innerAgg, wrapper.wrapperKey)
+            .toAggregateExpression(),
+          alias.name)()
+    }
+    val lowerGroupingOutputs =
+      partialAgg.aggregateExpressions.take(partialAgg.groupingExpressions.size)
+    val lowerAgg = Aggregate(
+      groupingExpressions = partialAgg.groupingExpressions,
+      aggregateExpressions = lowerGroupingOutputs ++ lowerAliases,
+      child = partialAgg.child)
+    val lowerBuffersByExprId = wrapperAliases.zip(
+      lowerAgg.output.drop(lowerGroupingOutputs.size))
+      .map { case ((alias, _), attr) => alias.exprId.id -> cleanAttr(attr) }
+      .toMap
+    val partialMergeExpressions = partialAgg.aggregateExpressions.map {
+      case alias @ Alias(
+            AggregateExpression(wrapper: JoinAggregateFunctionWrapper, _, _, 
_, _),
+            _)
+          if wrapper.targetPhase == JoinAggregateFunctionWrapper.PartialPhase 
=>
+        val lowerBuffer = lowerBuffersByExprId.getOrElse(
+          alias.exprId.id,
+          throw new IllegalStateException(s"Cannot resolve pushed buffer for 
${alias.sql}"))
+        val partialMerge = JoinAggregateFunctionWrapper
+          .wrapperPartialMerge(wrapper.innerAgg, lowerBuffer, 
wrapper.wrapperKey)
+          .toAggregateExpression()
+        Alias(partialMerge, alias.name)(
+          exprId = alias.exprId,
+          qualifier = alias.qualifier,
+          explicitMetadata = alias.explicitMetadata,
+          nonInheritableMetadataKeys = alias.nonInheritableMetadataKeys
+        )
+      case other => other
+    }
+    (partialAgg.copy(aggregateExpressions = partialMergeExpressions, child = 
lowerAgg), lowerAgg)
+  }
+
   private def isPurePartialWrapperAggregate(agg: Aggregate): Boolean = {
     val wrapperAliases = collectPartialWrapperAliases(agg.aggregateExpressions)
     wrapperAliases.nonEmpty && agg.aggregateExpressions.forall {
diff --git 
a/gluten-substrait/src/test/scala/org/apache/gluten/execution/PushAggregateThroughJoinSuite.scala
 
b/gluten-substrait/src/test/scala/org/apache/gluten/execution/PushAggregateThroughJoinSuite.scala
index aac2376429..44e654e790 100644
--- 
a/gluten-substrait/src/test/scala/org/apache/gluten/execution/PushAggregateThroughJoinSuite.scala
+++ 
b/gluten-substrait/src/test/scala/org/apache/gluten/execution/PushAggregateThroughJoinSuite.scala
@@ -37,7 +37,7 @@ class PushAggregateThroughJoinSuite extends PlanTest with 
SharedSparkSession {
   private val joinAggregateRule = PushAggregateThroughJoin(spark)
   private val debugMode: Boolean = true
 
-  private case class PushdownCase(inputSql: String, expectedAggCount: Int)
+  private case class PushdownCase(inputSql: String)
 
   override protected def sparkConf: SparkConf = {
     // Avoid Janino projection codegen here because Spark 4's 
QueryExecutionErrors
@@ -92,10 +92,15 @@ class PushAggregateThroughJoinSuite extends PlanTest with 
SharedSparkSession {
   private def runCaseWithMaxDepth(
       testCase: PushdownCase,
       maxDepth: Int,
-      expectedPushCount: Int): Unit = {
+      expectedPushCount: Int,
+      expectedAggCount: Int,
+      partialMergeEnabled: Boolean): Unit = {
     withSQLConf(
       GlutenConfig.PUSH_AGGREGATE_THROUGH_JOIN_ENABLED.key -> "true",
-      GlutenConfig.PUSH_AGGREGATE_THROUGH_JOIN_MAX_DEPTH.key -> 
maxDepth.toString) {
+      GlutenConfig.PUSH_AGGREGATE_THROUGH_JOIN_MAX_DEPTH.key -> 
maxDepth.toString,
+      GlutenConfig.PUSH_AGGREGATE_THROUGH_JOIN_PARTIAL_MERGE_ENABLED.key ->
+        partialMergeEnabled.toString
+    ) {
       val (withoutRuleRows, withoutRuleLogicalPlan, withoutRulePhysicalPlan) =
         withExtraPlanning(Nil, Nil) {
           val df = spark.sql(testCase.inputSql)
@@ -128,7 +133,7 @@ class PushAggregateThroughJoinSuite extends PlanTest with 
SharedSparkSession {
                 .map(_.treeString)
                 .mkString("\n---\n")}")
           assert(joinAggregateRule.getSuccessfulPushCount == expectedPushCount)
-          assert(aggregateNodeCount == testCase.expectedAggCount)
+          assert(aggregateNodeCount == expectedAggCount)
           (withRuleRows, withRulePlan, withRulePhysicalPlan)
         }
 
@@ -213,10 +218,20 @@ class PushAggregateThroughJoinSuite extends PlanTest with 
SharedSparkSession {
                    |WHERE d_year IN (1999, 2000, 2001, 2002)
                    |GROUP BY substring(i_item_desc, 1, 30), i_item_sk, d_date
                    |HAVING count(1) > 4
-                   |""".stripMargin,
-      expectedAggCount = 2
+                   |""".stripMargin
     )
-    runCaseWithMaxDepth(pushdownCase, maxDepth = Int.MaxValue, 
expectedPushCount = 2)
+    runCaseWithMaxDepth(
+      pushdownCase,
+      maxDepth = Int.MaxValue,
+      expectedPushCount = 2,
+      expectedAggCount = 2,
+      partialMergeEnabled = false)
+    runCaseWithMaxDepth(
+      pushdownCase,
+      maxDepth = Int.MaxValue,
+      expectedPushCount = 2,
+      expectedAggCount = 4,
+      partialMergeEnabled = true)
   }
 
   test("pre-aggregate store_sales for sum") {
@@ -228,10 +243,20 @@ class PushAggregateThroughJoinSuite extends PlanTest with 
SharedSparkSession {
                    |FROM store_sales
                    |JOIN item ON ss_item_sk = i_item_sk
                    |GROUP BY i_item_sk
-                   |""".stripMargin,
-      expectedAggCount = 2
+                   |""".stripMargin
     )
-    runCaseWithMaxDepth(pushdownCase, maxDepth = Int.MaxValue, 
expectedPushCount = 1)
+    runCaseWithMaxDepth(
+      pushdownCase,
+      maxDepth = Int.MaxValue,
+      expectedPushCount = 1,
+      expectedAggCount = 2,
+      partialMergeEnabled = false)
+    runCaseWithMaxDepth(
+      pushdownCase,
+      maxDepth = Int.MaxValue,
+      expectedPushCount = 1,
+      expectedAggCount = 3,
+      partialMergeEnabled = true)
   }
 
   test("pre-aggregate store_sales for avg") {
@@ -243,10 +268,20 @@ class PushAggregateThroughJoinSuite extends PlanTest with 
SharedSparkSession {
                    |FROM store_sales
                    |JOIN item ON ss_item_sk = i_item_sk
                    |GROUP BY i_item_sk
-                   |""".stripMargin,
-      expectedAggCount = 2
+                   |""".stripMargin
     )
-    runCaseWithMaxDepth(pushdownCase, maxDepth = Int.MaxValue, 
expectedPushCount = 1)
+    runCaseWithMaxDepth(
+      pushdownCase,
+      maxDepth = Int.MaxValue,
+      expectedPushCount = 1,
+      expectedAggCount = 2,
+      partialMergeEnabled = false)
+    runCaseWithMaxDepth(
+      pushdownCase,
+      maxDepth = Int.MaxValue,
+      expectedPushCount = 1,
+      expectedAggCount = 3,
+      partialMergeEnabled = true)
   }
 
   test("pre-aggregate store_sales for sum on fact table") {
@@ -258,10 +293,20 @@ class PushAggregateThroughJoinSuite extends PlanTest with 
SharedSparkSession {
                    |FROM store_sales
                    |JOIN item ON ss_item_sk = i_item_sk
                    |GROUP BY ss_sold_date_sk
-                   |""".stripMargin,
-      expectedAggCount = 2
+                   |""".stripMargin
     )
-    runCaseWithMaxDepth(pushdownCase, maxDepth = Int.MaxValue, 
expectedPushCount = 1)
+    runCaseWithMaxDepth(
+      pushdownCase,
+      maxDepth = Int.MaxValue,
+      expectedPushCount = 1,
+      expectedAggCount = 2,
+      partialMergeEnabled = false)
+    runCaseWithMaxDepth(
+      pushdownCase,
+      maxDepth = Int.MaxValue,
+      expectedPushCount = 1,
+      expectedAggCount = 3,
+      partialMergeEnabled = true)
   }
 
   test("pre-aggregate store_sales for avg on fact table") {
@@ -273,10 +318,20 @@ class PushAggregateThroughJoinSuite extends PlanTest with 
SharedSparkSession {
                    |FROM store_sales
                    |JOIN item ON ss_item_sk = i_item_sk
                    |GROUP BY ss_sold_date_sk
-                   |""".stripMargin,
-      expectedAggCount = 2
+                   |""".stripMargin
     )
-    runCaseWithMaxDepth(pushdownCase, maxDepth = Int.MaxValue, 
expectedPushCount = 1)
+    runCaseWithMaxDepth(
+      pushdownCase,
+      maxDepth = Int.MaxValue,
+      expectedPushCount = 1,
+      expectedAggCount = 2,
+      partialMergeEnabled = false)
+    runCaseWithMaxDepth(
+      pushdownCase,
+      maxDepth = Int.MaxValue,
+      expectedPushCount = 1,
+      expectedAggCount = 3,
+      partialMergeEnabled = true)
   }
 
   test("pre-aggregate store_sales for sum on three-way join") {
@@ -290,10 +345,20 @@ class PushAggregateThroughJoinSuite extends PlanTest with 
SharedSparkSession {
                    |JOIN date_dim ON ss_sold_date_sk = d_date_sk
                    |JOIN item ON ss_item_sk = i_item_sk
                    |GROUP BY item_desc, d_date
-                   |""".stripMargin,
-      expectedAggCount = 2
+                   |""".stripMargin
     )
-    runCaseWithMaxDepth(pushdownCase, maxDepth = Int.MaxValue, 
expectedPushCount = 2)
+    runCaseWithMaxDepth(
+      pushdownCase,
+      maxDepth = Int.MaxValue,
+      expectedPushCount = 2,
+      expectedAggCount = 2,
+      partialMergeEnabled = false)
+    runCaseWithMaxDepth(
+      pushdownCase,
+      maxDepth = Int.MaxValue,
+      expectedPushCount = 2,
+      expectedAggCount = 4,
+      partialMergeEnabled = true)
   }
 
   test("pre-aggregate store_sales for sum and avg on different fact columns on 
three-way join") {
@@ -308,10 +373,20 @@ class PushAggregateThroughJoinSuite extends PlanTest with 
SharedSparkSession {
                    |JOIN date_dim ON ss_sold_date_sk = d_date_sk
                    |JOIN item ON ss_item_sk = i_item_sk
                    |GROUP BY item_desc, d_date
-                   |""".stripMargin,
-      expectedAggCount = 2
+                   |""".stripMargin
     )
-    runCaseWithMaxDepth(pushdownCase, maxDepth = Int.MaxValue, 
expectedPushCount = 2)
+    runCaseWithMaxDepth(
+      pushdownCase,
+      maxDepth = Int.MaxValue,
+      expectedPushCount = 2,
+      expectedAggCount = 2,
+      partialMergeEnabled = false)
+    runCaseWithMaxDepth(
+      pushdownCase,
+      maxDepth = Int.MaxValue,
+      expectedPushCount = 2,
+      expectedAggCount = 4,
+      partialMergeEnabled = true)
   }
 
   test("pre-aggregate store_sales for sum and avg on same fact column on 
three-way join") {
@@ -326,10 +401,20 @@ class PushAggregateThroughJoinSuite extends PlanTest with 
SharedSparkSession {
                    |JOIN date_dim ON ss_sold_date_sk = d_date_sk
                    |JOIN item ON ss_item_sk = i_item_sk
                    |GROUP BY item_desc, d_date
-                   |""".stripMargin,
-      expectedAggCount = 2
+                   |""".stripMargin
     )
-    runCaseWithMaxDepth(pushdownCase, maxDepth = Int.MaxValue, 
expectedPushCount = 2)
+    runCaseWithMaxDepth(
+      pushdownCase,
+      maxDepth = Int.MaxValue,
+      expectedPushCount = 2,
+      expectedAggCount = 2,
+      partialMergeEnabled = false)
+    runCaseWithMaxDepth(
+      pushdownCase,
+      maxDepth = Int.MaxValue,
+      expectedPushCount = 2,
+      expectedAggCount = 4,
+      partialMergeEnabled = true)
   }
 
   test("pre-aggregate store_sales by i_item_desc") {
@@ -341,10 +426,20 @@ class PushAggregateThroughJoinSuite extends PlanTest with 
SharedSparkSession {
                    |FROM store_sales
                    |JOIN item ON ss_item_sk = i_item_sk
                    |GROUP BY item_desc
-                   |""".stripMargin,
-      expectedAggCount = 2
+                   |""".stripMargin
     )
-    runCaseWithMaxDepth(pushdownCase, maxDepth = Int.MaxValue, 
expectedPushCount = 1)
+    runCaseWithMaxDepth(
+      pushdownCase,
+      maxDepth = Int.MaxValue,
+      expectedPushCount = 1,
+      expectedAggCount = 2,
+      partialMergeEnabled = false)
+    runCaseWithMaxDepth(
+      pushdownCase,
+      maxDepth = Int.MaxValue,
+      expectedPushCount = 1,
+      expectedAggCount = 3,
+      partialMergeEnabled = true)
   }
 
   test("pre-aggregate store_sales by substr(i_item_desc, 3), 3 ways") {
@@ -358,10 +453,20 @@ class PushAggregateThroughJoinSuite extends PlanTest with 
SharedSparkSession {
                    |JOIN date_dim ON ss_sold_date_sk = d_date_sk
                    |JOIN item ON ss_item_sk = i_item_sk
                    |GROUP BY d_date, item_desc
-                   |""".stripMargin,
-      expectedAggCount = 2
+                   |""".stripMargin
     )
-    runCaseWithMaxDepth(pushdownCase, maxDepth = Int.MaxValue, 
expectedPushCount = 2)
+    runCaseWithMaxDepth(
+      pushdownCase,
+      maxDepth = Int.MaxValue,
+      expectedPushCount = 2,
+      expectedAggCount = 2,
+      partialMergeEnabled = false)
+    runCaseWithMaxDepth(
+      pushdownCase,
+      maxDepth = Int.MaxValue,
+      expectedPushCount = 2,
+      expectedAggCount = 4,
+      partialMergeEnabled = true)
   }
 
   test("pre-aggregate store_sales for sum with item filter") {
@@ -372,10 +477,20 @@ class PushAggregateThroughJoinSuite extends PlanTest with 
SharedSparkSession {
                    |FROM store_sales
                    |JOIN item ON ss_item_sk = i_item_sk
                    |WHERE i_category_id IN (1, 2, 3, 4, 5)
-                   |""".stripMargin,
-      expectedAggCount = 2
+                   |""".stripMargin
     )
-    runCaseWithMaxDepth(pushdownCase, maxDepth = Int.MaxValue, 
expectedPushCount = 1)
+    runCaseWithMaxDepth(
+      pushdownCase,
+      maxDepth = Int.MaxValue,
+      expectedPushCount = 1,
+      expectedAggCount = 2,
+      partialMergeEnabled = false)
+    runCaseWithMaxDepth(
+      pushdownCase,
+      maxDepth = Int.MaxValue,
+      expectedPushCount = 1,
+      expectedAggCount = 3,
+      partialMergeEnabled = true)
   }
 
   test("pre-aggregate three-way joins independently under union all") {
@@ -412,12 +527,44 @@ class PushAggregateThroughJoinSuite extends PlanTest with 
SharedSparkSession {
           |  JOIN item ON ss_item_sk = i_item_sk
           |  GROUP BY concat('year-', cast(d_year AS string), '-', 
cast(i_item_sk AS string))
           |)
-          |""".stripMargin,
-      expectedAggCount = 6
+          |""".stripMargin
     )
-    runCaseWithMaxDepth(pushdownCase, maxDepth = 1, expectedPushCount = 3)
-    runCaseWithMaxDepth(pushdownCase, maxDepth = 2, expectedPushCount = 6)
-    runCaseWithMaxDepth(pushdownCase, maxDepth = Int.MaxValue, 
expectedPushCount = 6)
+    runCaseWithMaxDepth(
+      pushdownCase,
+      maxDepth = 1,
+      expectedPushCount = 3,
+      expectedAggCount = 6,
+      partialMergeEnabled = false)
+    runCaseWithMaxDepth(
+      pushdownCase,
+      maxDepth = 1,
+      expectedPushCount = 3,
+      expectedAggCount = 9,
+      partialMergeEnabled = true)
+    runCaseWithMaxDepth(
+      pushdownCase,
+      maxDepth = 2,
+      expectedPushCount = 6,
+      expectedAggCount = 6,
+      partialMergeEnabled = false)
+    runCaseWithMaxDepth(
+      pushdownCase,
+      maxDepth = 2,
+      expectedPushCount = 6,
+      expectedAggCount = 12,
+      partialMergeEnabled = true)
+    runCaseWithMaxDepth(
+      pushdownCase,
+      maxDepth = Int.MaxValue,
+      expectedPushCount = 6,
+      expectedAggCount = 6,
+      partialMergeEnabled = false)
+    runCaseWithMaxDepth(
+      pushdownCase,
+      maxDepth = Int.MaxValue,
+      expectedPushCount = 6,
+      expectedAggCount = 12,
+      partialMergeEnabled = true)
   }
 
   test("pre-aggregate store_sales for sum on three-way join with maxDepth=1 / 
maxDepth=2") {
@@ -431,12 +578,44 @@ class PushAggregateThroughJoinSuite extends PlanTest with 
SharedSparkSession {
                    |JOIN date_dim ON ss_sold_date_sk = d_date_sk
                    |JOIN item ON ss_item_sk = i_item_sk
                    |GROUP BY item_desc, d_date
-                   |""".stripMargin,
-      expectedAggCount = 2
+                   |""".stripMargin
     )
-    runCaseWithMaxDepth(pushdownCase, maxDepth = 1, expectedPushCount = 1)
-    runCaseWithMaxDepth(pushdownCase, maxDepth = 2, expectedPushCount = 2)
-    runCaseWithMaxDepth(pushdownCase, maxDepth = Int.MaxValue, 
expectedPushCount = 2)
+    runCaseWithMaxDepth(
+      pushdownCase,
+      maxDepth = 1,
+      expectedPushCount = 1,
+      expectedAggCount = 2,
+      partialMergeEnabled = false)
+    runCaseWithMaxDepth(
+      pushdownCase,
+      maxDepth = 1,
+      expectedPushCount = 1,
+      expectedAggCount = 3,
+      partialMergeEnabled = true)
+    runCaseWithMaxDepth(
+      pushdownCase,
+      maxDepth = 2,
+      expectedPushCount = 2,
+      expectedAggCount = 2,
+      partialMergeEnabled = false)
+    runCaseWithMaxDepth(
+      pushdownCase,
+      maxDepth = 2,
+      expectedPushCount = 2,
+      expectedAggCount = 4,
+      partialMergeEnabled = true)
+    runCaseWithMaxDepth(
+      pushdownCase,
+      maxDepth = Int.MaxValue,
+      expectedPushCount = 2,
+      expectedAggCount = 2,
+      partialMergeEnabled = false)
+    runCaseWithMaxDepth(
+      pushdownCase,
+      maxDepth = Int.MaxValue,
+      expectedPushCount = 2,
+      expectedAggCount = 4,
+      partialMergeEnabled = true)
   }
 
   test("pre-aggregate with filter inside inner equi-join") {
@@ -448,9 +627,19 @@ class PushAggregateThroughJoinSuite extends PlanTest with 
SharedSparkSession {
                    |FROM store_sales
                    |JOIN item ON ss_item_sk = i_item_sk AND ss_quantity > 1
                    |GROUP BY i_item_sk
-                   |""".stripMargin,
-      expectedAggCount = 2
+                   |""".stripMargin
     )
-    runCaseWithMaxDepth(pushdownCase, maxDepth = Int.MaxValue, 
expectedPushCount = 1)
+    runCaseWithMaxDepth(
+      pushdownCase,
+      maxDepth = Int.MaxValue,
+      expectedPushCount = 1,
+      expectedAggCount = 2,
+      partialMergeEnabled = false)
+    runCaseWithMaxDepth(
+      pushdownCase,
+      maxDepth = Int.MaxValue,
+      expectedPushCount = 1,
+      expectedAggCount = 3,
+      partialMergeEnabled = true)
   }
 }


---------------------------------------------------------------------
To unsubscribe, e-mail: [email protected]
For additional commands, e-mail: [email protected]

Reply via email to