Skip to content
Merged
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
1 change: 1 addition & 0 deletions .github/workflows/pr_build_linux.yml
Original file line number Diff line number Diff line change
Expand Up @@ -526,6 +526,7 @@ jobs:
- name: "exec"
value: |
org.apache.comet.exec.CometAggregateSuite
org.apache.comet.exec.CometFirstLastBoolAggFuzzSuite
org.apache.comet.exec.CometExec3_4PlusSuite
org.apache.comet.exec.CometExecSuite
org.apache.comet.exec.CometTaskBinarySizeSuite
Expand Down
1 change: 1 addition & 0 deletions .github/workflows/pr_build_macos.yml
Original file line number Diff line number Diff line change
Expand Up @@ -174,6 +174,7 @@ jobs:
- name: "exec"
value: |
org.apache.comet.exec.CometAggregateSuite
org.apache.comet.exec.CometFirstLastBoolAggFuzzSuite
org.apache.comet.exec.CometExec3_4PlusSuite
org.apache.comet.exec.CometExecSuite
org.apache.comet.exec.CometTaskBinarySizeSuite
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -47,14 +47,14 @@ import org.apache.comet.serde.OperatorOuterClass.Operator
* `TreeNode.makeCopy` on MERGE re-planning (the CometIcebergNativeScanExec lesson).
*/
case class CometDeltaNativeScanExec(
override val nativeOp: Operator,
@transient override val nativeOp: Operator,
override val output: Seq[Attribute],
requiredSchema: StructType,
runtimeFilters: Seq[Expression],
dataFilters: Seq[Expression],
@transient relation: HadoopFsRelation,
originalPlan: FileSourceScanExec,
override val serializedPlanOpt: SerializedPlan,
@transient override val serializedPlanOpt: SerializedPlan,
sourceKey: String)
extends CometLeafExec
with CometScanWithPlanData {
Expand Down
6 changes: 5 additions & 1 deletion docs/source/user-guide/latest/tuning.md
Original file line number Diff line number Diff line change
Expand Up @@ -602,7 +602,11 @@ prices per row from a table of measurements: `c0 + k0*L + k1*L*min(L, 600)` ns n
Spark one.
- `sort` over every leaf of the sorted rows (Comet also 0.15 ns per byte beyond 12 per leaf); `sortSpill` adds the
price of a spill to a fraction `sortSpillFraction` of rows, none by default. `smj` prices a sort-merge join and `bhj`
the probe side of a broadcast hash join, over every output leaf.
the probe side of a broadcast hash join, over every output leaf. A sort-merge join with a join condition adds
`smjCondition`, also over every output leaf: Comet builds every pair of rows of equal keys before the condition drops
them, about 3.5 times Spark's price on band joins. A condition that is one validity interval, `L <= V < U` with `V`
from one input and `L` and `U` from the other (casts, date truncations, `COALESCE(U, literal)` and `U IS NULL OR`
allowed), adds nothing when `U` is not `L` shifted by a constant through the projections below the join.
- `predicate` for filters, over the leaves their predicate references, plus a pass-through of 1.5 ns per output leaf
natively and none in Spark. A native filter over a native scan, and the native projects over it, stay native
whatever their prices (`keepFiltersOverNativeScans=false` lets them move): the rows a filter drops are not estimated,
Expand Down
167 changes: 167 additions & 0 deletions native/core/src/execution/jni_api.rs
Original file line number Diff line number Diff line change
Expand Up @@ -4225,3 +4225,170 @@ mod aggregate_offset_overflow_tests {
);
}
}

