Copilot commented on code in PR #13086:
URL: https://github.com/apache/gluten/pull/13086#discussion_r4080951572


##########
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.                                                
                                                                                
                                                                                
                                                                                
                                                              |

Review Comment:
   The PR description states this option is intended “when eager aggregation 
(experimental) is enabled”, but the documentation entry doesn’t mention that 
relationship/limitation. If partial-merge is only effective (or only supported) 
with eager aggregation, the docstring should state the prerequisite/interaction 
so users don’t enable it expecting behavior changes under the non-eager path.



##########
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)
       }

Review Comment:
   This PR introduces a new logical phase (`PartialMergePhase`) with a distinct 
physical planning path (unpack + merge + repack). The updated suites validate 
higher-level aggregate counts/semantics, but there’s no targeted assertion that 
`PartialMergePhase` is actually lowered as intended (e.g., that the unpack 
projection is present and that the planned aggregate mode corresponds to 
merge). Adding a focused test that inspects the physical plan for the 
partial-merge path would help prevent regressions in the unpack/pack logic and 
phase-to-mode mapping.



##########
gluten-substrait/src/main/scala/org/apache/gluten/extension/joinagg/JoinAggregateFunctionWrapper.scala:
##########
@@ -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)
+  }

Review Comment:
   `PartialMergePhase` (and `FinalPhase`) semantically require `inputBuffer` to 
be defined, but that invariant is only implied by helper constructors. Adding a 
hard validation (e.g., a `require(...)` in the case class constructor based on 
`targetPhase`, or stricter constructors) would prevent accidental construction 
of invalid wrappers that would otherwise fall back to 
`CreateStruct(innerAgg.inputAggBufferAttributes)` via `outputBufferExpr`, 
potentially masking logical bugs and producing hard-to-debug plans.



##########
gluten-substrait/src/main/scala/org/apache/gluten/extension/joinagg/ImplementJoinAggregate.scala:
##########
@@ -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) =>

Review Comment:
   `ExprId` in Spark includes more identity than just `exprId.id`; using `Long` 
here bakes in an assumption about uniqueness and makes the code harder to 
reason about. Consider tracking `ExprId` directly (or `bufferAttr.exprId`) in 
`seenExprIds` to better reflect intent and avoid accidental collisions if plan 
fragments ever combine ExprIds from different origins.



##########
gluten-substrait/src/main/scala/org/apache/gluten/extension/joinagg/PushAggregateThroughJoin.scala:
##########
@@ -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)

Review Comment:
   The `lowerGroupingOutputs` selection assumes the first 
`groupingExpressions.size` entries in `partialAgg.aggregateExpressions` are 
exactly the grouping outputs. That positional contract is fragile and can 
silently break if the aggregate output ordering changes (e.g., due to analyzer 
rewrites or a different construction path), causing incorrect output schemas or 
mis-binding in later phases. Prefer deriving the grouping output 
`NamedExpression`s explicitly (e.g., by matching the grouping expressions to 
corresponding output attributes/aliases via semantic equality or a dedicated 
helper that constructs grouping outputs), and/or add an assertion 
documenting/enforcing the ordering invariant at the point where `partialAgg` is 
created.



-- 
This is an automated message from the Apache Git Service.
To respond to the message, please log on to GitHub and use the
URL above to go to the specific comment.

To unsubscribe, e-mail: [email protected]

For queries about this service, please contact Infrastructure at:
[email protected]


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

Reply via email to