rluvaton commented on code in PR #24015:
URL: https://github.com/apache/datafusion/pull/24015#discussion_r4144865458
##########
datafusion/physical-plan/src/aggregates/partial_reduce_stream.rs:
##########
@@ -225,327 +152,491 @@ impl PartialReduceHashAggregateStream {
self.input = Box::pin(EmptyRecordBatchStream::new(input_schema));
}
- fn break_with_err(
- error: DataFusionError,
- ) -> PartialReduceHashAggregateStateTransition {
- ControlFlow::Break((
- Poll::Ready(Some(Err(error))),
- PartialReduceHashAggregateState::Error,
- ))
+ pub(crate) fn into_stream(self) -> SendableRecordBatchStream {
+ let schema = Arc::clone(&self.schema);
+
+ Box::pin(RecordBatchStreamAdapter::new(schema, self.create_stream()))
}
- /// Handle ReadingInput state - aggregate partial state batches into the
hash table.
- ///
- /// See comments at `poll_next()` for details.
- ///
- /// Returns the next operator state with control flow decision.
- fn handle_reading_input(
- &mut self,
- cx: &mut Context<'_>,
- mut original_state: PartialReduceHashAggregateState,
- ) -> PartialReduceHashAggregateStateTransition {
- debug_assert!(matches!(
- &original_state,
- PartialReduceHashAggregateState::ReadingInput { .. }
- ));
- debug_assert!(original_state.hash_table().is_building());
-
- match self.input.poll_next_unpin(cx) {
- Poll::Pending => ControlFlow::Break((Poll::Pending,
original_state)),
- // Get a new input batch, aggregate it in the hash table
- Poll::Ready(Some(Ok(batch))) => {
- let elapsed_compute =
self.baseline_metrics.elapsed_compute().clone();
- let timer = elapsed_compute.timer();
- let result =
original_state.hash_table_mut().aggregate_batch(&batch);
- timer.done();
+ /// Entry point for the partial reduce hash aggregate.
+ fn create_stream(mut self) -> impl Stream<Item = Result<RecordBatch>> {
+ async_try_stream(|mut emitter| async move {
+ let mut hash_table: AggregateHashTable<PartialReduceMarker> =
+ self.hash_table.take().expect("must have hash table");
- if let Err(e) = result {
- return Self::break_with_err(e);
- }
+ debug_assert!(hash_table.is_building());
+ let elapsed_compute =
self.baseline_metrics.elapsed_compute().clone();
- // Update the memory reservation. If OOM, do early emit.
- self.resize_or_emit_early(original_state)
- }
- Poll::Ready(Some(Err(e))) => Self::break_with_err(e),
- // Input ends, move to output state
- Poll::Ready(None) => {
- self.close_input();
- let elapsed_compute =
self.baseline_metrics.elapsed_compute().clone();
+ while let Some(batch) = self.input.next().await.transpose()? {
let timer = elapsed_compute.timer();
- let result = original_state.hash_table_mut().start_output();
- timer.done();
- match result {
- Ok(()) => {
-
ControlFlow::Continue(original_state.into_producing_output())
+ let state = self.handle_input_batch(&batch, &mut hash_table)?;
+ // Avoid holding on the batch
+ drop(batch);
+
+ match state {
+ HandleInputResult::ProcessNext => {}
+ HandleInputResult::OOM => {
+ let materialized_group_states =
hash_table.take_state_batch()?.ok_or_else(|| {
+ internal_datafusion_err!(
+ "Partial reduce hash aggregate ran out of
memory with no aggregated groups"
+ )
+ })?;
+
+ self.early_emit_count.add(1);
+ timer.done();
+ self.emit_on_memory_pressure(
+ materialized_group_states,
+ &mut emitter,
+ hash_table.memory_size(),
+ )
+ .await?;
}
- Err(e) => Self::break_with_err(e),
}
}
+
+ let timer = elapsed_compute.timer();
+
+ self.close_input();
+ hash_table.start_output()?;
+
+ timer.done();
+
+ self.produce_output(hash_table, emitter).await?;
+
+ Ok(())
+ })
+ }
+
+ /// Aggregate partial state batch into the hash table
+ fn handle_input_batch(
+ &mut self,
+ batch: &RecordBatch,
+ hash_table: &mut AggregateHashTable<PartialReduceMarker>,
+ ) -> Result<HandleInputResult> {
+ debug_assert!(hash_table.is_building());
+ hash_table.aggregate_batch(batch)?;
+
+ let resize_result =
self.reservation.try_resize(hash_table.memory_size());
+ match resize_result {
+ Ok(()) => Ok(HandleInputResult::ProcessNext),
+ Err(DataFusionError::ResourcesExhausted(_)) =>
Ok(HandleInputResult::OOM),
+ Err(e) => Err(e),
}
}
- /// Update the memory reservation. If the reservation succeeds, continue
reading
- /// input. If OOM, clear the aggregated states in the hash table, and
early emit
- /// them immediately.
- ///
- /// Returns the next state; the caller finishes the intended task based on
it.
- ///
- /// The reservation is left at its pre-emission size while the states are
being
- /// emitted, because the cleared states are still held in memory as
- /// `remaining_groups`. The reservation will be reset after exiting the
- /// `EmittingOnMemoryPressure` state.
+ /// emit a materialized partial-state on memory pressure
+ /// batch in `batch_size`(from configuration) slices
///
/// # Implementation Note
/// All accumulated states are materialized at once, and then sliced into
- /// `batch_size` output batches. Emit them incrementally after blocked
state
- /// management is ready.
+ /// `batch_size` output batches (in case we have enough memory to hold on
them while slicing).
+ /// Emit them incrementally after blocked state management is ready.
///
/// Issue: <https://github.com/apache/datafusion/issues/7065>
- fn resize_or_emit_early(
+ async fn emit_on_memory_pressure(
&mut self,
- mut original_state: PartialReduceHashAggregateState,
- ) -> PartialReduceHashAggregateStateTransition {
- let elapsed_compute = self.baseline_metrics.elapsed_compute().clone();
- let _timer = elapsed_compute.timer(); // Stop on drop
- let resize_result = self
+ remaining_groups: RecordBatch,
+ emitter: &mut TryEmitter<RecordBatch, DataFusionError>,
+ hash_table_mem_size: usize,
+ ) -> Result<()> {
+ let remaining_groups_memory = remaining_groups.get_array_memory_size();
+
+ // Emitting clears the aggregate table and releases its
+ // accumulated memory. Update the reservation accordingly.
+ // We account here for the remaining groups memory to see if we can
return batch size states
+ // if there is not enough memory, fallback to emit large batch
+ match self
.reservation
- .try_resize(original_state.hash_table().memory_size());
-
- let oom = match resize_result {
- Ok(()) => return ControlFlow::Continue(original_state),
- Err(e @ DataFusionError::ResourcesExhausted(_)) => e,
- Err(e) => return Self::break_with_err(e),
- };
-
- let state_batch_result =
original_state.hash_table_mut().take_state_batch();
-
- match state_batch_result {
- Ok(Some(remaining_groups)) => {
- self.early_emit_count.add(1);
- ControlFlow::Continue(
- PartialReduceHashAggregateState::EmittingOnMemoryPressure {
- hash_table: original_state.into_hash_table(),
- remaining_groups,
- },
- )
+ .try_resize(hash_table_mem_size + remaining_groups_memory)
+ {
+ Ok(_) => {
+ // Continue with slicing
}
- // No accumulated group to emit, so early emission cannot release
any
- // memory: report the original error.
- Ok(None) => Self::break_with_err(oom),
- Err(e) => Self::break_with_err(e),
+ Err(DataFusionError::ResourcesExhausted(_)) => {
+ // Fail to reserve memory for the hash table + state batch
while slicing so emit a huge batch
+
+ // Try resize without holding the state batch, if it fails
there is nothing we can do
+ self.reservation.try_resize(hash_table_mem_size)?;
Review Comment:
this fails in CI since there is a test that created memory pool with size 1
byte and we don't reserve initial table allocation on creation, fixed in:
- https://github.com/apache/datafusion/pull/25900
--
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]