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
Original file line number Diff line number Diff line change
Expand Up @@ -31,7 +31,7 @@ use auron_jni_bridge::{
use datafusion::{
common::Result,
execution::{SendableRecordBatchStream, TaskContext},
physical_expr::{EquivalenceProperties, Partitioning, PhysicalExprRef},
physical_expr::{EquivalenceProperties, Partitioning, PhysicalExprRef, expressions::lit},
physical_plan::{
DisplayAs, DisplayFormatType, ExecutionPlan, ExecutionPlanProperties, PlanProperties,
execution_plan::{Boundedness, EmissionType},
Expand Down Expand Up @@ -135,6 +135,16 @@ impl ExecutionPlan for BroadcastJoinBuildHashMapExec {
}
}

/// Use a constant key for keyless joins so SMJ fallback forms one Cartesian
/// group.
pub(crate) fn smj_fallback_keys(keys: &[PhysicalExprRef]) -> Vec<PhysicalExprRef> {
if keys.is_empty() {
vec![lit(0i32)]
} else {
keys.to_vec()
}
}

pub fn execute_build_hash_map(
mut input: SendableRecordBatchStream,
keys: Vec<PhysicalExprRef>,
Expand Down Expand Up @@ -205,7 +215,7 @@ pub fn execute_build_hash_map(
let input_exec = create_record_batch_stream_exec(input, exec_ctx.partition_id())?;
let sort_exec = create_default_ascending_sort_exec(
input_exec,
&keys,
&smj_fallback_keys(&keys),
Some(exec_ctx.execution_plan_metrics().clone()),
false, // do not record output metric
);
Expand Down Expand Up @@ -234,3 +244,215 @@ pub fn execute_build_hash_map(
Ok(())
}))
}

#[cfg(test)]
mod tests {
use arrow::array::Int32Array;
use arrow_schema::{Field, Schema};
use auron_memmgr::MemManager;
use datafusion::{
common::{JoinSide, ScalarValue},
logical_expr::Operator,
physical_expr::expressions::{BinaryExpr, Column},
physical_plan::{common::collect, joins::utils::build_join_schema, test::TestMemoryExec},
prelude::SessionContext,
};

use super::*;
use crate::{
broadcast_join_exec::BroadcastJoinExec,
common::column_pruning::ExecuteWithColumnPruning,
joins::{ColumnIndex, JoinFilter, join_utils::JoinType},
};

fn memory(batch: RecordBatch) -> Result<Arc<dyn ExecutionPlan>> {
Ok(Arc::new(TestMemoryExec::try_new(
&[vec![batch.clone()]],
batch.schema(),
None,
)?))
}

#[tokio::test]
async fn nested_loop_sorted_fallback_preserves_condition_and_projection() -> Result<()> {
MemManager::init(1000000);
let values = [
vec![None, Some(1), Some(3), Some(8)],
vec![None, Some(2), Some(4)],
];
let batches = values
.iter()
.enumerate()
.map(|(side, values)| {
RecordBatch::try_new(
Arc::new(Schema::new(vec![Field::new(
if side == 0 { "l" } else { "r" },
DataType::Int32,
true,
)])),
vec![Arc::new(Int32Array::from(values.clone()))],
)
.expect("valid test batch")
})
.collect::<Vec<_>>();
for (jt, build_left) in [
(JoinType::Inner, false),
(JoinType::Inner, true),
(JoinType::Left, false),
(JoinType::Right, true),
(JoinType::LeftSemi, false),
(JoinType::LeftAnti, false),
(JoinType::Existence, false),
] {
let schema = if jt == JoinType::Existence {
Arc::new(Schema::new(vec![
batches[0].schema().field(0).clone(),
Field::new("exists", DataType::Boolean, false),
]))
} else {
Arc::new(
build_join_schema(&batches[0].schema(), &batches[1].schema(), &jt.try_into()?)
.0,
)
};
let mut expected = vec![];
let mut right_matched = vec![false; values[1].len()];
for &l in &values[0] {
let mut matched = false;
for (ri, &r) in values[1].iter().enumerate() {
if l.zip(r).is_some_and(|(l, r)| l < r) {
matched = true;
right_matched[ri] = true;
if matches!(jt, JoinType::Inner | JoinType::Left | JoinType::Right) {
expected.push(vec![ScalarValue::Int32(l), ScalarValue::Int32(r)]);
}
}
}
match jt {
JoinType::Left if !matched => {
expected.push(vec![ScalarValue::Int32(l), ScalarValue::Int32(None)])
}
JoinType::LeftSemi if matched => expected.push(vec![ScalarValue::Int32(l)]),
JoinType::LeftAnti if !matched => expected.push(vec![ScalarValue::Int32(l)]),
JoinType::Existence => expected.push(vec![
ScalarValue::Int32(l),
ScalarValue::Boolean(Some(matched)),
]),
_ => {}
}
}
if jt == JoinType::Right {
for (&r, matched) in values[1].iter().zip(right_matched) {
if !matched {
expected.push(vec![ScalarValue::Int32(None), ScalarValue::Int32(r)]);
}
}
}
let ctx = SessionContext::new().task_ctx();
let build_batch = batches[usize::from(!build_left)].clone();
// A null table column marks the build stream as sorted for SMJ fallback.
let sorted = create_default_ascending_sort_exec(
memory(build_batch)?,
&smj_fallback_keys(&[]),
None,
false,
);
let sorted_batches = collect(sorted.execute(0, ctx.clone())?).await?;
let mut spill_batches = vec![];
for batch in sorted_batches {
let schema = join_hash_map_schema(&batch.schema());
let cols = [
batch.columns().to_vec(),
vec![new_null_array(&DataType::Binary, batch.num_rows())],
]
.concat();
spill_batches.push(RecordBatch::try_new(schema, cols)?);
}
let spill_schema = spill_batches[0].schema();
let built = Arc::new(TestMemoryExec::try_new(
&[spill_batches],
spill_schema,
None,
)?) as Arc<dyn ExecutionPlan>;
let (left, right) = if build_left {
(built, memory(batches[1].clone())?)
} else {
(memory(batches[0].clone())?, built)
};
let filter = JoinFilter {
expression: Arc::new(BinaryExpr::new(
Arc::new(Column::new("l", 0)),
Operator::Lt,
Arc::new(Column::new("r", 1)),
)),
column_indices: vec![
ColumnIndex {
side: JoinSide::Left,
index: 0,
},
ColumnIndex {
side: JoinSide::Right,
index: 0,
},
],
schema: Arc::new(Schema::new(vec![
batches[0].schema().field(0).clone(),
batches[1].schema().field(0).clone(),
])),
};
let join = BroadcastJoinExec::try_new(
schema.clone(),
left,
right,
vec![],
jt,
if build_left {
JoinSide::Left
} else {
JoinSide::Right
},
true,
None,
false,
Some(filter),
)?;
for projection in [
(0..schema.fields().len()).collect::<Vec<_>>(),
vec![schema.fields().len() - 1],
vec![],
] {
let output = collect(join.execute_projected(0, ctx.clone(), &projection)?).await?;
let mut actual = vec![];
for batch in output {
for row in 0..batch.num_rows() {
actual.push(
batch
.columns()
.iter()
.map(|col| {
ScalarValue::try_from_array(col, row).map(|v| v.to_string())
})
.collect::<Result<Vec<_>>>()?,
);
}
}
let mut wanted = expected
.iter()
.map(|row| {
projection
.iter()
.map(|&i| row[i].to_string())
.collect::<Vec<_>>()
})
.collect::<Vec<_>>();
actual.sort();
wanted.sort();
assert_eq!(
actual, wanted,
"join={jt:?} build_left={build_left} projection={projection:?}"
);
}
}
Ok(())
}
}
18 changes: 8 additions & 10 deletions native-engine/datafusion-ext-plans/src/broadcast_join_exec.rs
Original file line number Diff line number Diff line change
Expand Up @@ -50,7 +50,7 @@ use once_cell::sync::OnceCell;
use parking_lot::Mutex;

use crate::{
broadcast_join_build_hash_map_exec::execute_build_hash_map,
broadcast_join_build_hash_map_exec::{execute_build_hash_map, smj_fallback_keys},
common::{
column_pruning::ExecuteWithColumnPruning,
execution_context::{ExecutionContext, WrappedRecordBatchSender},
Expand Down Expand Up @@ -458,21 +458,24 @@ async fn execute_join_with_smj_fallback(
create_record_batch_stream_exec(remoted_stream, exec_ctx.partition_id())?
};

let left_keys = smj_fallback_keys(&join_params.left_keys);
let right_keys = smj_fallback_keys(&join_params.right_keys);

// create sorted streams, build side is already sorted
let (left_exec, right_exec) = match broadcast_side {
JoinSide::Left => (
built_sorted,
create_default_ascending_sort_exec(
probed_plan,
&join_params.right_keys,
&right_keys,
Some(exec_ctx.execution_plan_metrics().clone()),
false, // do not record output metric
),
),
JoinSide::Right => (
create_default_ascending_sort_exec(
probed_plan,
&join_params.left_keys,
&left_keys,
Some(exec_ctx.execution_plan_metrics().clone()),
false, // do not record output metric
),
Expand All @@ -486,15 +489,10 @@ async fn execute_join_with_smj_fallback(
join_params.output_schema,
left_exec.clone(),
right_exec.clone(),
join_params
.left_keys
.to_vec()
.into_iter()
.zip(join_params.right_keys.to_vec())
.collect(),
left_keys.iter().cloned().zip(right_keys).collect(),
join_params.join_type,
join_params.join_filter.clone(),
vec![SortOptions::default(); join_params.left_keys.len()],
vec![SortOptions::default(); left_keys.len()],
)?);
let mut projection = match join_params.join_type {
Inner | Left | Right | Full => join_params
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -18,12 +18,14 @@ package org.apache.spark.sql.execution.joins.auron.plan

import org.apache.spark.sql.auron.join.JoinBuildSides.{JoinBuildLeft, JoinBuildRight, JoinBuildSide}
import org.apache.spark.sql.catalyst.expressions.Expression
import org.apache.spark.sql.catalyst.expressions.SortOrder
import org.apache.spark.sql.catalyst.plans.JoinType
import org.apache.spark.sql.catalyst.plans.physical.Partitioning
import org.apache.spark.sql.execution.SparkPlan
import org.apache.spark.sql.execution.auron.plan.NativeBroadcastJoinBase
import org.apache.spark.sql.execution.joins.HashJoin

import org.apache.auron.spark.configuration.SparkAuronConfiguration
import org.apache.auron.sparkver

case class NativeBroadcastJoinExec(
Expand All @@ -48,6 +50,16 @@ case class NativeBroadcastJoinExec(
isNullAwareAntiJoin)
with HashJoin {

// Keyless joins and SMJ fallback cannot guarantee the original probe ordering.
override def outputOrdering: Seq[SortOrder] = {
if ((leftKeys.isEmpty && rightKeys.isEmpty) || SparkAuronConfiguration.SMJ_FALLBACK_ENABLE
.get()) {
Nil
} else {
super.outputOrdering
}
}

@sparkver("3.1 / 3.2 / 3.3 / 3.4 / 3.5 / 4.0 / 4.1 / 4.2")
override def buildSide: org.apache.spark.sql.catalyst.optimizer.BuildSide =
broadcastSide match {
Expand Down
Loading
Loading