This is an automated email from the ASF dual-hosted git repository.
github-merge-queue[bot] pushed a commit to branch main
in repository https://gitbox.apache.org/repos/asf/datafusion.git
The following commit(s) were added to refs/heads/main by this push:
new 1e09a2a3d7 perf: Avoid copying when materializing output in
OrderedPartialAggregateStream (#25312)
1e09a2a3d7 is described below
commit 1e09a2a3d76c03f469e80bed712a6bada45e77c8
Author: Yongting You <[email protected]>
AuthorDate: Mon Sep 21 00:06:54 2026 +0000
perf: Avoid copying when materializing output in
OrderedPartialAggregateStream (#25312)
## Which issue does this PR close?
<!--
We generally require a GitHub issue to be filed for all bug fixes and
enhancements and this helps us generate change logs for our releases.
You can link an issue to this PR using the GitHub syntax. For example
`Closes #123` indicates that this PR will close issue #123.
-->
part of https://github.com/apache/datafusion/issues/25157
## Rationale for this change
<!--
Why are you proposing this change? If this is already explained clearly
in the issue then this section is not needed.
Explaining clearly why changes are proposed helps reviewers understand
your changes and offer better suggestions for fixes.
Please explain the problem you are trying to solve in terms of the
user-visible
behavior, rather than the implementation.
For example, "The code in `foo.rs` doesn't handle nulls" is a symptom of
the
implementation. "COUNT(DISTINCT) returns wrong results when the column
contains
nulls" is the user-visible problem.
-->
### Cause
See issue for the target query.
The query plan looks like
<details>
<summary>Query plan, Click to expand</summary>
```
> explain SELECT count(*) FROM (
SELECT DISTINCT d_year, brand, class, cat, manu, cnt, amt FROM src
);
+---------------+-------------------------------+
| plan_type | plan |
+---------------+-------------------------------+
| physical_plan | ┌───────────────────────────┐ |
| | │ ProjectionExec │ |
| | │ -------------------- │ |
| | │ count(*): │ |
| | │ count(Int64(1)) │ |
| | └─────────────┬─────────────┘ |
| | ┌─────────────┴─────────────┐ |
| | │ AggregateExec │ |
| | │ -------------------- │ |
| | │ aggr: count(1) │ |
| | │ mode: Final │ |
| | └─────────────┬─────────────┘ |
| | ┌─────────────┴─────────────┐ |
| | │ CoalescePartitionsExec │ |
| | └─────────────┬─────────────┘ |
| | ┌─────────────┴─────────────┐ |
| | │ AggregateExec │ |
| | │ -------------------- │ |
| | │ aggr: count(1) │ |
| | │ mode: Partial │ |
| | └─────────────┬─────────────┘ |
| | ┌─────────────┴─────────────┐ |
| | │ ProjectionExec │ |
| | └─────────────┬─────────────┘ |
| | ┌─────────────┴─────────────┐ |
| | │ AggregateExec │ |
| | │ -------------------- │ |
| | │ group_by: │ |
| | │ d_year, brand, class, cat,│ |
| | │ manu, cnt, amt │ |
| | │ │ |
| | │ mode: │ |
| | │ FinalPartitioned │ |
| | └─────────────┬─────────────┘ |
| | ┌─────────────┴─────────────┐ |
| | │ RepartitionExec │ |
| | │ -------------------- │ |
| | │ partition_count(in->out): │ |
| | │ 14 -> 14 │ |
| | │ │ |
| | │ partitioning_scheme: │ |
| | │ Hash([d_year@0, brand@1, │ |
| | │ class@2, cat@3, manu@4 │ |
| | │ , cnt@5, amt@6], 14) │ |
| | │ │ |
| | │ preserve_order: true │ |
| | └─────────────┬─────────────┘ |
| | ┌─────────────┴─────────────┐ |
| | │ AggregateExec │ |
| | │ -------------------- │ |
| | │ group_by: │ |
| | │ d_year, brand, class, cat,│ |
| | │ manu, cnt, amt │ |
| | │ │ |
| | │ mode: Partial │ |
| | └─────────────┬─────────────┘ |
| | ┌─────────────┴─────────────┐ |
| | │ DataSourceExec │ |
| | │ -------------------- │ |
| | │ files: 14 │ |
| | │ format: parquet │ |
| | └───────────────────────────┘ |
| | |
+---------------+-------------------------------+
```
</details>
It's slow due to inefficient output materializing in partial and final
aggregation
For internal mechanism, this comment explains 'why not X, and do Y
instead' -- X is the existing impl, Y is what this PR does.
-
https://github.com/2010YOUY01/arrow-datafusion/blob/0f165d0129820a3bd66641609d9efb954c0cc018/datafusion/physical-plan/src/aggregates/ordered_partial_stream.rs#L219-L248
### Fix
To fully restore the performance, we have to fix:
1. Ordered partial aggregation (this PR)
2. Ordered final aggregation (maybe a follow-up PR)
2 uses almost the same mechanism as 1, so once this PR is reviewed, we
can apply the pattern mechanically.
After this PR, the query runs in: (on an M4 Pro MacBook Pro)
```
-- Still some gap due to final aggregation is not fixed yet
Current main: 3.5s
PR: 0.45s
DataFusion 54.0: 0.37s
```
## What changes are included in this PR?
<!--
There is no need to duplicate the description in the issue here, but it
is sometimes worth providing a summary of the individual changes in this
PR.
-->
1. Refactor the ordered-partial aggregation, so it's easier to implement
incremental outputting with slicing
2. Implement the output materializing strategy mentioned above
Note to read this PR, I suggest directly reading the new impl start from
the entry point of state machine (`into_stream()`), instead of the diff,
due to a large refactor.
This refactor is necessary because its easier to implement this feature
with a different state machine pattern.
## What is the testing strategy for this PR?
<!--
We typically require tests for all PRs in order to:
1. Prevent the code from being accidentally broken by subsequent changes
4. Serve as another way to document the expected behavior of the code
Briefly describe how this PR is tested, and point to the specific tests
you added. For example: 'This new feature is covered by the
`sqllogictest` cases added in `foo.slt`'.
If this PR does not add tests, explain why. For example, if the change
is already covered by existing tests, please mention it.
You should also check the `codecov` bot reply on this PR to confirm the
changed code is exercised.
-->
For correctness, existing tests have covered it.
To prevent similar perf regression, we can do
- https://github.com/apache/datafusion/issues/25310
## Are there any user-facing changes?
<!--
If there are user-facing changes then we may require documentation to be
updated before approving the PR.
If there are any breaking changes to public APIs, please add the `api
change` label.
-->
---
.../aggregate_hash_table/common_ordered.rs | 54 ++-
.../aggregate_hash_table/ordered_partial_table.rs | 28 +-
.../src/aggregates/ordered_partial_stream.rs | 396 ++++++++++++---------
3 files changed, 271 insertions(+), 207 deletions(-)
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 26a1beb0a2..9476c0991a 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
@@ -362,24 +362,6 @@ impl<AggrMode> OrderedAggregateTable<AggrMode> {
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.
///
@@ -455,20 +437,34 @@ impl<AggrMode> OrderedAggregateTable<AggrMode> {
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<RecordBatch> {
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() {
@@ -484,6 +480,6 @@ impl<AggrMode> OrderedAggregateTable<AggrMode> {
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 6ed93e59f3..b53012d8bd 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<PartialMarker> {
)
}
- /// 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<Option<RecordBatch>> {
- 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 8c2315588a..dd845353c7 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<PartialMarker>,
+}
+
+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<Aggregating>,
+}
+
+/// 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<OrderedAggregateTable<PartialMarker>>,
}
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<Item = Result<RecordBatch>> {
- 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:
+ /// <https://github.com/apache/datafusion/issues/24704>
+ ///
+ /// [`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<PartialMarker>,
- emitter: &mut TryEmitter<RecordBatch, DataFusionError>,
- ) -> 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<Option<ExecutionStage>> {
+ 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: <https://github.com/apache/datafusion/issues/7065>
- fn resize_or_take_state_batch(
- &mut self,
- table: &mut OrderedAggregateTable<PartialMarker>,
- ) -> Result<Option<RecordBatch>> {
- 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<PartialMarker>,
+ /// See [`OrderedPartialAggregateStream::into_stream`] for stage
transitions
+ /// and output memory accounting.
+ async fn handle_stage(
+ self,
+ context: &OrderedPartialAggregateContext,
+ reservation: &MemoryReservation,
emitter: &mut TryEmitter<RecordBatch, DataFusionError>,
- ) -> Result<()> {
- let elapsed_compute = self.baseline_metrics.elapsed_compute().clone();
+ ) -> Result<Option<ExecutionStage>> {
+ 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)
}
}
---------------------------------------------------------------------
To unsubscribe, e-mail: [email protected]
For additional commands, e-mail: [email protected]