Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
16 changes: 16 additions & 0 deletions datafusion/common/src/config.rs
Original file line number Diff line number Diff line change
Expand Up @@ -1136,6 +1136,22 @@ config_namespace! {
/// aggregation ratio check and trying to switch to skipping aggregation mode
pub skip_partial_aggregation_probe_rows_threshold: usize, default = 100_000

/// (experimental) Allocated-byte threshold for flushing unordered partial
/// hash aggregation tables. Checked after each input batch; the table can
/// exceed this size by one batch. Flushed groups count toward the existing
/// skip-partial probe. Repeated keys do not disable flushing. Aggregations
/// with a soft group limit or nested aggregate state are excluded.
/// Set to 0 to disable this threshold. If both flush thresholds are
/// enabled, reaching either triggers a flush. Final aggregation is unchanged.
pub partial_aggregation_flush_bytes: usize, default = 0

/// (experimental) Number of distinct group rows above which an unordered
/// partial hash aggregation table is flushed. This counts groups held in
/// the table, not input rows. Uses the same eligibility and skip-partial
/// accounting as partial_aggregation_flush_bytes. Set to 0 to disable this
/// threshold. If both thresholds are enabled, reaching either triggers a flush.
pub partial_aggregation_flush_rows: usize, default = 0

/// Should DataFusion use row number estimates at the input to decide
/// whether increasing parallelism is beneficial or not. By default,
/// only exact row numbers (not estimates) are used for this decision.
Expand Down
119 changes: 103 additions & 16 deletions datafusion/physical-plan/src/aggregates/hash_stream.rs
Original file line number Diff line number Diff line change
Expand Up @@ -159,6 +159,9 @@ pub(crate) struct PartialHashAggregateStream {
/// Number of times accumulated states were emitted due to memory pressure.
early_emit_count: metrics::Count,

/// Optional thresholds for emitting partial states before memory pressure.
table_flush: Option<PartialTableFlush>,

/// Tracks whether partial aggregation should switch to direct state conversion.
skip_aggregation_probe: Option<SkipAggregationProbe>,

Expand Down Expand Up @@ -207,9 +210,22 @@ enum HandleInputResult {
ReachedLimit,
#[expect(clippy::upper_case_acronyms)]
OOM,
FlushThresholdReached,
SwitchToSkipAggregation,
}

/// Emits partial states at size thresholds while preserving skip-probe counts.
struct PartialTableFlush {
/// Allocated bytes above which the current table is emitted.
byte_threshold: usize,
/// Distinct group rows above which the current table is emitted; zero disables.
group_threshold: usize,
/// Previously emitted groups, included in the skip-partial reduction ratio.
flushed_groups: usize,
/// Number of emissions triggered by either size threshold.
flush_count: metrics::Count,
}

impl PartialHashAggregateStream {
pub fn new(
agg: &AggregateExec,
Expand All @@ -232,6 +248,26 @@ impl PartialHashAggregateStream {
let early_emit_count =
MetricBuilder::new(&agg.metrics).counter("early_emit_count", partition);

let execution_options = &context.session_config().options().execution;
let byte_threshold = execution_options.partial_aggregation_flush_bytes;
let group_threshold = execution_options.partial_aggregation_flush_rows;
let group_values_soft_limit = agg.limit_options().map(|config| config.limit());
let has_nested_state = schema
.fields()
.iter()
.skip(agg.group_by().num_group_exprs())
.any(|field| field.data_type().is_nested());
let table_flush = ((byte_threshold > 0 || group_threshold > 0)
&& group_values_soft_limit.is_none()
&& !has_nested_state)
.then(|| PartialTableFlush {
byte_threshold,
group_threshold,
flushed_groups: 0,
flush_count: MetricBuilder::new(&agg.metrics)
.counter("table_flush_count", partition),
});

let hash_table = AggregateHashTable::<PartialMarker>::new(
agg,
partition,
Expand Down Expand Up @@ -279,8 +315,9 @@ impl PartialHashAggregateStream {
reservation,
reduction_factor,
early_emit_count,
table_flush,
skip_aggregation_probe,
group_values_soft_limit: agg.limit_options().map(|config| config.limit()),
group_values_soft_limit,
hash_table: Some(hash_table),
})
}
Expand Down Expand Up @@ -315,14 +352,23 @@ impl PartialHashAggregateStream {
| HandleInputResult::SwitchToSkipAggregation => {
break;
}
HandleInputResult::OOM => {
let materialized_group_states = hash_table.take_state_batch()?.ok_or_else(|| {
internal_datafusion_err!(
"Partial hash aggregate ran out of memory with no aggregated groups"
)
})?;

self.early_emit_count.add(1);
HandleInputResult::OOM | HandleInputResult::FlushThresholdReached => {
let materialized_group_states =
hash_table.take_state_batch()?.ok_or_else(|| {
internal_datafusion_err!(
"Partial hash aggregate tried to flush an empty table"
)
})?;

if let Some(table_flush) = self.table_flush.as_mut()
&& last_state == HandleInputResult::FlushThresholdReached
{
table_flush.flushed_groups +=
materialized_group_states.num_rows();
table_flush.flush_count.add(1);
} else {
self.early_emit_count.add(1);
}
timer.done();
self.emit_on_memory_pressure(
materialized_group_states,
Expand Down Expand Up @@ -410,7 +456,14 @@ impl PartialHashAggregateStream {
// ----------------------------------------------
// Step 3: Skip partial aggregation optimization
// ----------------------------------------------
self.update_skip_aggregation_probe(input_rows, hash_table.building_group_count());
let flushed_groups = self
.table_flush
.as_ref()
.map_or(0, |table_flush| table_flush.flushed_groups);
self.update_skip_aggregation_probe(
input_rows,
flushed_groups + hash_table.building_group_count(),
);

// True branch: a decision has been made to skip partial aggregation.
if self.should_skip_aggregation() {
Expand All @@ -420,12 +473,26 @@ impl PartialHashAggregateStream {
// -------------------------------------------------
// Step 4: Larger-than-memory execution (early emit)
// -------------------------------------------------
let resize_result = self.reservation.try_resize(hash_table.memory_size());
let allocated_bytes = hash_table.memory_size();
let resize_result = self.reservation.try_resize(allocated_bytes);
match resize_result {
Ok(()) => Ok(HandleInputResult::ProcessNext),
Err(DataFusionError::ResourcesExhausted(_)) => Ok(HandleInputResult::OOM),
Err(e) => Err(e),
Ok(()) => {}
Err(DataFusionError::ResourcesExhausted(_)) => {
return Ok(HandleInputResult::OOM);
}
Err(e) => return Err(e),
}

// Step 5: Emit partial states once either configured threshold is reached.
if self.table_flush.as_ref().is_some_and(|table_flush| {
(table_flush.byte_threshold > 0
&& allocated_bytes >= table_flush.byte_threshold)
|| (table_flush.group_threshold > 0
&& hash_table.building_group_count() >= table_flush.group_threshold)
}) {
return Ok(HandleInputResult::FlushThresholdReached);
}
Ok(HandleInputResult::ProcessNext)
}

/// Emit a materialized partial-state batch in `batch_size` slices.
Expand Down Expand Up @@ -1021,9 +1088,16 @@ mod tests {
Ok(())
}

#[rstest::rstest]
#[case(0, 0)]
#[case(1, 0)]
#[case(0, 1)]
#[case(1, 1)]
#[tokio::test]
async fn test_partial_hash_stream_skip_aggregation_probe_not_locked_until_skip()
-> Result<()> {
async fn test_partial_hash_stream_skip_aggregation_probe_not_locked_until_skip(
#[case] flush_bytes: usize,
#[case] flush_rows: usize,
) -> Result<()> {
// Test that the probe is not locked until we actually decide to skip.
// This allows us to continue evaluating the skip condition across multiple batches.
//
Expand Down Expand Up @@ -1108,6 +1182,14 @@ mod tests {

// Configure skip aggregation settings
let mut session_config = task_ctx.session_config().clone();
session_config
.options_mut()
.execution
.partial_aggregation_flush_bytes = flush_bytes;
session_config
.options_mut()
.execution
.partial_aggregation_flush_rows = flush_rows;
session_config = session_config.set(
"datafusion.execution.skip_partial_aggregation_probe_rows_threshold",
&datafusion_common::ScalarValue::UInt64(Some(probe_rows_threshold)),
Expand Down Expand Up @@ -1166,6 +1248,11 @@ mod tests {
"Expected batch 3's rows ({batch3_rows}) to be skipped",
);

let flushes = metrics
.sum_by_name("table_flush_count")
.map_or(0, |value| value.as_usize());
assert_eq!(flushes, usize::from(flush_bytes > 0 || flush_rows > 0));

Ok(())
}

Expand Down
Loading
Loading