Skip to content
Open
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
20 changes: 3 additions & 17 deletions native-engine/datafusion-ext-plans/src/rss_shuffle_writer_exec.rs
Original file line number Diff line number Diff line change
Expand Up @@ -24,7 +24,7 @@ use datafusion::{
arrow::datatypes::SchemaRef,
error::{DataFusionError, Result},
execution::context::TaskContext,
physical_expr::{EquivalenceProperties, PhysicalExprRef, expressions::Column},
physical_expr::EquivalenceProperties,
physical_plan,
physical_plan::{
DisplayAs, DisplayFormatType, ExecutionPlan, PlanProperties, SendableRecordBatchStream,
Expand All @@ -42,7 +42,7 @@ use crate::{
rss_single_repartitioner::RssSingleShuffleRepartitioner,
rss_sort_repartitioner::RssSortShuffleRepartitioner,
},
sort_exec::create_default_ascending_sort_exec,
shuffle_writer_exec::create_round_robin_sort_exec,
};

/// The rss shuffle writer operator maps each input partition to M output
Expand Down Expand Up @@ -143,21 +143,7 @@ impl ExecutionPlan for RssShuffleWriterExec {
partitioner
}
Partitioning::RoundRobinPartitioning(..) => {
input = create_default_ascending_sort_exec(
input,
self.input
.schema()
.fields()
.iter()
.enumerate()
.map(|(index, field)| {
Arc::new(Column::new(&field.name(), index)) as PhysicalExprRef
})
.collect::<Vec<_>>()
.as_ref(),
None,
false, // do not record output metric
);
input = create_round_robin_sort_exec(input, partition)?;

let partitioner = Arc::new(RssSortShuffleRepartitioner::new(
partition,
Expand Down
85 changes: 68 additions & 17 deletions native-engine/datafusion-ext-plans/src/shuffle_writer_exec.rs
Original file line number Diff line number Diff line change
Expand Up @@ -17,20 +17,24 @@

use std::{any::Any, fmt::Debug, sync::Arc};

use arrow::datatypes::SchemaRef;
use arrow::datatypes::{DataType, Field, SchemaRef};
use async_trait::async_trait;
use auron_memmgr::MemManager;
use datafusion::{
error::Result,
execution::context::TaskContext,
physical_expr::{EquivalenceProperties, PhysicalExprRef, expressions::Column},
logical_expr::Volatility,
physical_expr::{
EquivalenceProperties, PhysicalExprRef, ScalarFunctionExpr, expressions::Column,
},
physical_plan,
physical_plan::{
DisplayAs, DisplayFormatType, ExecutionPlan, PlanProperties, SendableRecordBatchStream,
Statistics,
execution_plan::{Boundedness, EmissionType},
metrics::{ExecutionPlanMetricsSet, MetricsSet},
},
prelude::create_udf,
};
use datafusion_ext_commons::df_execution_err;
use once_cell::sync::OnceCell;
Expand All @@ -44,6 +48,67 @@ use crate::{
sort_exec::create_default_ascending_sort_exec,
};

fn contains_map(data_type: &DataType) -> bool {
match data_type {
DataType::Map(..) => true,
DataType::List(field) => contains_map(field.data_type()),
DataType::Struct(fields) => fields.iter().any(|field| contains_map(field.data_type())),
_ => false,
}
}

pub(crate) fn create_round_robin_sort_exec(
input: Arc<dyn ExecutionPlan>,
partition: usize,
) -> Result<Arc<dyn ExecutionPlan>> {
let sort_exprs = if input
.schema()
.fields()
.iter()
.any(|field| contains_map(field.data_type()))
{
let function_name = "Spark_Murmur3Hash";
let function =
datafusion_ext_functions::create_auron_ext_function(function_name, partition)?;
let args = input
.schema()
.fields()
.iter()
.enumerate()
.map(|(index, field)| Arc::new(Column::new(field.name(), index)) as PhysicalExprRef)
.collect::<Vec<_>>();
let udf = Arc::new(create_udf(
function_name,
args.iter()
.map(|expr| expr.data_type(&input.schema()))
.collect::<Result<Vec<_>>>()?,
DataType::Int32,
Volatility::Immutable,
function,
));
vec![Arc::new(ScalarFunctionExpr::new(
function_name,
udf,
args,
Arc::new(Field::new("round_robin_sort_hash", DataType::Int32, false)),
)) as PhysicalExprRef]
} else {
input
.schema()
.fields()
.iter()
.enumerate()
.map(|(index, field)| Arc::new(Column::new(field.name(), index)) as PhysicalExprRef)
.collect::<Vec<_>>()
};
Ok(create_default_ascending_sort_exec(
input,
&sort_exprs,
None,
false, // do not record output metric
))
}

/// The shuffle writer operator maps each input partition to M output partitions
/// based on a partitioning scheme. No guarantees are made about the order of
/// the resulting partitions.
Expand Down Expand Up @@ -137,21 +202,7 @@ impl ExecutionPlan for ShuffleWriterExec {
partitioner
}
Partitioning::RoundRobinPartitioning(..) => {
input = create_default_ascending_sort_exec(
input,
self.input
.schema()
.fields()
.iter()
.enumerate()
.map(|(index, field)| {
Arc::new(Column::new(&field.name(), index)) as PhysicalExprRef
})
.collect::<Vec<_>>()
.as_ref(),
None,
false, // do not record output metric
);
input = create_round_robin_sort_exec(input, partition)?;
let partitioner = Arc::new(SortShuffleRepartitioner::new(
exec_ctx.clone(),
self.output_data_file.clone(),
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -20,6 +20,7 @@ import org.apache.spark.sql.{AuronQueryTest, Row}
import org.apache.spark.sql.auron.join.JoinBuildSides.{JoinBuildLeft, JoinBuildRight}
import org.apache.spark.sql.execution.auron.plan.NativeFilterBase
import org.apache.spark.sql.execution.auron.plan.NativeShuffledHashJoinBase
import org.apache.spark.sql.execution.auron.plan.NativeShuffleExchangeExec
import org.apache.spark.sql.execution.auron.plan.NativeSortMergeJoinBase
import org.apache.spark.sql.execution.joins.auron.plan.NativeBroadcastJoinExec

Expand Down Expand Up @@ -119,7 +120,11 @@ class AuronQuerySuite extends AuronQueryTest with BaseAuronSQLSuite with AuronSQ
test("repartition over MapType") {
withTable("t_map") {
sql("create table t_map using parquet as select map('a', '1', 'b', '2') as data_map")
checkSparkAnswerAndOperator("SELECT /*+ repartition(10) */ data_map FROM t_map")
val df =
checkSparkAnswerAndOperator("SELECT /*+ repartition(10) */ data_map FROM t_map")
assert(collectFirst(df.queryExecution.executedPlan) { case e: NativeShuffleExchangeExec =>
e
}.isDefined)
}
}

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -128,7 +128,7 @@ object NativeConverters extends Logging {
}

def roundRobinTypeSupported(dataType: DataType): Boolean = dataType match {
case MapType(_, _, _) => false
case MapType(_, _, _) => true
case ArrayType(elementType, _) => roundRobinTypeSupported(elementType)
case StructType(fields) => fields.forall(f => roundRobinTypeSupported(f.dataType))
case _ => true
Expand Down
Loading