rluvaton commented on code in PR #24015:
URL: https://github.com/apache/datafusion/pull/24015#discussion_r4120560621
##########
datafusion/physical-plan/src/aggregates/partial_reduce_stream.rs:
##########
@@ -225,327 +152,481 @@ 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();
+ let mut last_state = HandleInputResult::ProcessNext;
+ 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())
+ last_state = self.handle_input_batch(batch, &mut hash_table)?;
+
+ match last_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
- .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),
- };
+ 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(hash_table_mem_size + remaining_groups_memory)
+ {
+ Ok(_) => {
+ // Continue with slicing
+ }
+ 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)?;
- 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,
- },
- )
+ emitter
+ .emit(remaining_groups.record_output(&self.baseline_metrics))
+ .await;
+
+ return Ok(());
}
- // 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(e) => return Err(e),
+ }
+
+ let mut index = 0;
+
+ while index + self.batch_size < remaining_groups.num_rows() {
+ // More batch to output
+ let output = remaining_groups.slice(index, index +
self.batch_size);
+ index += self.batch_size;
+
+ emitter
+ .emit(output.record_output(&self.baseline_metrics))
+ .await;
}
+
+ let last_batch = remaining_groups.slice(index,
remaining_groups.num_rows() - index);
+
+ debug_assert!(last_batch.num_rows() > 0);
+ debug_assert!(last_batch.num_rows() <= self.batch_size);
+
+ // We are no longer holding on the batch while slicing, so release the
memory.
+ // The memory will now equal to the hash table size
+ self.reservation.try_shrink(remaining_groups_memory)?;
+
+ emitter
+ .emit(remaining_groups.record_output(&self.baseline_metrics))
+ .await;
+
+ Ok(())
}
- /// Handle EmittingOnMemoryPressure state - emit a materialized
partial-state
- /// batch in `batch_size`(from configuration) slices. After all slices are
- /// emitted, update the memory reservation and resume reading input.
- ///
- /// See comments at `poll_next()` for details.
- ///
- /// Returns the next operator state with control flow decision.
- fn handle_emitting_on_memory_pressure(
+ /// Emit merged partial aggregate state batches.
+ async fn produce_output(
&mut self,
- original_state: PartialReduceHashAggregateState,
- ) -> PartialReduceHashAggregateStateTransition {
- let PartialReduceHashAggregateState::EmittingOnMemoryPressure {
- hash_table,
- remaining_groups: batch,
- } = original_state
- else {
- unreachable!("expected the EmittingOnMemoryPressure state")
- };
+ mut hash_table: AggregateHashTable<PartialReduceMarker>,
+ mut emitter: TryEmitter<RecordBatch, DataFusionError>,
+ ) -> Result<()> {
+ debug_assert!(!hash_table.is_building());
- let (output_batch, next_state) = if batch.num_rows() <=
self.batch_size {
- // Go back to `ReadingInput`
- (
- batch,
- PartialReduceHashAggregateState::ReadingInput { hash_table },
- )
- } else {
- // More batches to output, continue in the current state.
- let remaining =
- batch.slice(self.batch_size, batch.num_rows() -
self.batch_size);
- let output = batch.slice(0, self.batch_size);
- (
- output,
- PartialReduceHashAggregateState::EmittingOnMemoryPressure {
- hash_table,
- remaining_groups: remaining,
- },
- )
- };
+ let elapsed_compute = self.baseline_metrics.elapsed_compute().clone();
+
+ let mut timer = elapsed_compute.timer();
+
+ loop {
+ let Some(batch) = hash_table.next_output_batch()? else {
+ // Only reachable when the table held no groups at all: a
+ // non-empty table always reports its last batch together with
+ // the `Done` state, which the `try_resize` below already
zeroes.
+ self.reservation.try_resize(0)?;
+ return Ok(());
+ };
+
+ debug_assert!(batch.num_rows() > 0);
- debug_assert!(output_batch.num_rows() > 0);
- ControlFlow::Break((
-
Poll::Ready(Some(Ok(output_batch.record_output(&self.baseline_metrics)))),
- next_state,
- ))
+ // The table hands over its groups as they are materialized and
+ // reports a size of 0 once it reaches `Done`, so this releases the
+ // reservation before the final batch goes downstream.
+ // The output is already materialized, so a failed resize cannot
+ // be acted on: keep the reservation as is and finish the output.
+ let _ = self.reservation.try_resize(hash_table.memory_size());
+
+ timer.done();
+ emitter
+ .emit(batch.record_output(&self.baseline_metrics))
+ .await;
+ timer = elapsed_compute.timer();
+ };
}
+}
- /// Handle ProducingOutput state - emit merged partial aggregate state
batches.
- ///
- /// See comments at `poll_next()` for details.
+#[cfg(test)]
Review Comment:
Added the tests here since this change include the fix that I made in:
- https://github.com/apache/datafusion/pull/24852
--
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]