2010YOUY01 commented on code in PR #24015:
URL: https://github.com/apache/datafusion/pull/24015#discussion_r4140297918


##########
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)?;
+
+                emitter
+                    
.emit(remaining_groups.record_output(&self.baseline_metrics))
+                    .await;
+
+                return Ok(());
+            }
+            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, 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(last_batch.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")
-        };
-
-        let (output_batch, next_state) = if batch.num_rows() <= 
self.batch_size {
-            // Go back to `ReadingInput`
-            (
-                batch,
-                PartialReduceHashAggregateState::ReadingInput { hash_table },
+        mut hash_table: AggregateHashTable<PartialReduceMarker>,
+        mut emitter: TryEmitter<RecordBatch, DataFusionError>,
+    ) -> Result<()> {
+        debug_assert!(!hash_table.is_building());
+
+        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);
+
+            // 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();
+        }
+    }
+}
+
+#[cfg(test)]
+mod tests {
+    use crate::ExecutionPlan;
+    use 
crate::aggregates::partial_reduce_stream::PartialReduceHashAggregateStream;
+    use crate::aggregates::{AggregateExec, AggregateMode, PhysicalGroupBy};
+    use crate::test::exec::BarrierExec;
+    use arrow::array::{AsArray, Int32Array, Int64Array, RecordBatch};
+    use arrow::datatypes::Int32Type;
+    use arrow_schema::{DataType, Field, Schema};
+    use datafusion_execution::runtime_env::RuntimeEnvBuilder;
+    use datafusion_execution::{SendableRecordBatchStream, TaskContext};
+    use datafusion_functions_aggregate::count::count_udaf;
+    use datafusion_physical_expr::aggregate::AggregateExprBuilder;
+    use datafusion_physical_expr::expressions::col;
+    use futures::StreamExt;
+    use std::sync::Arc;
+    use std::time::Duration;
+
+    /// Builds a partial reduce hash aggregate stream over a single input 
batch of
+    /// `num_groups` distinct groups, running under `memory_limit` bytes.
+    ///
+    /// The input does not signal end-of-stream until `wait_finish` is called
+    /// on the returned [`BarrierExec`], so any output produced before that can
+    /// only come from the memory pressure emission path (normal output waits
+    /// for all input).
+    fn partial_reduce_stream_under_memory_limit(
+        memory_limit: usize,
+        batch_size: usize,
+        num_groups: usize,
+    ) -> datafusion_common::Result<(
+        SendableRecordBatchStream,
+        Arc<BarrierExec>,
+        Arc<datafusion_execution::runtime_env::RuntimeEnv>,
+    )> {
+        let schema = Arc::new(Schema::new(vec![
+            Field::new("group_col", DataType::Int32, false),
+            Field::new("value_col_state", DataType::Int64, false),
+        ]));
+
+        let group_ids: Vec<i32> = (0..num_groups as i32).collect();
+        let values: Vec<i64> = vec![1; num_groups];
+
+        let batch = RecordBatch::try_new(
+            Arc::clone(&schema),
+            vec![
+                Arc::new(Int32Array::from(group_ids)),
+                Arc::new(Int64Array::from(values)),
+            ],
+        )?;
+        let input_partitions = vec![vec![batch]];
+
+        let runtime = RuntimeEnvBuilder::default()
+            .with_memory_limit(memory_limit, 1.0)
+            .build_arc()?;
+
+        let mut task_ctx = 
TaskContext::default().with_runtime(Arc::clone(&runtime));
+        let session_config = task_ctx.session_config().clone().set(
+            "datafusion.execution.batch_size",
+            &datafusion_common::ScalarValue::UInt64(Some(batch_size as u64)),
+        );
+        task_ctx = task_ctx.with_session_config(session_config);
+        let task_ctx = Arc::new(task_ctx);
+
+        // Create aggregate: COUNT(*) GROUP BY group_col
+        let group_expr = vec![(col("group_col", &schema)?, 
"group_col".to_string())];
+        let aggr_expr = vec![Arc::new(
+            AggregateExprBuilder::new(
+                count_udaf(),
+                vec![col("value_col_state", &schema)?],
             )
-        } 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,
-                },
+            .schema(Arc::clone(&schema))
+            .alias("count_value")
+            .build()?,
+        )];
+
+        let input = Arc::new(
+            BarrierExec::new(input_partitions, Arc::clone(&schema))
+                .without_start_barrier()
+                .with_finish_barrier()
+                .with_log(false),
+        );
+
+        let aggregate_exec = AggregateExec::try_new(
+            AggregateMode::PartialReduce,
+            PhysicalGroupBy::new_single(group_expr),
+            aggr_expr,
+            vec![None],
+            Arc::clone(&input) as Arc<dyn ExecutionPlan>,
+            Arc::clone(&schema),
+        )?;
+
+        let stream =
+            PartialReduceHashAggregateStream::new(&aggregate_exec, &task_ctx, 
0)?
+                .into_stream();
+
+        Ok((stream, input, runtime))
+    }
+
+    #[tokio::test]
+    async fn 
test_partial_reduce_hash_stream_accounts_held_batch_on_memory_pressure_while_slicing()
+    -> datafusion_common::Result<()> {
+        // When memory pressure triggers early emission, the materialized state
+        // batch is held while it is sliced into `batch_size` outputs. The
+        // stream must keep that held batch accounted for in its memory
+        // reservation until the last slice is emitted; before the fix the
+        // reservation was resized down to just the (emptied) hash table size,
+        // leaving the held batch unaccounted.
+
+        let batch_size = 1024;
+        // One row per group so the state batch is emitted in 4 slices
+        let num_groups = 4 * batch_size;
+
+        // Smaller than the building hash table (so pressure triggers) but 
large
+        // enough to hold the materialized state batch (so slicing can proceed)
+        let memory_limit = 100 * 1024;
+        let (mut stream, input, runtime) = 
partial_reduce_stream_under_memory_limit(
+            memory_limit,
+            batch_size,
+            num_groups,
+        )?;
+
+        // The first output batch must be a pressure-emitted slice, with the 
rest
+        // of the materialized state batch still held by the stream
+        let first = tokio::time::timeout(Duration::from_secs(5), stream.next())
+            .await
+            .expect(
+                "did not get early emit due to OOM, this probably means that 
the \
+                 memory limit is too high to trigger the OOM",
             )
-        };
+            .expect("stream ended early")?;
+        assert_eq!(first.num_rows(), batch_size);
+
+        // The emitted slice shares buffers with the held state batch, so its
+        // array memory size reflects the full held allocation
+        let held_size = first.get_array_memory_size();
+        let reserved = runtime.memory_pool.reserved();
+        assert!(

Review Comment:
   Is there anyway we can convert those tests into e2e tests, like this query 
should finish under x MB rss?
   
   Here it's asserting how should output physically layout, which is a 
implementation detail. If there are some future refactor (e.g. blocked state 
management), I think those tests have to be reworked.



-- 
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]

Reply via email to