#[cfg(test)]
mod partial_merge_spill_tests {
use super::*;
use crate::execution::memory_pools::fair_unified_pool_with_fake_spark;
use crate::execution::operators::InputBatch;
use arrow::array::{Int32Array, Int64Array};
use datafusion_comet_proto::spark_expression::{self, Expr};
use datafusion_comet_proto::spark_operator::{self, operator::OpStruct};
use std::collections::HashMap as StdHashMap;

const KEYS: i32 = 4;
const VALUES: i32 = 6000;
const COPIES: usize = 3;
const BATCH_ROWS: usize = 1024;

fn data_type(type_id: i32) -> spark_expression::DataType {
spark_expression::DataType {
type_id,
type_info: None,
}
}

fn bound(index: i32, type_id: i32) -> Expr {
Expr {
expr_struct: Some(spark_expression::expr::ExprStruct::Bound(
spark_expression::BoundReference {
index,
datatype: Some(data_type(type_id)),
},
)),
query_context: None,
expr_id: None,
}
}

/// `HashAggregate(keys = [k, x], functions = [merge count])`: the de-duplicating stage
/// Spark plans under a single COUNT(DISTINCT x) GROUP BY k, fed `(k, x, count)` states.
fn partial_merge_count() -> Operator {
let scan = Operator {
op_struct: Some(OpStruct::Scan(spark_operator::Scan {
fields: vec![data_type(3), data_type(3), data_type(4)],
source: "states".to_string(),
})),
..Default::default()
};
let count = spark_expression::AggExpr {
expr_struct: Some(AggExprStruct::Count(spark_expression::Count {
children: vec![bound(1, 3)],
})),
..Default::default()
};
Operator {
children: vec![scan],
op_struct: Some(OpStruct::HashAgg(spark_operator::HashAggregate {
grouping_exprs: vec![bound(0, 3), bound(1, 3)],
agg_exprs: vec![count],
mode: AggregateMode::PartialMerge as i32,
expr_modes: vec![AggregateMode::PartialMerge as i32],
initial_input_buffer_offset: 2,
})),
..Default::default()
}
}

fn input_batches() -> Vec<InputBatch> {
let mut rows: Vec<(i32, i32)> = (0..COPIES)
.flat_map(|_| (0..KEYS).flat_map(|k| (0..VALUES).map(move |x| (k, x))))
.collect();
let mut seed = 0x9e37_79b9_7f4a_7c15u64;
for i in (1..rows.len()).rev() {
seed ^= seed << 13;
seed ^= seed >> 7;
seed ^= seed << 17;
rows.swap(i, (seed % (i as u64 + 1)) as usize);
}
rows.chunks(BATCH_ROWS)
.map(|chunk| {
InputBatch::Batch(
vec![
Arc::new(Int32Array::from_iter_values(chunk.iter().map(|r| r.0))),
Arc::new(Int32Array::from_iter_values(chunk.iter().map(|r| r.1))),
Arc::new(Int64Array::from(vec![1i64; chunk.len()])),
],
chunk.len(),
)
})
.chain(std::iter::once(InputBatch::EOF))
.collect()
}

/// Under memory pressure a PartialMerge must spill rather than emit a group twice, or the
/// COUNT(DISTINCT) above it counts the repeated value again.
#[tokio::test]
async fn partial_merge_spills_instead_of_emitting_a_group_twice() {
let share = 256 * 1024;
let (pool, _spark) = fair_unified_pool_with_fake_spark(share, share);
let spill_dir = tempfile::tempdir().unwrap();
let operator = partial_merge_count();
let session = Arc::new(
prepare_datafusion_session_context(
BATCH_ROWS,
PlanCancellation::new(),
pool,
vec![spill_dir.path().to_string_lossy().into_owned()],
u64::MAX,
1,
&StdHashMap::new(),
&operator,
Some(share),
)
.unwrap(),
);
let planner = PhysicalPlanner::new(Arc::clone(&session), 0);
let (mut scans, _, plan) = planner.create_plan(&operator, &mut vec![], 1).unwrap();
let mut stream = plan.native_plan.execute(0, session.task_ctx()).unwrap();
let mut input = input_batches().into_iter();

let mut counts: StdHashMap<(i32, i32), Vec<i64>> = StdHashMap::new();
while let Some(batch) = futures::future::poll_fn(|cx| {
let result = stream.poll_next_unpin(cx);
if result.is_pending() && scans[0].batch.try_lock().unwrap().is_none() {
if let Some(batch) = input.next() {
scans[0].set_input_batch(batch);
cx.waker().wake_by_ref();
}
}
result
})
.await
{
let batch = batch.unwrap();
let k = batch
.column(0)
.as_any()
.downcast_ref::<Int32Array>()
.unwrap();
let x = batch
.column(1)
.as_any()
.downcast_ref::<Int32Array>()
.unwrap();
let c = batch
.column(2)
.as_any()
.downcast_ref::<Int64Array>()
.unwrap();
for row in 0..batch.num_rows() {
counts
.entry((k.value(row), x.value(row)))
.or_default()
.push(c.value(row));
}
}

let spills = plan
.native_plan
.metrics()
.and_then(|m| m.spill_count())
.unwrap_or(0);
let repeated = counts.values().filter(|c| c.len() > 1).count();
assert_eq!(repeated, 0, "{repeated} groups emitted more than once");
assert_eq!(counts.len(), (KEYS * VALUES) as usize);
assert!(counts.values().all(|c| c == &[COPIES as i64]));
assert!(spills > 0, "the PartialMerge aggregate did not spill");
}
}
39 changes: 25 additions & 14 deletions native/core/src/execution/planner.rs
Original file line number Diff line number Diff line change
Expand Up @@ -61,6 +61,7 @@ use arrow::datatypes::{
};
use arrow::ffi_stream::FFI_ArrowArrayStream;
use datafusion::functions_aggregate::bit_and_or_xor::{bit_and_udaf, bit_or_udaf, bit_xor_udaf};
use datafusion::functions_aggregate::bool_and_or::{bool_and_udaf, bool_or_udaf};
use datafusion::functions_aggregate::count::count_udaf;
use datafusion::functions_aggregate::min_max::max_udaf;
use datafusion::functions_aggregate::min_max::min_udaf;
Expand All @@ -73,7 +74,6 @@ use datafusion::{
common::DataFusionError,
config::ConfigOptions,
execution::FunctionRegistry,
functions_aggregate::first_last::{FirstValue, LastValue},
logical_expr::Operator as DataFusionOperator,
physical_expr::{
expressions::{
Expand Down Expand Up @@ -155,8 +155,8 @@ use datafusion_comet_spark_expr::{
jvm_udf::JvmScalarUdfExpr, spark_in_list, ApproxPercentile, ArrayInsert, Avg, AvgDecimal, Cast,
CheckOverflow, Correlation, Covariance, CreateNamedStruct, DecimalRescaleCheckOverflow,
GetArrayStructFields, GetStructField, HllPlusPlus, IfExpr, ListExtract, MaxMinBy, Mode,
NormalizeNaNAndZero, Regr, RegrType, SparkCastOptions, Stddev, SumDecimal, ToJson,
UnboundColumn, Variance, WideDecimalBinaryExpr, WideDecimalOp,
NormalizeNaNAndZero, Regr, RegrType, SparkCastOptions, SparkFirstLast, Stddev, SumDecimal,
ToJson, UnboundColumn, Variance, WideDecimalBinaryExpr, WideDecimalOp,
};
use itertools::Itertools;
use jni::objects::{Global, JObject};
Expand Down Expand Up @@ -1544,19 +1544,20 @@ impl PhysicalPlanner {
agg.mode
))
})?;
// A PartialMerge feeds groups it must not repeat (the distinct values of a
// single COUNT(DISTINCT)), so it runs as PartialReduce, which spills under
// memory pressure where Partial emits groups early.
let mode = match proto_mode {
ProtoAggregateMode::Partial => DFAggregateMode::Partial,
ProtoAggregateMode::Final => DFAggregateMode::Final,
// PartialMerge: Partial + MergeAsPartial
ProtoAggregateMode::PartialMerge => DFAggregateMode::Partial,
ProtoAggregateMode::PartialMerge => DFAggregateMode::PartialReduce,
};

