Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
72 changes: 66 additions & 6 deletions datafusion/physical-plan/src/aggregates/group_values/metrics.rs
Original file line number Diff line number Diff line change
Expand Up @@ -426,7 +426,58 @@ mod tests {
}

#[tokio::test]
async fn test_groupby_metrics_final_mode() -> Result<()> {
async fn test_legacy_groupby_aggregate_accumulator_metrics() -> Result<()> {
let schema = Arc::new(Schema::new(vec![
Field::new("k", DataType::UInt32, false),
Field::new("a", DataType::Float64, false),
Field::new("b", DataType::Float64, false),
]));
let batch = RecordBatch::try_new(
Arc::clone(&schema),
vec![
Arc::new(UInt32Array::from(vec![1, 2, 1, 2])),
Arc::new(Float64Array::from(vec![1.0, 2.0, 3.0, 4.0])),
Arc::new(Float64Array::from(vec![5.0, 6.0, 7.0, 8.0])),
],
)?;
let input =
TestMemoryExec::try_new_exec(&[vec![batch]], Arc::clone(&schema), None)?;
let group_by =
PhysicalGroupBy::new_single(vec![(col("k", &schema)?, "k".to_string())]);
let aggregates = vec![
sum_aggregate(&schema, "a", "SUM(a)")?,
sum_aggregate(&schema, "b", "SUM(b)")?,
];
let aggregate_exec = Arc::new(AggregateExec::try_new(
AggregateMode::Partial,
group_by,
aggregates,
vec![None, None],
input,
schema,
)?);
let task_ctx = Arc::new(
TaskContext::default().with_session_config(
SessionConfig::new()
.set_bool("datafusion.execution.enable_migration_aggregate", false),
),
);
let _result =
collect(Arc::clone(&aggregate_exec) as _, Arc::clone(&task_ctx)).await?;

let metrics = aggregate_exec.metrics().unwrap();
assert_aggregate_metric_labels(&metrics, "arguments_time");
assert_aggregate_metric_labels(&metrics, "update_time");
assert_aggregate_metric_labels(&metrics, "state_time");
assert_aggregate_metric_times_positive(&metrics, "update_time");
assert_aggregate_metric_times_positive(&metrics, "state_time");

Ok(())
}

async fn assert_groupby_metrics_final_mode(
enable_migration_aggregate: bool,
) -> Result<()> {
let schema = Arc::new(Schema::new(vec![
Field::new("a", DataType::UInt32, false),
Field::new("b", DataType::Float64, false),
Expand Down Expand Up @@ -486,12 +537,12 @@ mod tests {
schema,
)?);

let task_ctx = Arc::new(
TaskContext::default().with_session_config(
SessionConfig::new()
.set_bool("datafusion.execution.enable_migration_aggregate", true),
let task_ctx = Arc::new(TaskContext::default().with_session_config(
SessionConfig::new().set_bool(
"datafusion.execution.enable_migration_aggregate",
enable_migration_aggregate,
),
);
));
let _result =
collect(Arc::clone(&final_aggregate) as _, Arc::clone(&task_ctx)).await?;

Expand All @@ -516,4 +567,13 @@ mod tests {

Ok(())
}

#[tokio::test]
async fn test_groupby_metrics_final_mode() -> Result<()> {
for enable_migration_aggregate in [true, false] {
assert_groupby_metrics_final_mode(enable_migration_aggregate).await?;
}

Ok(())
}
}
82 changes: 64 additions & 18 deletions datafusion/physical-plan/src/aggregates/grouped_hash_stream.rs
Original file line number Diff line number Diff line change
Expand Up @@ -21,11 +21,13 @@ use std::sync::Arc;
use std::task::{Context, Poll};
use std::vec;

use super::aggregate_hash_table::accumulator_phases;
use super::order::GroupOrdering;
use super::skip_partial::SkipAggregationProbe;
use super::{AggregateExec, format_human_display};
use crate::aggregates::group_values::{
AggregateArgumentMetrics, GroupByMetrics, GroupValues, new_group_values,
AccumulatorPhase, AggregateAccumulatorMetrics, AggregateArgumentMetrics,
GroupByMetrics, GroupValues, new_group_values,
};
use crate::aggregates::order::GroupOrderingFull;
use crate::aggregates::{
Expand Down Expand Up @@ -378,6 +380,9 @@ pub(crate) struct GroupedHashAggregateStream {
/// Per-aggregate timing metrics for evaluating aggregate arguments.
aggregate_argument_metrics: AggregateArgumentMetrics,

/// Per-aggregate timing metrics for accumulator phases.
aggregate_accumulator_metrics: AggregateAccumulatorMetrics,

/// Reduction factor metric, calculated as `output_rows/input_rows` (only for partial aggregation)
reduction_factor: Option<metrics::RatioMetrics>,
}
Expand All @@ -398,12 +403,21 @@ impl GroupedHashAggregateStream {
let input = agg.input.execute(partition, Arc::clone(context))?;
let baseline_metrics = BaselineMetrics::new(&agg.metrics, partition);
let group_by_metrics = GroupByMetrics::new(&agg.metrics, partition);
let aggregate_labels = agg
.aggr_expr
.iter()
.map(|agg_expr| aggregate_metric_label(agg_expr))
.collect::<Vec<_>>();
let aggregate_argument_metrics = AggregateArgumentMetrics::new(
&agg.metrics,
partition,
agg.aggr_expr
.iter()
.map(|agg_expr| aggregate_metric_label(agg_expr)),
aggregate_labels.iter().cloned(),
);
let aggregate_accumulator_metrics = AggregateAccumulatorMetrics::new(
&agg.metrics,
partition,
aggregate_labels,
accumulator_phases(&agg.mode),
);

let timer = baseline_metrics.elapsed_compute().timer();
Expand Down Expand Up @@ -615,6 +629,7 @@ impl GroupedHashAggregateStream {
baseline_metrics,
group_by_metrics,
aggregate_argument_metrics,
aggregate_accumulator_metrics,
batch_size,
group_ordering,
input_done: false,
Expand Down Expand Up @@ -929,19 +944,25 @@ impl GroupedHashAggregateStream {
.zip(input_values.iter())
.zip(filter_values.iter());

for ((acc, values), opt_filter) in t {
for (idx, ((acc, values), opt_filter)) in t.enumerate() {
let opt_filter = opt_filter.as_ref().map(|filter| filter.as_boolean());

// Call the appropriate method on each aggregator with
// the entire input row and the relevant group indexes
if self.mode.input_mode() == AggregateInputMode::Raw
&& !self.spill_state.is_stream_merging
{
acc.update_batch(
values,
group_indices,
opt_filter,
total_num_groups,
self.aggregate_accumulator_metrics.time(
idx,
AccumulatorPhase::Update,
|| {
acc.update_batch(
values,
group_indices,
opt_filter,
total_num_groups,
)
},
)?;
} else {
assert_or_internal_err!(
Expand All @@ -951,7 +972,11 @@ impl GroupedHashAggregateStream {

// if aggregation is over intermediate states,
// use merge
acc.merge_batch(values, group_indices, total_num_groups)?;
self.aggregate_accumulator_metrics.time(
idx,
AccumulatorPhase::Merge,
|| acc.merge_batch(values, group_indices, total_num_groups),
)?;
}
self.group_by_metrics
.aggregation_time
Expand Down Expand Up @@ -1058,13 +1083,21 @@ impl GroupedHashAggregateStream {
}

// Next output each aggregate value
for acc in self.accumulators.iter_mut() {
for (idx, acc) in self.accumulators.iter_mut().enumerate() {
if self.mode.output_mode() == AggregateOutputMode::Final && !spilling {
output.push(acc.evaluate(emit_to)?)
output.push(self.aggregate_accumulator_metrics.time(
idx,
AccumulatorPhase::Evaluate,
|| acc.evaluate(emit_to),
)?)
} else {
// Output partial state: either because we're in a non-final mode,
// or because we're spilling and will merge/re-evaluate later.
output.extend(acc.state(emit_to)?)
output.extend(self.aggregate_accumulator_metrics.time(
idx,
AccumulatorPhase::State,
|| acc.state(emit_to),
)?)
}
}
drop(timer);
Expand Down Expand Up @@ -1178,8 +1211,17 @@ impl GroupedHashAggregateStream {
})
.collect::<Result<Vec<_>>>()?;
let false_filter = BooleanArray::from(vec![false]);
for (acc, args) in self.accumulators.iter_mut().zip(null_args.iter()) {
acc.update_batch(args, &[0], Some(&false_filter), total_groups)?;
for (idx, (acc, args)) in self
.accumulators
.iter_mut()
.zip(null_args.iter())
.enumerate()
{
self.aggregate_accumulator_metrics.time(
idx,
AccumulatorPhase::Update,
|| acc.update_batch(args, &[0], Some(&false_filter), total_groups),
)?;
}
}

Expand Down Expand Up @@ -1419,9 +1461,13 @@ impl GroupedHashAggregateStream {
.zip(input_values.iter())
.zip(filter_values.iter());

for ((acc, values), opt_filter) in iter {
for (idx, ((acc, values), opt_filter)) in iter.enumerate() {
let opt_filter = opt_filter.as_ref().map(|filter| filter.as_boolean());
output.extend(acc.convert_to_state(values, opt_filter)?);
output.extend(self.aggregate_accumulator_metrics.time(
idx,
AccumulatorPhase::ConvertToState,
|| acc.convert_to_state(values, opt_filter),
)?);
}

let states_batch = RecordBatch::try_new(self.schema(), output)?;
Expand Down
Loading