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 @@ -529,6 +529,7 @@ jobs:
org.apache.comet.exec.CometFirstLastBoolAggFuzzSuite
org.apache.comet.exec.CometExec3_4PlusSuite
org.apache.comet.exec.CometExecSuite
org.apache.comet.exec.CometFilterShortCircuitSuite
org.apache.comet.exec.CometTaskBinarySizeSuite
org.apache.spark.sql.comet.CometNativeTaskKillSuite
org.apache.comet.exec.CometEmptyRelationExecSuite
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 @@ -177,6 +177,7 @@ jobs:
org.apache.comet.exec.CometFirstLastBoolAggFuzzSuite
org.apache.comet.exec.CometExec3_4PlusSuite
org.apache.comet.exec.CometExecSuite
org.apache.comet.exec.CometFilterShortCircuitSuite
org.apache.comet.exec.CometTaskBinarySizeSuite
org.apache.spark.sql.comet.CometNativeTaskKillSuite
org.apache.comet.exec.CometEmptyRelationExecSuite
Expand Down
36 changes: 33 additions & 3 deletions native/core/src/execution/jni_api.rs
Original file line number Diff line number Diff line change
Expand Up @@ -1241,6 +1241,7 @@ pub unsafe extern "system" fn Java_org_apache_comet_Native_executePlan(
PhysicalPlanner::new(Arc::clone(&exec_context.session_ctx), partition)
.with_exec_id(exec_context_id)
.with_sql_text_pool(&exec_context.spark_plan)
.with_jvm_output_sorts(&exec_context.spark_plan)
.with_task_context(exec_context.task_context.clone())
.with_class_loader(exec_context.class_loader.clone())
.with_shuffle_partition_pusher(
Expand Down Expand Up @@ -3657,6 +3658,7 @@ mod native_sort_spill_tests {
share: usize,
executor_cores: usize,
spill_before_output: Option<usize>,
jvm_consumer: bool,
) -> OutputTrace {
let shape = source.shape;
let rows = source.rows;
Expand Down Expand Up @@ -3708,7 +3710,8 @@ mod native_sort_spill_tests {
},
)])
.unwrap();
let sort: Arc<dyn ExecutionPlan> = Arc::new(SortExec::new(ordering, child));
let sort: Arc<dyn ExecutionPlan> =
Arc::new(SortExec::new(ordering, child).with_spill_before_output(jvm_consumer));
let mut stream = sort.execute(0, session.task_ctx()).unwrap();
let mut trace = OutputTrace {
reserved: vec![],
Expand Down Expand Up @@ -3744,8 +3747,14 @@ mod native_sort_spill_tests {
let share = SPILL_BEFORE_OUTPUT_SHARE;
for rows_per_batch in [8192, 3] {
let case = format!("{shape:?} rows={rows} rows_per_batch={rows_per_batch}");
let held =
sort_and_trace(KibRows::new(shape, rows, rows_per_batch), share, 8, None).await;
let held = sort_and_trace(
KibRows::new(shape, rows, rows_per_batch),
share,
8,
None,
true,
)
.await;
eprintln!(
"{case} off: spill_count={} peak={}MiB reserved at 0%={}MiB 50%={}MiB 90%={}MiB",
held.spill_count(),
Expand All @@ -3762,6 +3771,7 @@ mod native_sort_spill_tests {
share,
8,
Some(share / 4),
true,
)
.await;
eprintln!(
Expand All @@ -3782,11 +3792,31 @@ mod native_sort_spill_tests {
share,
8,
Some(held.peak_reserved),
true,
)
.await;
assert_eq!(below.spill_count(), 0, "{case}");
assert_eq!(below.reserved.len(), held.reserved.len(), "{case}");
assert!(below.reserved_at(0.9) >= input_bytes(rows) / 2, "{case}");

let native_consumer = sort_and_trace(
KibRows::new(shape, rows, rows_per_batch),
share,
8,
Some(share / 4),
false,
)
.await;
assert_eq!(native_consumer.spill_count(), 0, "{case}");
assert_eq!(
native_consumer.reserved.len(),
held.reserved.len(),
"{case}"
);
assert!(
native_consumer.reserved_at(0.9) >= input_bytes(rows) / 2,
"{case}"
);
}
}

Expand Down
166 changes: 164 additions & 2 deletions native/core/src/execution/planner.rs
Original file line number Diff line number Diff line change
Expand Up @@ -314,6 +314,29 @@ pub struct PhysicalPlanner {
/// Task-owned destination for remote shuffle blocks, registered on the driving Spark task
/// thread before native planning. Only explicit RSS destinations may use it.
shuffle_partition_pusher: Option<Arc<dyn ShufflePartitionPusher>>,
jvm_output_sorts: std::collections::HashSet<u32>,
}

pub(crate) fn jvm_output_sorts(root: &Operator) -> std::collections::HashSet<u32> {
let mut sorts = std::collections::HashSet::new();
let mut pending = vec![root];
while let Some(op) = pending.pop() {
match op.op_struct.as_ref() {
Some(OpStruct::Sort(_)) => {
sorts.insert(op.plan_id);
}
Some(
OpStruct::Projection(_)
| OpStruct::Filter(_)
| OpStruct::Limit(_)
| OpStruct::Sample(_)
| OpStruct::Expand(_)
| OpStruct::Explode(_),
) => pending.extend(op.children.iter()),
_ => {}
}
}
sorts
}

impl Default for PhysicalPlanner {
Expand All @@ -333,12 +356,18 @@ impl PhysicalPlanner {
task_context: None,
class_loader: None,
shuffle_partition_pusher: None,
jvm_output_sorts: std::collections::HashSet::new(),
}
}

/// Load the SQL text pool from the root operator of the plan about to be planned. Must be
/// called with the *root* operator: the JVM only populates the pool there, and
/// `QueryContext.sql_text_idx` values are indices into it.
pub fn with_jvm_output_sorts(mut self, root: &Operator) -> Self {
self.jvm_output_sorts = jvm_output_sorts(root);
self
}

pub fn with_sql_text_pool(mut self, root: &Operator) -> Self {
self.sql_text_pool = root
.sql_text_pool
Expand Down Expand Up @@ -1783,7 +1812,8 @@ impl PhysicalPlanner {
LexOrdering::new(exprs?).unwrap(),
Arc::clone(&child.native_plan),
)
.with_fetch(fetch),
.with_fetch(fetch)
.with_spill_before_output(self.jvm_output_sorts.contains(&spark_plan.plan_id)),
);

if let Some(skip) = sort.skip.filter(|&n| n > 0).map(|n| n as usize) {
Expand Down Expand Up @@ -5391,7 +5421,7 @@ mod tests {

use crate::execution::operators::{ExecutionError, PartitionedRankLimitExec, WindowFnKind};
use crate::execution::planner::{
convert_spark_types_to_arrow_schema, literal_to_array_ref,
convert_spark_types_to_arrow_schema, jvm_output_sorts, literal_to_array_ref,
parse_file_scan_tasks_from_common,
};
use crate::execution::shuffle::CometPartitioning;
Expand Down Expand Up @@ -6755,6 +6785,138 @@ mod tests {
}
}

fn op(plan_id: u32, op_struct: OpStruct, children: Vec<Operator>) -> Operator {
Operator {
plan_id,
sql_text_pool: vec![],
children,
op_struct: Some(op_struct),
}
}

fn sort_op(plan_id: u32, child: Operator) -> Operator {
let order = spark_expression::Expr {
expr_struct: Some(ExprStruct::SortOrder(Box::new(
spark_expression::SortOrder {
child: Some(Box::new(spark_expression::Expr {
expr_struct: Some(Bound(spark_expression::BoundReference {
index: 0,
datatype: Some(create_proto_datatype()),
})),
query_context: None,
expr_id: None,
})),
direction: spark_expression::SortDirection::Ascending as i32,
null_ordering: spark_expression::NullOrdering::NullsFirst as i32,
},
))),
query_context: None,
expr_id: None,
};
op(
plan_id,
OpStruct::Sort(spark_operator::Sort {
sort_orders: vec![order],
fetch: None,
skip: None,
}),
vec![child],
)
}

fn leaf(plan_id: u32) -> Operator {
let mut scan = create_scan();
scan.plan_id = plan_id;
scan
}

#[test]
fn jvm_output_sorts_are_sorts_reached_from_the_root_through_streaming_operators() {
use std::collections::HashSet;
let ids = |v: &[u32]| v.iter().copied().collect::<HashSet<u32>>();

assert_eq!(jvm_output_sorts(&sort_op(1, leaf(0))), ids(&[1]));

let projected = op(
3,
OpStruct::Projection(Default::default()),
vec![op(
2,
OpStruct::Filter(Default::default()),
vec![sort_op(1, leaf(0))],
)],
);
assert_eq!(jvm_output_sorts(&projected), ids(&[1]));

let nested = sort_op(
2,
op(
3,
OpStruct::Window(Default::default()),
vec![sort_op(1, leaf(0))],
),
);
assert_eq!(jvm_output_sorts(&nested), ids(&[2]));

for consumer in [
OpStruct::ShuffleWriter(Default::default()),
OpStruct::Window(Default::default()),
OpStruct::HashAgg(Default::default()),
OpStruct::ParquetWriter(Default::default()),
OpStruct::WindowGroupLimit(Default::default()),
] {
let plan = op(2, consumer, vec![sort_op(1, leaf(0))]);
assert!(jvm_output_sorts(&plan).is_empty(), "{plan:?}");
}

let join = op(
5,
OpStruct::SortMergeJoin(Default::default()),
vec![sort_op(1, leaf(0)), sort_op(2, leaf(3))],
);
assert!(jvm_output_sorts(&join).is_empty());
assert!(
jvm_output_sorts(&op(6, OpStruct::Projection(Default::default()), vec![join]))
.is_empty()
);
}

fn sort_execs(plan: &Arc<dyn ExecutionPlan>, out: &mut Vec<bool>) {
if let Some(sort) = plan.downcast_ref::<SortExec>() {
out.push(sort.spill_before_output());
}
for child in plan.children() {
sort_execs(child, out);
}
}

fn planned_sort_flags(root: &Operator) -> Vec<bool> {
let planner = PhysicalPlanner::default().with_jvm_output_sorts(root);
let (_, _, plan) = planner.create_plan(root, &mut vec![], 1).unwrap();
let mut flags = vec![];
sort_execs(&plan.native_plan, &mut flags);
flags
}

#[test]
fn planner_lets_only_a_sort_read_by_the_jvm_spill_before_output() {
assert_eq!(planned_sort_flags(&sort_op(1, leaf(0))), vec![true]);
assert_eq!(
planned_sort_flags(&create_filter(sort_op(1, leaf(0)), 3)),
vec![true]
);
assert_eq!(
planned_sort_flags(&sort_op(2, sort_op(1, leaf(0)))),
vec![true, false]
);
let sort = sort_op(1, leaf(0));
let planner = PhysicalPlanner::default();
let (_, _, plan) = planner.create_plan(&sort, &mut vec![], 1).unwrap();
let mut flags = vec![];
sort_execs(&plan.native_plan, &mut flags);
assert_eq!(flags, vec![false]);
}

fn create_scan() -> Operator {
Operator {
plan_id: 0,
Expand Down
30 changes: 24 additions & 6 deletions native/vendor/datafusion-physical-plan/src/sorts/sort.rs
Original file line number Diff line number Diff line change
Expand Up @@ -1241,6 +1241,7 @@ pub struct SortExec {
/// If `fetch` is `Some`, this will also be set and a TopK operator may be used.
/// If `fetch` is `None`, this will be `None`.
filter: Option<Arc<RwLock<TopKDynamicFilters>>>,
spill_before_output: bool,
}

impl SortExec {
Expand All @@ -1260,9 +1261,19 @@ impl SortExec {
common_sort_prefix: sort_prefix,
cache: Arc::new(cache),
filter: None,
spill_before_output: false,
}
}

pub fn with_spill_before_output(mut self, spill_before_output: bool) -> Self {
self.spill_before_output = spill_before_output;
self
}

pub fn spill_before_output(&self) -> bool {
self.spill_before_output
}

/// Whether this `SortExec` preserves partitioning of the children
pub fn preserve_partitioning(&self) -> bool {
self.preserve_partitioning
Expand Down Expand Up @@ -1336,6 +1347,7 @@ impl SortExec {
fetch: self.fetch,
cache: Arc::clone(&self.cache),
filter: self.filter.clone(),
spill_before_output: self.spill_before_output,
}
}

Expand Down Expand Up @@ -1724,10 +1736,14 @@ impl ExecutionPlan for SortExec {
execution_options.sort_spill_reservation_bytes;
let in_place_bytes =
execution_options.sort_in_place_threshold_bytes;
let spill_before_output = context
.session_config()
.get_extension::<SpillBeforeOutputThreshold>()
.map_or(0, |threshold| threshold.0);
let spill_before_output = if self.spill_before_output {
context
.session_config()
.get_extension::<SpillBeforeOutputThreshold>()
.map_or(0, |threshold| threshold.0)
} else {
0
};
let compression = context.session_config().spill_compression();
let runtime = context.runtime_env();
Ok(Box::pin(RecordBatchStreamAdapter::new(
Expand Down Expand Up @@ -1867,7 +1883,8 @@ impl ExecutionPlan for SortExec {
Ok(Some(Arc::new(
SortExec::new(updated_exprs, make_with_child(projection, self.input())?)
.with_fetch(self.fetch())
.with_preserve_partitioning(self.preserve_partitioning()),
.with_preserve_partitioning(self.preserve_partitioning())
.with_spill_before_output(self.spill_before_output),
)))
}

Expand Down Expand Up @@ -1947,7 +1964,8 @@ impl ExecutionPlan for SortExec {
let new_sort = Arc::new(
SortExec::new(self.expr.clone(), new_child)
.with_fetch(self.fetch())
.with_preserve_partitioning(self.preserve_partitioning()),
.with_preserve_partitioning(self.preserve_partitioning())
.with_spill_before_output(self.spill_before_output),
) as Arc<dyn ExecutionPlan>;

Ok(FilterPushdownPropagation {
Expand Down
Loading
Loading