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]