diff --git a/datafusion/physical-plan/src/aggregates/aggregate_hash_table/common_ordered.rs b/datafusion/physical-plan/src/aggregates/aggregate_hash_table/common_ordered.rs index 97ef898b51f4d..a5a81bc511e24 100644 --- a/datafusion/physical-plan/src/aggregates/aggregate_hash_table/common_ordered.rs +++ b/datafusion/physical-plan/src/aggregates/aggregate_hash_table/common_ordered.rs @@ -352,24 +352,6 @@ impl OrderedAggregateTable { Ok(Some(batch)) } - /// Returns the [`EmitTo`], clamped to the specified batch size - /// - /// Returns `(emit_to, should_remove_groups)`, where `emit_to` is the number - /// of groups to emit from `GroupValues` / accumulators, and - /// `should_remove_groups` indicates whether `GroupOrdering` must also shift - /// its tracked indexes. - pub(super) fn clamp_emit_to( - &self, - group_count: usize, - emit_to: EmitTo, - ) -> (EmitTo, bool) { - match emit_to { - EmitTo::First(n) => (EmitTo::First(n.min(self.batch_size)), true), - EmitTo::All if group_count <= self.batch_size => (EmitTo::All, false), - EmitTo::All => (EmitTo::First(self.batch_size), false), - } - } - /// Aggregates one evaluated input batch after selecting the mode-specific /// accumulator operation. /// @@ -445,20 +427,34 @@ impl OrderedAggregateTable { let Some(emit_to) = self.buffer.group_ordering.emit_to() else { return Ok(None); }; - let (emit_to, should_remove_groups) = - self.clamp_emit_to(self.buffer.group_values.len(), emit_to); + let emit_to = match emit_to { + EmitTo::First(n) => EmitTo::First(n.min(self.batch_size)), + EmitTo::All if self.num_groups() > self.batch_size => { + EmitTo::First(self.batch_size) + } + EmitTo::All => EmitTo::All, + }; + self.materialize_groups(emit_to, materialize_accumulator_fn, accumulator_phase) + .map(Some) + } + /// Removes the selected groups once and materializes their output columns. + /// The caller chooses the completed prefix and any output-size limit. + pub(super) fn materialize_groups( + &mut self, + emit_to: EmitTo, + materialize_accumulator_fn: MaterializeAccumulatorFn, + accumulator_phase: AccumulatorPhase, + ) -> Result { let accumulator_metrics = Arc::clone(&self.aggregate_accumulator_metrics); let output = self.group_by_metrics.time_emitting(|| { let mut output = self.buffer.group_values.emit(emit_to)?; - if should_remove_groups { - match emit_to { - EmitTo::First(n) => self.buffer.group_ordering.remove_groups(n), - // `EmitTo::All` is only used after `input_done`, when all - // buffered groups are known complete and the ordering state is - // no longer needed. - EmitTo::All => {} - } + // EOF can also emit a prefix when a caller limits its batch size, + // but the completed ordering state no longer tracks group indexes. + if let EmitTo::First(n) = emit_to + && matches!(self.buffer.group_ordering.emit_to(), Some(EmitTo::First(_))) + { + self.buffer.group_ordering.remove_groups(n); } for (idx, acc) in self.buffer.accumulators.iter_mut().enumerate() { @@ -474,6 +470,6 @@ impl OrderedAggregateTable { let batch = RecordBatch::try_new(Arc::clone(&self.output_schema), output)?; debug_assert!(batch.num_rows() > 0); - Ok(Some(batch)) + Ok(batch) } } diff --git a/datafusion/physical-plan/src/aggregates/aggregate_hash_table/ordered_partial_table.rs b/datafusion/physical-plan/src/aggregates/aggregate_hash_table/ordered_partial_table.rs index 6ed93e59f3296..b53012d8bd1e4 100644 --- a/datafusion/physical-plan/src/aggregates/aggregate_hash_table/ordered_partial_table.rs +++ b/datafusion/physical-plan/src/aggregates/aggregate_hash_table/ordered_partial_table.rs @@ -89,26 +89,22 @@ impl OrderedAggregateTable { ) } - /// Emits the next batch of partial state rows for groups proven complete by - /// the input ordering. - /// - /// For example, when the query is `GROUP BY a` and the input is ordered by - /// `a`, seeing a latest input row with `a = 3` means all groups with `a < 3` - /// are complete and safe to emit. - /// - /// Key steps: - /// 1. Ask `group_ordering` to decide how many groups can be emitted eagerly. - /// 2. Remove the emitted groups from `group_ordering`, `GroupValues`, and - /// all `GroupsAccumulator`s. - /// - /// This may output small batches. Avoiding tiny batches is left to future - /// ordered-aggregation optimizations. - pub(in crate::aggregates) fn next_output_batch( + /// Materializes all groups proven complete by the input ordering, leaving + /// the active ordered-key range in the table. + pub(in crate::aggregates) fn take_completed_state_batch( &mut self, ) -> Result> { - self.next_output_batch_inner( + if self.is_empty() { + return Ok(None); + } + let Some(emit_to) = self.group_ordering().emit_to() else { + return Ok(None); + }; + self.materialize_groups( + emit_to, HashAggregateAccumulator::state, AccumulatorPhase::State, ) + .map(Some) } } diff --git a/datafusion/physical-plan/src/aggregates/ordered_partial_stream.rs b/datafusion/physical-plan/src/aggregates/ordered_partial_stream.rs index 8c2315588ad7a..dd845353c77eb 100644 --- a/datafusion/physical-plan/src/aggregates/ordered_partial_stream.rs +++ b/datafusion/physical-plan/src/aggregates/ordered_partial_stream.rs @@ -24,14 +24,14 @@ use arrow::record_batch::RecordBatch; use datafusion_common::{DataFusionError, Result}; use datafusion_execution::memory_pool::{MemoryConsumer, MemoryReservation}; use datafusion_execution::{TaskContext, TryEmitter, async_try_stream}; -use futures::stream::{Stream, StreamExt}; +use futures::stream::StreamExt; use super::AggregateExec; use super::aggregate_hash_table::{OrderedAggregateTable, PartialMarker}; use crate::aggregates::AggregateMode; use crate::aggregates::order::GroupOrdering; use crate::metrics::{BaselineMetrics, MetricBuilder, SpillMetrics}; -use crate::stream::{EmptyRecordBatchStream, ObservedStream, RecordBatchStreamAdapter}; +use crate::stream::{ObservedStream, RecordBatchStreamAdapter}; use crate::{InputOrderMode, SendableRecordBatchStream, metrics}; /// Partial aggregate stream for `InputOrderMode::Sorted` and @@ -66,7 +66,9 @@ use crate::{InputOrderMode, SendableRecordBatchStream, metrics}; /// After each input batch, check whether any groups can be emitted eagerly to /// improve memory efficiency. For example, if the last group key seen is /// `k = 100`, it is safe to emit all groups with keys less than 100 because the -/// input is ordered. +/// input is ordered. Materialize that entire completed prefix once, then emit +/// slices of it before reading more input. This avoids repeatedly removing small +/// batches of groups and shifting the remaining hash table and accumulator state. /// /// # Memory Pressure and Spilling /// @@ -103,12 +105,36 @@ use crate::{InputOrderMode, SendableRecordBatchStream, metrics}; /// remaining states, performs a sort-preserving merge of all runs, and feeds the /// merged input into a fully ordered final aggregate stream. pub(crate) struct OrderedPartialAggregateStream { - schema: SchemaRef, - input: SendableRecordBatchStream, reservation: MemoryReservation, + context: OrderedPartialAggregateContext, + stage: ExecutionStage, +} + +/// Execution stages described in [`OrderedPartialAggregateStream::into_stream`]. +enum ExecutionStage { + Aggregating(Aggregating), + Outputting(Outputting), +} + +struct Aggregating { + input: SendableRecordBatchStream, + table: OrderedAggregateTable, +} + +struct Outputting { + /// Materialized aggregate states. Each iteration emits the first `batch_size` + /// rows and replaces this batch with the remaining slice. + batch: RecordBatch, + /// Aggregation stage to resume after output; `None` after EOF. + resume: Option, +} + +/// Immutable execution context shared by aggregation and output emission. +struct OrderedPartialAggregateContext { + schema: SchemaRef, + batch_size: usize, baseline_metrics: BaselineMetrics, reduction_factor: metrics::RatioMetrics, - table: Option>, } impl OrderedPartialAggregateStream { @@ -146,199 +172,245 @@ impl OrderedPartialAggregateStream { .register(context.memory_pool()); Ok(Self { - schema, - input, reservation, - baseline_metrics, - reduction_factor, - table: Some(table), + context: OrderedPartialAggregateContext { + schema, + batch_size, + baseline_metrics, + reduction_factor, + }, + stage: ExecutionStage::Aggregating(Aggregating { input, table }), }) } - pub(crate) fn into_stream(self) -> SendableRecordBatchStream { - let schema_clone = Arc::clone(&self.schema); - - let cloned_metrics = self.baseline_metrics.clone(); - let stream = Box::pin(RecordBatchStreamAdapter::new( - schema_clone, - self.create_stream(), - )); - - Box::pin(ObservedStream::new(stream, cloned_metrics, None)) - } - - /// Entry point for the ordered partial aggregate state machine. - /// - /// See comments in [`OrderedPartialAggregateStream`] for high-level ideas. + /// Entry point for the ordered partial aggregate execution stages. /// - /// State transitions are implemented using the generator pattern; see the comments in [`async_try_stream`]. + /// See [`OrderedPartialAggregateStream`] for high-level ideas. /// - /// Conceptual state-transition graph: + /// # Stage transition graph: /// /// ```text - /// (start) - /// -> ReadingInput - /// The stream starts by polling ordered input and aggregating batches - /// into the ordered partial aggregate table. + /// +----[2]----+ +----[5]----+ + /// | | | | + /// v | v | + /// +-------------------+ +-------------------+ + /// | | | | + /// (start)-[1]->| Aggregating |-----[3]---->| Outputting | + /// | |<----[6]-----| | + /// +-------------------+ +-------------------+ + /// | [4] | [7] + /// | | + /// +----------------+----------------+ + /// | + /// v + /// +---------+ + /// | Done |--[8]--> (end) + /// +---------+ + /// ``` + /// + /// ## Stages + /// + /// - [`Aggregating`]: Aggregates raw input and materializes one batch of partial + /// states. + /// - [`Outputting`]: Emits slices of one materialized batch. If the materialized + /// buffers cannot be reserved while slicing, hand off the whole batch, then + /// resume aggregation or finish as described below. /// - /// ReadingInput - /// -> ReadingInput - /// Aggregate one input batch. If the ordering proves some groups are - /// complete, yield one partial-state batch immediately, then continue - /// reading input. Otherwise continue directly with the next input batch. - /// -> DrainingFinal - /// Input was exhausted. Mark the table input as done so every remaining - /// group is safe to emit. + /// ### Incremental output /// - /// DrainingFinal - /// -> DrainingFinal - /// One remaining partial-state batch was yielded; repeat to continue - /// draining the table. - /// -> Done - /// All remaining groups were emitted. + /// Consider this query with input ordered only by `k1`: /// - /// Done - /// -> (end) + /// ```sql + /// SELECT k1, k2, AVG(v) + /// FROM table_with_order_k1 + /// GROUP BY k1, k2 /// ``` - fn create_stream(mut self) -> impl Stream> { - async_try_stream(|mut emitter| async move { - let mut table = self - .table - .take() - .expect("OrderedPartialAggregateStream state should not be None"); - - self.handle_reading_input(&mut table, &mut emitter).await?; - - // Input has exhausted, move to the final draining stage. - self.close_input(); - table.input_done(); - - self.handle_draining_final(table, &mut emitter).await?; - + /// + /// Suppose one `k1` value spans 1M rows with distinct, unordered `k2` + /// values. Ordering only proves these 1M `(k1, k2)` groups complete when + /// `k1` changes, so a single early emission can produce far more than + /// `batch_size` rows. + /// + /// Emitting those groups in small batches through [EmitTo::First] would + /// repeatedly remove a prefix from [`GroupValues`]. Because the group + /// values are stored contiguously, each removal copies the remaining values + /// and updates their group indexes. + /// + /// To avoid repeating that work, this stream: + /// + /// 1. Materializes all completed groups into one large batch. + /// 2. Emits `batch_size` slices that share the batch's buffers. + /// + /// Blocked aggregate state management may simplify this approach: + /// + /// + /// [`GroupValues`]: crate::aggregates::group_values::GroupValues + /// [EmitTo::First]: datafusion_expr::EmitTo::First + /// + /// + /// ## Transition Edges + /// + /// 1. Start. + /// 2. Aggregate one input batch. If memory fits and no groups are complete, + /// continue reading input. + /// 3. Prepare output: + /// - Ordering proves a prefix complete: materialize the entire prefix once, + /// retaining the input and active groups to resume aggregation. + /// - On memory pressure with partial ordering, materialize all current + /// states instead, including incomplete groups, and reset the table. + /// - At EOF, materialize all remaining states and prepare to output. + /// 4. Input was exhausted with no remaining groups, directly end. + /// 5. Yield one slice without materializing the table again. Keep the shared + /// buffers reserved until handing off the last slice. + /// 6. The batch was fully emitted and retained aggregation can resume. + /// 7. The output batch was fully emitted. + /// 8. End. + pub(crate) fn into_stream(self) -> SendableRecordBatchStream { + let Self { + reservation, + context, + stage, + } = self; + let schema = Arc::clone(&context.schema); + let metrics = context.baseline_metrics.clone(); + let stream = async_try_stream(|mut emitter| async move { + let mut stage = Some(stage); + while let Some(current_stage) = stage { + stage = match current_stage { + ExecutionStage::Aggregating(aggregating) => { + aggregating.handle_stage(&context, &reservation).await? + } + ExecutionStage::Outputting(outputting) => { + outputting + .handle_stage(&context, &reservation, &mut emitter) + .await? + } + }; + } Ok(()) - }) - } - - fn close_input(&mut self) { - let input_schema = self.input.schema(); - self.input = Box::pin(EmptyRecordBatchStream::new(input_schema)); + }); + let stream = Box::pin(RecordBatchStreamAdapter::new(schema, stream)); + Box::pin(ObservedStream::new(stream, metrics, None)) } +} - /// Consumes one ordered input batch, then immediately emits completed groups - /// if the ordering proves any group is ready. +impl Aggregating { + /// Aggregates raw input and materializes one batch of partial states. /// - /// See comments at [`Self::create_stream`] for details. - async fn handle_reading_input( - &mut self, - table: &mut OrderedAggregateTable, - emitter: &mut TryEmitter, - ) -> Result<()> { - let elapsed_compute = self.baseline_metrics.elapsed_compute().clone(); + /// See [`OrderedPartialAggregateStream::into_stream`] for stage transitions. + async fn handle_stage( + mut self, + context: &OrderedPartialAggregateContext, + reservation: &MemoryReservation, + ) -> Result> { + let elapsed_compute = context.baseline_metrics.elapsed_compute(); while let Some(batch) = self.input.next().await.transpose()? { - let input_rows = batch.num_rows(); - self.reduction_factor.add_total(input_rows); - + context.reduction_factor.add_total(batch.num_rows()); let timer = elapsed_compute.timer(); - - table.aggregate_batch(&batch)?; - - // Check memory reservation. See function comments for details. - if let Some(batch) = self.resize_or_take_state_batch(table)? { - self.reduction_factor.add_part(batch.num_rows()); - drop(timer); - emitter.emit(batch).await; - continue; - } - - let Some(batch) = table.next_output_batch()? else { - // Can't do early emit, continue aggregating. + self.table.aggregate_batch(&batch)?; + + let output = match reservation.try_resize(self.table.memory_size()) { + Ok(()) => self.table.take_completed_state_batch()?, + Err(oom @ DataFusionError::ResourcesExhausted(_)) => { + // Partial ordering may have an unbounded active key range. + // The final stage can merge incomplete states emitted here. + if matches!(self.table.group_ordering(), GroupOrdering::Full(_)) { + return Err(oom); + } + let Some(batch) = self.table.take_state_batch()? else { + return Err(oom); + }; + Some(batch) + } + Err(e) => return Err(e), + }; + let Some(batch) = output else { continue; }; - self.reduction_factor.add_part(batch.num_rows()); - self.reservation.try_resize(table.memory_size())?; - - drop(timer); - emitter.emit(batch).await; - } - - Ok(()) - } - - /// Update the memory reservation, and: - /// - If memory reservation succeed, returns `Ok(None)` - /// - If memory reservation failed, - /// - If input is partially ordered, materialize all the output, and - /// directly send them to the final aggregation stage. - /// Returns `Ok(Some(batch))` - /// - If input is fully ordered, directly return error. It's not - /// expected to use more than constant memory. - /// Returns `Err(..)` - /// - /// # Implementation Note - /// Incrementally output it after the blocked state management is ready, keep - /// it simple for now. - /// - /// Issue: - fn resize_or_take_state_batch( - &mut self, - table: &mut OrderedAggregateTable, - ) -> Result> { - let oom = match self.reservation.try_resize(table.memory_size()) { - Ok(()) => return Ok(None), - Err(e @ DataFusionError::ResourcesExhausted(_)) => e, - Err(e) => return Err(e), - }; + timer.done(); - if matches!(table.group_ordering(), GroupOrdering::Full(_)) { - return Err(oom); + // OOM, do early emit next, and go back to the current state to continue + // aggregating + return Ok(Some(ExecutionStage::Outputting(Outputting { + batch, + resume: Some(self), + }))); } - let Some(batch) = table.take_state_batch()? else { - return Err(oom); + // Release upstream resources before draining the remaining states. + drop(self.input); + self.table.input_done(); + let timer = elapsed_compute.timer(); + let output = self.table.take_completed_state_batch()?; + drop(self.table); + timer.done(); + + let Some(batch) = output else { + reservation.try_resize(0)?; + return Ok(None); }; - self.reservation.try_resize(table.memory_size())?; - Ok(Some(batch)) + Ok(Some(ExecutionStage::Outputting(Outputting { + batch, + resume: None, + }))) } +} - /// Emits one batch after input is exhausted. - /// - /// `table.input_done()` has already made every remaining group safe to emit, - /// so this state keeps draining until the table is empty. - /// - /// See comments at [`Self::create_stream`] for details. +impl Outputting { + /// Emits slices of one materialized batch without touching the hash table. /// - async fn handle_draining_final( - &mut self, - mut table: OrderedAggregateTable, + /// See [`OrderedPartialAggregateStream::into_stream`] for stage transitions + /// and output memory accounting. + async fn handle_stage( + self, + context: &OrderedPartialAggregateContext, + reservation: &MemoryReservation, emitter: &mut TryEmitter, - ) -> Result<()> { - let elapsed_compute = self.baseline_metrics.elapsed_compute().clone(); + ) -> Result> { + let Self { mut batch, resume } = self; + let elapsed_compute = context.baseline_metrics.elapsed_compute(); let mut timer = elapsed_compute.timer(); - - while let Some(batch) = table.next_output_batch()? { - self.reduction_factor.add_part(batch.num_rows()); - - if table.is_empty() { - // Clear memory before emitting last batch so we don't have to wait for next poll to clear - drop(table); - let _ = self.reservation.try_resize(0); - drop(timer); - + let (table_memory, next_stage) = match resume { + Some(aggregating) => ( + aggregating.table.memory_size(), + Some(ExecutionStage::Aggregating(aggregating)), + ), + None => (0, None), + }; + let batch_memory = batch.get_array_memory_size(); + match reservation.try_resize(table_memory + batch_memory) { + Ok(()) => {} + Err(DataFusionError::ResourcesExhausted(_)) => { + // If we cannot hold the batch while slicing, hand it off whole. + // Only the retained table needs to remain reserved. + reservation.try_resize(table_memory)?; + context.reduction_factor.add_part(batch.num_rows()); + timer.done(); emitter.emit(batch).await; - - return Ok(()); + return Ok(next_stage); } + Err(e) => return Err(e), + } - self.reservation.try_resize(table.memory_size())?; - + while batch.num_rows() > context.batch_size { + // 1. Emit first `batch_size` rows from `batch` + // 2. Update `batch`` with the remaining tail + let output = batch.slice(0, context.batch_size); + batch = + batch.slice(context.batch_size, batch.num_rows() - context.batch_size); + context.reduction_factor.add_part(output.num_rows()); timer.done(); - emitter.emit(batch).await; + emitter.emit(output).await; timer = elapsed_compute.timer(); } - // was empty - Ok(()) + // The final slice transfers ownership of the buffers to the consumer. + reservation.try_shrink(batch_memory)?; + context.reduction_factor.add_part(batch.num_rows()); + timer.done(); + emitter.emit(batch).await; + Ok(next_stage) } }