cloud-fan commented on code in PR #57742:
URL: https://github.com/apache/spark/pull/57742#discussion_r3711016977


##########
sql/core/src/main/scala/org/apache/spark/sql/execution/aggregate/TungstenAggregationIterator.scala:
##########
@@ -191,29 +208,62 @@ class TungstenAggregationIterator(
       }
     } else {
       var i = 0
-      while (inputIter.hasNext) {
+      var processedRows = 0L
+      // The next row count at which the no-spill tier re-evaluates the 
reduction ratio. It starts
+      // at `sampleRows` and doubles after each sub-threshold check, so the 
ratio is checked only
+      // rarely once the input proves low-cardinality.
+      var nextSampleRow = 
adaptivePartialAggConfig.map(_.sampleRows.toLong).getOrElse(0L)
+      while (inputIter.hasNext && !passThrough) {
         val newInput = inputIter.next()
         val groupingKey = groupingProjection.apply(newInput)
         var buffer: UnsafeRow = null
         if (i < fallbackStartsAt._2) {
           buffer = hashMap.getAggregationBufferFromUnsafeRow(groupingKey)
         }
         if (buffer == null) {
-          val sorter = hashMap.destructAndCreateExternalSorter()
-          if (externalSorter == null) {
-            externalSorter = sorter
+          // The map is full and would normally spill. On the first spill, 
adaptive partial
+          // aggregation may instead bypass: keep the in-memory map as-is, 
pass this row and all
+          // remaining rows through, and skip the spill entirely.
+          if (adaptivePartialAggConfig.isDefined && externalSorter == null && 
processedRows > 0 &&
+              hashMap.getNumKeys().toDouble >=
+                processedRows * 
adaptivePartialAggConfig.get.spillReductionRatioThreshold) {
+            passThrough = true
+            // `newInput` could not be inserted; stash a copy as the first 
pass-through row so it
+            // is not lost when we drain the rest of `inputIter`.
+            pendingPassThroughRow = newInput.copy()
           } else {
-            externalSorter.merge(sorter)
+            val sorter = hashMap.destructAndCreateExternalSorter()
+            if (externalSorter == null) {
+              externalSorter = sorter
+            } else {
+              externalSorter.merge(sorter)
+            }
+            i = 0
+            buffer = hashMap.getAggregationBufferFromUnsafeRow(groupingKey)
+            if (buffer == null) {
+              // failed to allocate the first page
+              throw QueryExecutionErrors.aggregateOutOfMemoryError()
+            }
           }
-          i = 0
-          buffer = hashMap.getAggregationBufferFromUnsafeRow(groupingKey)
-          if (buffer == null) {
-            // failed to allocate the first page
-            throw QueryExecutionErrors.aggregateOutOfMemoryError()
+        }
+        if (!passThrough) {
+          processRow(buffer, newInput)
+          i += 1
+          processedRows += 1
+          // No-spill tier: from the sampling window on, if the map is still 
fully in memory and
+          // the reduction ratio is too high to be worthwhile, bypass partial 
aggregation for the
+          // rest. The window doubles after each sub-threshold check so 
low-cardinality input is
+          // re-evaluated only rarely while a late high-cardinality tail can 
still be caught.
+          if (adaptivePartialAggConfig.isDefined && externalSorter == null &&

Review Comment:
   Extract stable primitive settings before entering the input loop. The 
current `Option.isDefined`/`get` checks repeat for every aggregated row even 
though eligibility and thresholds cannot change during the iterator's lifetime.



##########
sql/core/src/main/scala/org/apache/spark/sql/execution/aggregate/HashAggregateExec.scala:
##########
@@ -663,46 +855,122 @@ case class HashAggregateExec(
       case _ => ("true", "", "")
     }
 
-    val findOrInsertRegularHashMap: String =
-      s"""
-         |// generate grouping key
-         |${unsafeRowKeyCode.code}
-         |int $unsafeRowKeyHash = ${unsafeRowKeyCode.value}.hashCode();
-         |if ($checkFallbackForBytesToBytesMap) {
-         |  // try to get the buffer from hash map
-         |  $unsafeRowBuffer =
-         |    $hashMapTerm.getAggregationBufferFromUnsafeRow($unsafeRowKeys, 
$unsafeRowKeyHash);
-         |}
-         |// Can't allocate buffer from the hash map. Spill the map and 
fallback to sort-based
-         |// aggregation after processing all input rows.
-         |if ($unsafeRowBuffer == null) {
-         |  if ($sorterTerm == null) {
-         |    $sorterTerm = $hashMapTerm.destructAndCreateExternalSorter();
-         |  } else {
-         |    
$sorterTerm.merge($hashMapTerm.destructAndCreateExternalSorter());
-         |  }
-         |  $resetCounter
-         |  // the hash map had be spilled, it should have enough memory now,
-         |  // try to allocate buffer again.
-         |  $unsafeRowBuffer = $hashMapTerm.getAggregationBufferFromUnsafeRow(
-         |    $unsafeRowKeys, $unsafeRowKeyHash);
-         |  if ($unsafeRowBuffer == null) {
-         |    // failed to allocate the first page
-         |    throw QueryExecutionErrors.aggregateOutOfMemoryError();
-         |  }
-         |}
-       """.stripMargin
+    val findOrInsertRegularHashMap: String = {
+      // Assumes the grouping key projection (`unsafeRowKeyCode.code`) has 
already run for this row,
+      // so `unsafeRowKeyCode.value` holds the current key. The projection is 
emitted exactly once
+      // per regular-map row (see below); emitting it in more than one runtime 
branch is unsafe
+      // because the projection's subexpression/writer state assigned in one 
branch would be read
+      // stale from another (e.g. the adaptive pass-through path would reuse 
the last probed key).
+      val probeRegularMap =
+        s"""
+           |int $unsafeRowKeyHash = ${unsafeRowKeyCode.value}.hashCode();
+           |if ($checkFallbackForBytesToBytesMap) {
+           |  // try to get the buffer from hash map
+           |  $unsafeRowBuffer =
+           |    $hashMapTerm.getAggregationBufferFromUnsafeRow($unsafeRowKeys, 
$unsafeRowKeyHash);
+           |}
+         """.stripMargin
+
+      val spillMap =
+        s"""
+           |if ($sorterTerm == null) {
+           |  $sorterTerm = $hashMapTerm.destructAndCreateExternalSorter();
+           |} else {
+           |  
$sorterTerm.merge($hashMapTerm.destructAndCreateExternalSorter());
+           |}
+           |$resetCounter
+           |// the hash map had be spilled, it should have enough memory now,
+           |// try to allocate buffer again.
+           |$unsafeRowBuffer = $hashMapTerm.getAggregationBufferFromUnsafeRow(
+           |  $unsafeRowKeys, $unsafeRowKeyHash);
+           |if ($unsafeRowBuffer == null) {
+           |  // failed to allocate the first page
+           |  throw QueryExecutionErrors.aggregateOutOfMemoryError();
+           |}
+         """.stripMargin
+
+      if (adaptivePartialAggConfig.isDefined) {
+        val cfg = adaptivePartialAggConfig.get
+        // Adaptive partial aggregation governs only this regular 
(second-level) map. Count the
+        // rows that enter it (a fast-map miss, or every row when the fast map 
is off) and use
+        // `regularMap.getNumKeys() / regularRows` as the pre-shuffle 
reduction ratio.
+        //   - Tier 2 (on-spill): when the map cannot allocate for a new key 
(it would otherwise
+        //     spill), bypass instead if the ratio is at least 
`spillReductionRatioThreshold`.
+        //   - Tier 1 (no-spill): from `sampleRows` regular rows on, bypass if 
the ratio is at
+        //     least `noSpillReductionRatioThreshold`. The sampling window 
doubles after each
+        //     sub-threshold check, so low-cardinality input is re-evaluated 
only rarely while a
+        //     late high-cardinality tail can still trigger the bypass.
+        // Both tiers fire only before any spill (`sorter == null`): once the 
map has spilled, the
+        // reduction-ratio estimate no longer covers the spilled rows, and 
pass-through must never
+        // coexist with sort-based aggregation. When the map is full after a 
spill, the map spills
+        // again as usual.
+        // The key projection runs once here so `unsafeRowKeyCode.value` is 
valid for both the
+        // probe below and the pass-through buffer built by the caller.
+        s"""
+           |// generate grouping key
+           |${unsafeRowKeyCode.code}
+           |if (!$adaptivePassThroughTerm) {
+           |  $regularMapRowCountTerm += 1;

Review Comment:
   Evaluate the spill ratio over the same row set in both execution paths. This 
increment counts the failed insertion that becomes the first pass-through row, 
while the interpreted path evaluates before counting it. Move the increment 
after the spill decision and add an exact threshold-boundary test with codegen 
enabled and disabled.



##########
sql/core/src/main/scala/org/apache/spark/sql/execution/aggregate/HashAggregateExec.scala:
##########
@@ -663,46 +855,122 @@ case class HashAggregateExec(
       case _ => ("true", "", "")
     }
 
-    val findOrInsertRegularHashMap: String =
-      s"""
-         |// generate grouping key
-         |${unsafeRowKeyCode.code}
-         |int $unsafeRowKeyHash = ${unsafeRowKeyCode.value}.hashCode();
-         |if ($checkFallbackForBytesToBytesMap) {
-         |  // try to get the buffer from hash map
-         |  $unsafeRowBuffer =
-         |    $hashMapTerm.getAggregationBufferFromUnsafeRow($unsafeRowKeys, 
$unsafeRowKeyHash);
-         |}
-         |// Can't allocate buffer from the hash map. Spill the map and 
fallback to sort-based
-         |// aggregation after processing all input rows.
-         |if ($unsafeRowBuffer == null) {
-         |  if ($sorterTerm == null) {
-         |    $sorterTerm = $hashMapTerm.destructAndCreateExternalSorter();
-         |  } else {
-         |    
$sorterTerm.merge($hashMapTerm.destructAndCreateExternalSorter());
-         |  }
-         |  $resetCounter
-         |  // the hash map had be spilled, it should have enough memory now,
-         |  // try to allocate buffer again.
-         |  $unsafeRowBuffer = $hashMapTerm.getAggregationBufferFromUnsafeRow(
-         |    $unsafeRowKeys, $unsafeRowKeyHash);
-         |  if ($unsafeRowBuffer == null) {
-         |    // failed to allocate the first page
-         |    throw QueryExecutionErrors.aggregateOutOfMemoryError();
-         |  }
-         |}
-       """.stripMargin
+    val findOrInsertRegularHashMap: String = {
+      // Assumes the grouping key projection (`unsafeRowKeyCode.code`) has 
already run for this row,
+      // so `unsafeRowKeyCode.value` holds the current key. The projection is 
emitted exactly once
+      // per regular-map row (see below); emitting it in more than one runtime 
branch is unsafe
+      // because the projection's subexpression/writer state assigned in one 
branch would be read
+      // stale from another (e.g. the adaptive pass-through path would reuse 
the last probed key).
+      val probeRegularMap =
+        s"""
+           |int $unsafeRowKeyHash = ${unsafeRowKeyCode.value}.hashCode();
+           |if ($checkFallbackForBytesToBytesMap) {
+           |  // try to get the buffer from hash map
+           |  $unsafeRowBuffer =
+           |    $hashMapTerm.getAggregationBufferFromUnsafeRow($unsafeRowKeys, 
$unsafeRowKeyHash);
+           |}
+         """.stripMargin
+
+      val spillMap =
+        s"""
+           |if ($sorterTerm == null) {
+           |  $sorterTerm = $hashMapTerm.destructAndCreateExternalSorter();
+           |} else {
+           |  
$sorterTerm.merge($hashMapTerm.destructAndCreateExternalSorter());
+           |}
+           |$resetCounter
+           |// the hash map had be spilled, it should have enough memory now,

Review Comment:
   ```suggestion
              |// the hash map had been spilled, so it should have enough 
memory now,
   ```



##########
sql/catalyst/src/main/scala/org/apache/spark/sql/internal/SQLConf.scala:
##########
@@ -4156,6 +4156,66 @@ object SQLConf {
     .booleanConf
     .createWithDefault(false)
 
+  val ADAPTIVE_PARTIAL_AGGREGATION_ENABLED =
+    
buildConf("spark.sql.execution.aggregate.adaptivePartialAggregation.enabled")
+      .doc("When true, hash aggregation adaptively bypasses the pre-shuffle 
partial aggregation " +
+        "at runtime when it observes that the partial aggregation is not 
reducing the number of " +
+        "rows enough to be worthwhile. Once bypassed, the remaining input rows 
are passed " +
+        "through as single-row partial aggregation buffers for the final 
aggregation to merge, " +
+        "which avoids the cost of maintaining and spilling a large aggregation 
map with little " +
+        "reduction benefit. This applies only to hash aggregation with 
grouping keys.")
+      .version("4.3.0")
+      .withBindingPolicy(ConfigBindingPolicy.SESSION)
+      .booleanConf
+      .createWithDefault(true)
+
+  val ADAPTIVE_PARTIAL_AGGREGATION_SAMPLE_ROWS =
+    
buildConf("spark.sql.execution.aggregate.adaptivePartialAggregation.sampleRows")
+      .doc("The number of input rows to sample before evaluating the reduction 
ratio for the " +
+        s"no-spill tier of adaptive partial aggregation (see " +
+        s"'${ADAPTIVE_PARTIAL_AGGREGATION_ENABLED.key}'). From this many rows 
on, if the ratio " +
+        "of distinct grouping keys to processed rows is at least " +
+        s"'spark.sql.execution.aggregate.adaptivePartialAggregation." +
+        "noSpillReductionRatioThreshold', partial aggregation is bypassed for 
the rest of the " +
+        "input. When the ratio is below the threshold, the next evaluation 
happens after twice " +
+        "as many rows, so low-cardinality input is re-checked only rarely.")
+      .version("4.3.0")
+      .withBindingPolicy(ConfigBindingPolicy.SESSION)
+      .intConf
+      .checkValue(_ > 0, "The sample row count must be positive.")
+      .createWithDefault(100000)
+
+  val ADAPTIVE_PARTIAL_AGGREGATION_NO_SPILL_REDUCTION_RATIO_THRESHOLD =

Review Comment:
   Could we start with one minimum-row setting and one reduction threshold, 
applying the same policy both periodically and immediately before spilling? The 
separate no-spill and spill thresholds plus exponential resampling add 
configuration and behavioral complexity without showing that these dimensions 
must be independently tunable. At the spill boundary, switch to pass-through 
when the common policy says aggregation is ineffective; otherwise spill 
normally. This also gives the code-generated and interpreted paths one 
invariant to implement and test.



##########
sql/core/src/test/scala/org/apache/spark/sql/execution/benchmark/AdaptivePartialAggregationBenchmark.scala:
##########
@@ -0,0 +1,137 @@
+/*
+ * 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.spark.sql.execution.benchmark
+
+import org.apache.spark.benchmark.Benchmark
+import org.apache.spark.sql.DataFrame
+import org.apache.spark.sql.internal.SQLConf
+
+/**
+ * Benchmark comparing runtime adaptive partial aggregation (see
+ * [[SQLConf.ADAPTIVE_PARTIAL_AGGREGATION_ENABLED]]) against the static 
pre-shuffle partial
+ * aggregation. When the partial aggregation is not reducing rows, the 
operator streams the
+ * remaining rows through as single-row partial buffers instead of maintaining 
(and possibly
+ * spilling) a large aggregation map.
+ *
+ * Each scenario runs the query across the full matrix of whole-stage codegen 
on/off and the
+ * feature disabled (`adaptive = F`, the pre-change baseline) vs enabled 
(`adaptive = T`), over a
+ * {high, low}-cardinality x {no-spill, on-spill} grid:
+ *   - high-cardinality, no spill: the no-spill tier bypasses, which should 
win.
+ *   - low-cardinality, no spill: nothing bypasses, which must not regress.
+ *   - high-cardinality, forced regular-map spill: the on-spill tier bypasses 
instead of spilling,
+ *     which should win.
+ *   - low-cardinality, forced regular-map spill: the ratio is too low for the 
on-spill tier to
+ *     bypass, so both runs spill identically (no regression).
+ *
+ * To run this benchmark:
+ * {{{
+ *   1. build/sbt "sql/Test/runMain
+ *        
org.apache.spark.sql.execution.benchmark.AdaptivePartialAggregationBenchmark"
+ *   2. generate result: SPARK_GENERATE_BENCHMARK_FILES=1 build/sbt 
"sql/Test/runMain
+ *        
org.apache.spark.sql.execution.benchmark.AdaptivePartialAggregationBenchmark"
+ *      Results will be written to 
"benchmarks/AdaptivePartialAggregationBenchmark-results.txt".
+ * }}}
+ */
+object AdaptivePartialAggregationBenchmark extends SqlBasedBenchmark {
+
+  override def runBenchmarkSuite(mainArgs: Array[String]): Unit = {
+    // The upstream `CombineAdjacentAggregation` and `ReplaceHashWithSortAgg` 
rules would collapse
+    // or convert these single-partition hash aggregates, so both are disabled 
to keep the
+    // Partial+Final `HashAggregateExec` structure the adaptive feature 
governs.
+    val fixedPlanConfs = Seq(
+      SQLConf.COMBINE_ADJACENT_AGGREGATION_ENABLED.key -> "false",
+      SQLConf.REPLACE_HASH_WITH_SORT_AGG_ENABLED.key -> "false")
+
+    // Adds the (whole-stage codegen, adaptive switch) matrix for `query`. 
`extraConf` is applied
+    // to all four cases so the only differences are the two axes.
+    def addCodegenAdaptiveCases(
+        benchmark: Benchmark,
+        query: () => DataFrame,
+        extraConf: Seq[(String, String)] = Nil): Unit = {
+      for {
+        wholeStage <- Seq(true, false)
+        adaptive <- Seq(false, true)
+      } {
+        val adaptiveLabel = if (adaptive) "T" else "F"
+        val label = s"codegen = $wholeStage, adaptive = $adaptiveLabel"
+        benchmark.addCase(label) { _ =>
+          withSQLConf(
+            (Seq(
+              SQLConf.WHOLESTAGE_CODEGEN_ENABLED.key -> wholeStage.toString,
+              SQLConf.ADAPTIVE_PARTIAL_AGGREGATION_ENABLED.key -> 
adaptive.toString) ++
+              fixedPlanConfs ++ extraConf): _*) {
+            query().noop()
+          }
+        }
+      }
+    }
+
+    // Fully distinct keys make partial aggregation useless, so the no-spill 
(Tier 1) sampling tier
+    // bypasses: the feature should be faster than the baseline that maintains 
a map entry per row.
+    runBenchmark("high-cardinality input, no-spill pass-through (Tier 1)") {
+      val N = 8L << 20
+      val benchmark = new Benchmark("adaptive partial agg, high card, no 
spill", N,
+        output = output)
+      addCodegenAdaptiveCases(benchmark, () => distinctKeyedDf(N))
+      benchmark.run()
+    }
+
+    // 1000 distinct keys over a large input: partial aggregation reduces a 
lot, the no-spill tier
+    // never fires, and the two runs must match (no regression).
+    runBenchmark("low-cardinality input, no-spill pass-through (Tier 1)") {
+      val N = 16L << 20
+      val benchmark = new Benchmark("adaptive partial agg, low card, no 
spill", N,
+        output = output)
+      addCodegenAdaptiveCases(benchmark, () =>
+        spark.range(N).selectExpr("id % 1000 as k", "id as 
v").groupBy("k").agg("v" -> "sum"))
+      benchmark.run()
+    }
+
+    // Force the regular map to spill quickly and disable the no-spill tier 
(huge sample). With
+    // fully distinct keys the reduction ratio is 1.0, so at the spill 
boundary the on-spill
+    // (Tier 2) tier bypasses instead of spilling; the baseline spills 
repeatedly and falls back

Review Comment:
   ```suggestion
       // tier (Tier 2) bypasses instead of spilling; the baseline spills 
repeatedly and falls back
   ```



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