// Check if any expression uses PartialMerge mode. When present,
// those expressions are wrapped with MergeAsPartial to get merge
// semantics inside a Partial-mode AggregateExec.
// A mixed {Partial, PartialMerge} aggregate runs as Partial and wraps its
// PartialMerge expressions with MergeAsPartial to get merge semantics.
let partial_merge_value = ProtoAggregateMode::PartialMerge as i32;
let has_partial_merge = proto_mode == ProtoAggregateMode::PartialMerge
|| agg.expr_modes.contains(&partial_merge_value);
let has_partial_merge = proto_mode == ProtoAggregateMode::Partial
&& agg.expr_modes.contains(&partial_merge_value);

let agg_exprs: PhyAggResult = agg
.agg_exprs
Expand Down Expand Up @@ -2937,8 +2938,13 @@ impl PhysicalPlanner {
let child = self.create_expr(expr.child.as_ref().unwrap(), Arc::clone(&schema))?;
let datatype = to_arrow_datatype(expr.datatype.as_ref().unwrap());
let child = Arc::new(CastExpr::new(child, datatype.clone(), None));
let func = if datatype == DataType::Boolean {
bool_and_udaf()
} else {
min_udaf()
};

AggregateExprBuilder::new(min_udaf(), vec![child])
AggregateExprBuilder::new(func, vec![child])
.schema(schema)
.alias("min")
.with_ignore_nulls(false)
Expand All @@ -2950,8 +2956,13 @@ impl PhysicalPlanner {
let child = self.create_expr(expr.child.as_ref().unwrap(), Arc::clone(&schema))?;
let datatype = to_arrow_datatype(expr.datatype.as_ref().unwrap());
let child = Arc::new(CastExpr::new(child, datatype.clone(), None));
let func = if datatype == DataType::Boolean {
bool_or_udaf()
} else {
max_udaf()
};

AggregateExprBuilder::new(max_udaf(), vec![child])
AggregateExprBuilder::new(func, vec![child])
.schema(schema)
.alias("max")
.with_ignore_nulls(false)
Expand Down Expand Up @@ -3032,7 +3043,7 @@ impl PhysicalPlanner {
}
AggExprStruct::First(expr) => {
let child = self.create_expr(expr.child.as_ref().unwrap(), Arc::clone(&schema))?;
let func = AggregateUDF::new_from_impl(FirstValue::new());
let func = AggregateUDF::new_from_impl(SparkFirstLast::first());

AggregateExprBuilder::new(Arc::new(func), vec![child])
.schema(schema)
Expand All @@ -3044,7 +3055,7 @@ impl PhysicalPlanner {
}
AggExprStruct::Last(expr) => {
let child = self.create_expr(expr.child.as_ref().unwrap(), Arc::clone(&schema))?;
let func = AggregateUDF::new_from_impl(LastValue::new());
let func = AggregateUDF::new_from_impl(SparkFirstLast::last());

AggregateExprBuilder::new(Arc::new(func), vec![child])
.schema(schema)
Expand Down
4 changes: 4 additions & 0 deletions native/spark-expr/Cargo.toml
Original file line number Diff line number Diff line change
Expand Up @@ -95,6 +95,10 @@ harness = false
name = "aggregate"
harness = false

[[bench]]
name = "first_last"
harness = false

[[bench]]
name = "approx_percentile"
harness = false
Expand Down
Loading
Loading