cloud-fan commented on code in PR #57742:
URL: https://github.com/apache/spark/pull/57742#discussion_r3713155707
##########
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:
Thanks, the shared default plus the concrete spill-cost rationale addresses
my concern about introducing two policies by default.
##########
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:
Confirmed: both paths now evaluate the spill ratio over the pre-failure
rows, and the boundary test covers codegen on and off.
##########
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:
Confirmed, the immutable adaptive settings are now extracted before the hot
loop.
##########
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:
Confirmed, thanks.
--
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]