use std::collections::VecDeque;
use std::fmt;
use std::sync::{Arc, Mutex, MutexGuard};
use datafusion::common::config::ConfigOptions;
use datafusion::common::JoinType;
use datafusion::common::stats::Precision;
use datafusion::common::tree_node::{Transformed, TreeNode};
use datafusion::common::Result;
use datafusion::physical_optimizer::sanity_checker::SanityCheckPlan;
use datafusion::physical_optimizer::PhysicalOptimizerRule;
use datafusion::physical_plan::aggregates::{AggregateExec, AggregateMode};
use datafusion::physical_plan::coalesce_partitions::CoalescePartitionsExec;
use datafusion::physical_expr::expressions::Column;
use datafusion::physical_plan::filter::FilterExec;
use datafusion::physical_plan::joins::HashJoinExec;
use datafusion::physical_plan::projection::ProjectionExec;
use datafusion::physical_plan::repartition::RepartitionExec;
use datafusion::physical_plan::sorts::sort::SortExec;
use datafusion::physical_plan::sorts::sort_preserving_merge::SortPreservingMergeExec;
use datafusion::physical_plan::{
displayable, Distribution, ExecutionPlan, ExecutionPlanProperties, Partitioning, StatisticsArgs,
StatisticsContext,
};
use crate::exec::{MetalExec, MetalOp};
use crate::translate;
#[derive(Debug, Clone)]
#[non_exhaustive]
pub struct ArrowMetalConfig {
pub min_rows: usize,
pub accept_inexact: bool,
pub take_when_unknown: bool,
pub sort: bool,
pub topk: bool,
pub aggregate: bool,
pub filter: bool,
pub aggregate_choice: AggregateChoice,
pub table_rows: Option<usize>,
pub join: bool,
pub join_choice: JoinChoice,
pub report_plans: usize,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum AggregateChoice {
#[default]
Measured,
ArrowMetal,
DataFusion,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
#[non_exhaustive]
pub enum JoinChoice {
#[default]
Measured,
ArrowMetal,
}
impl Default for ArrowMetalConfig {
fn default() -> Self {
Self {
min_rows: 250_000,
accept_inexact: false,
take_when_unknown: false,
sort: true,
topk: false,
aggregate: true,
filter: false,
aggregate_choice: AggregateChoice::Measured,
table_rows: None,
join: true,
join_choice: JoinChoice::Measured,
report_plans: 64,
}
}
}
impl ArrowMetalConfig {
pub fn all() -> Self {
Self { topk: true, aggregate: true, filter: true, ..Self::default() }
}
pub fn with_min_rows(mut self, rows: usize) -> Self {
self.min_rows = rows;
self
}
pub fn with_accept_inexact(mut self, on: bool) -> Self {
self.accept_inexact = on;
self
}
pub fn with_take_when_unknown(mut self, on: bool) -> Self {
self.take_when_unknown = on;
self
}
pub fn with_sort(mut self, on: bool) -> Self {
self.sort = on;
self
}
pub fn with_topk(mut self, on: bool) -> Self {
self.topk = on;
self
}
pub fn with_aggregate(mut self, on: bool) -> Self {
self.aggregate = on;
self
}
pub fn with_filter(mut self, on: bool) -> Self {
self.filter = on;
self
}
pub fn with_aggregate_choice(mut self, choice: AggregateChoice) -> Self {
self.aggregate_choice = choice;
self
}
pub fn with_table_rows(mut self, rows: Option<usize>) -> Self {
self.table_rows = rows;
self
}
pub fn with_join(mut self, on: bool) -> Self {
self.join = on;
self
}
pub fn with_join_choice(mut self, choice: JoinChoice) -> Self {
self.join_choice = choice;
self
}
pub fn with_report_plans(mut self, plans: usize) -> Self {
self.report_plans = plans;
self
}
}
#[derive(Debug, Clone, PartialEq)]
#[non_exhaustive]
pub struct Decision {
pub node: String,
pub taken: bool,
pub reason: String,
pub runtime_fallback: bool,
pub groups: Option<GroupChoice>,
}
#[derive(Debug, Clone, PartialEq)]
#[non_exhaustive]
pub struct GroupChoice {
pub rows: usize,
pub estimate: Option<crate::probe::GroupEstimate>,
}
impl Decision {
pub(crate) fn runtime_fallback(op: &MetalOp, msg: &str) -> Self {
let reason = if msg.starts_with(crate::gpu::DATA_DEPENDENT) {
format!("{msg}; ran the DataFusion plan instead")
} else {
format!("ArrowMetal error at run time, ran the DataFusion plan instead: {msg}")
};
Decision { node: format!("MetalExec {op:?}"), taken: false, reason, runtime_fallback: true, groups: None }
}
pub(crate) fn memory_hand_back(op: &MetalOp, what: &str, err: &str) -> Self {
Decision {
node: format!("MetalExec {op:?}"),
taken: false,
reason: format!("the memory pool refused the reservation for {what}, ran the DataFusion plan instead: {err}"),
runtime_fallback: true,
groups: None,
}
}
pub(crate) fn runtime_choice(
op: &MetalOp,
on_arrowmetal: bool,
reason: String,
rows: usize,
estimate: Option<crate::probe::GroupEstimate>,
) -> Self {
let reason = format!("{}: {reason}", if on_arrowmetal { "ran on ArrowMetal" } else { "handed back to DataFusion" });
Decision {
node: format!("MetalExec {op:?}"),
taken: on_arrowmetal,
reason,
runtime_fallback: false,
groups: Some(GroupChoice { rows, estimate }),
}
}
pub fn is_runtime_choice(&self) -> bool {
self.groups.is_some()
}
pub fn is_data_dependent(&self) -> bool {
self.runtime_fallback && self.reason.starts_with(crate::gpu::DATA_DEPENDENT)
}
}
impl fmt::Display for Decision {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
let tag = if self.runtime_fallback {
"FALLBACK"
} else if self.groups.is_some() {
if self.taken { "GPU" } else { "HANDBACK" }
} else if self.taken {
"TAKEN"
} else {
"LEFT"
};
write!(f, "{tag:8} {} -- {}", self.node, self.reason)
}
}
#[derive(Debug, Clone, Default)]
#[non_exhaustive]
pub struct Report(Vec<Decision>);
impl Report {
pub fn decisions(&self) -> &[Decision] {
&self.0
}
pub fn taken(&self) -> impl Iterator<Item = &Decision> {
self.0.iter().filter(|d| d.taken && d.groups.is_none())
}
pub fn left(&self) -> impl Iterator<Item = &Decision> {
self.0.iter().filter(|d| !d.taken && !d.runtime_fallback && d.groups.is_none())
}
pub fn runtime_choices(&self) -> impl Iterator<Item = &Decision> {
self.0.iter().filter(|d| d.groups.is_some())
}
pub fn runtime_fallbacks(&self) -> impl Iterator<Item = &Decision> {
self.0.iter().filter(|d| d.runtime_fallback)
}
}
impl fmt::Display for Report {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
for d in &self.0 {
writeln!(f, "{d}")?;
}
Ok(())
}
}
#[derive(Debug)]
pub(crate) struct Log {
entries: VecDeque<(u64, Decision)>,
plan: u64,
keep: u64,
}
impl Log {
fn new(keep: usize) -> Self {
Self { entries: VecDeque::new(), plan: 0, keep: keep.max(1) as u64 }
}
fn begin_plan(&mut self) -> u64 {
self.plan += 1;
let first = self.plan.saturating_sub(self.keep - 1);
while self.entries.front().is_some_and(|(p, _)| *p < first) {
self.entries.pop_front();
}
self.plan
}
pub(crate) fn push(&mut self, plan: u64, d: Decision) {
if plan + self.keep > self.plan {
self.entries.push_back((plan, d));
}
}
}
pub(crate) type SharedLog = Arc<Mutex<Log>>;
pub(crate) fn lock(log: &SharedLog) -> MutexGuard<'_, Log> {
log.lock().unwrap_or_else(|p| p.into_inner())
}
#[derive(Debug, Clone)]
pub struct ArrowMetalRule {
config: ArrowMetalConfig,
log: SharedLog,
}
struct Candidate {
op: std::result::Result<MetalOp, String>,
input: Arc<dyn ExecutionPlan>,
right: Option<Arc<dyn ExecutionPlan>>,
original: Option<Arc<dyn ExecutionPlan>>,
reparent: Option<Arc<dyn ExecutionPlan>>,
}
impl Candidate {
fn new(op: std::result::Result<MetalOp, String>, input: Arc<dyn ExecutionPlan>) -> Self {
Self { op, input, right: None, original: None, reparent: None }
}
}
fn below_exchanges(p: &Arc<dyn ExecutionPlan>) -> Arc<dyn ExecutionPlan> {
let mut cur = Arc::clone(p);
loop {
let next = if let Some(r) = cur.downcast_ref::<RepartitionExec>() {
Arc::clone(r.input())
} else if let Some(c) = cur.downcast_ref::<CoalescePartitionsExec>() {
if c.fetch().is_some() {
return cur;
}
Arc::clone(c.input())
} else {
return cur;
};
cur = next;
}
}
fn map_through_projection(
spm_expr: &datafusion::physical_expr::LexOrdering,
proj: &ProjectionExec,
) -> Option<Vec<(usize, arrow::compute::SortOptions)>> {
let mut out = Vec::new();
for s in spm_expr.iter() {
let c = s.expr.downcast_ref::<Column>()?;
let pe = proj.expr().get(c.index())?;
let inner = pe.expr.downcast_ref::<Column>()?;
out.push((inner.index(), s.options));
}
Some(out)
}
fn sort_keys(expr: &datafusion::physical_expr::LexOrdering) -> Option<Vec<(usize, arrow::compute::SortOptions)>> {
expr.iter().map(|s| s.expr.downcast_ref::<Column>().map(|c| (c.index(), s.options))).collect()
}
impl ArrowMetalRule {
pub fn new(config: ArrowMetalConfig) -> Self {
let keep = config.report_plans;
Self { config, log: Arc::new(Mutex::new(Log::new(keep))) }
}
pub fn config(&self) -> &ArrowMetalConfig {
&self.config
}
pub fn report(&self) -> Report {
Report(lock(&self.log).entries.iter().map(|(_, d)| d.clone()).collect())
}
pub fn clear_report(&self) {
lock(&self.log).entries.clear();
}
fn record(&self, plan: u64, node: &Arc<dyn ExecutionPlan>, taken: bool, reason: String) {
let node = displayable(node.as_ref()).one_line().to_string().trim_end().to_string();
lock(&self.log).push(plan, Decision { node, taken, reason, runtime_fallback: false, groups: None });
}
fn sort_disabled(&self, fetch: Option<usize>) -> Option<String> {
match fetch {
None if !self.config.sort => Some("sort disabled in config".into()),
Some(n) if !self.config.topk => Some(format!("top-k (sort with fetch {n}) disabled in config")),
_ => None,
}
}
fn candidate(&self, node: &Arc<dyn ExecutionPlan>) -> Option<Candidate> {
if let Some(spm) = node.downcast_ref::<SortPreservingMergeExec>() {
let (sort_node, proj) = if spm.input().downcast_ref::<SortExec>().is_some() {
(Arc::clone(spm.input()), None)
} else {
let p = spm.input().downcast_ref::<ProjectionExec>()?;
p.input().downcast_ref::<SortExec>()?;
(Arc::clone(p.input()), Some(Arc::clone(spm.input())))
};
let sort = sort_node.downcast_ref::<SortExec>()?;
let fetch = match (spm.fetch(), sort.fetch()) {
(Some(a), Some(b)) => Some(a.min(b)),
(a, b) => a.or(b),
};
if let Some(why) = self.sort_disabled(fetch) {
return Some(Candidate::new(Err(why), Arc::clone(sort.input())));
}
let same_order = match &proj {
None => sort.expr() == spm.expr(),
Some(p) => {
let p = p.downcast_ref::<ProjectionExec>()?;
let mapped = map_through_projection(spm.expr(), p);
mapped.is_some() && mapped == sort_keys(sort.expr())
}
};
if !same_order {
return Some(Candidate::new(
Err("merge ordering differs from the sort's".into()),
Arc::clone(sort.input()),
));
}
let input = Arc::clone(sort.input());
let op = translate::sort_op(sort.expr(), &input.schema(), fetch);
let mut c = Candidate::new(op, input);
if proj.is_some() {
c.original = Some(Arc::new(
SortPreservingMergeExec::new(sort.expr().clone(), Arc::clone(&sort_node)).with_fetch(fetch),
));
c.reparent = proj;
}
return Some(c);
}
if let Some(sort) = node.downcast_ref::<SortExec>() {
let input = Arc::clone(sort.input());
if let Some(why) = self.sort_disabled(sort.fetch()) {
return Some(Candidate::new(Err(why), input));
}
if sort.preserve_partitioning() && input.output_partitioning().partition_count() > 1 {
return Some(Candidate::new(Err("per-partition sort (preserve_partitioning) with no replaced merge above it".into()), input));
}
let op = translate::sort_op(sort.expr(), &input.schema(), sort.fetch());
return Some(Candidate::new(op, input));
}
if let Some(agg) = node.downcast_ref::<AggregateExec>() {
if !self.config.aggregate {
return Some(Candidate::new(Err("aggregate disabled in config".into()), Arc::clone(agg.input())));
}
if node.output_ordering().is_some() {
return Some(Candidate::new(
Err("the aggregate's output carries an ordering (sorted input), which the GPU group-by does not keep".into()),
Arc::clone(agg.input()),
));
}
return Some(match agg.mode() {
AggregateMode::Single | AggregateMode::SinglePartitioned => {
Candidate::new(translate::aggregate_op(agg), Arc::clone(agg.input()))
}
AggregateMode::Final | AggregateMode::FinalPartitioned => {
let mut cur = Arc::clone(agg.input());
loop {
if cur.downcast_ref::<RepartitionExec>().is_some()
|| cur.downcast_ref::<CoalescePartitionsExec>().is_some()
{
let next = Arc::clone(cur.children()[0]);
cur = next;
continue;
}
break;
}
match cur.downcast_ref::<AggregateExec>() {
Some(p) if *p.mode() == AggregateMode::Partial && translate::same_aggregates(agg, p) => {
Candidate::new(translate::aggregate_op(p), Arc::clone(p.input()))
}
_ => Candidate::new(Err("Final aggregate without a matching Partial below its exchange".into()), Arc::clone(agg.input())),
}
}
AggregateMode::Partial => Candidate::new(Err("Partial aggregate whose Final was not replaced".into()), Arc::clone(agg.input())),
AggregateMode::PartialReduce => Candidate::new(Err("PartialReduce aggregate".into()), Arc::clone(agg.input())),
});
}
if let Some(f) = node.downcast_ref::<FilterExec>() {
let input = Arc::clone(f.input());
if !self.config.filter {
return Some(Candidate::new(Err("filter disabled in config".into()), input));
}
if node.fetch().is_some() {
return Some(Candidate::new(Err("filter with a fetch limit".into()), input));
}
if node.output_ordering().is_some() && input.output_partitioning().partition_count() > 1 {
return Some(Candidate::new(Err("order-preserving filter over several partitions".into()), input));
}
let projection = f.projection().as_ref().map(|p| p.iter().copied().collect::<Vec<usize>>());
let op = translate::filter_op(f.predicate(), &input.schema(), projection);
return Some(Candidate::new(op, input));
}
if let Some(j) = node.downcast_ref::<HashJoinExec>() {
let (left, right) = (below_exchanges(j.left()), below_exchanges(j.right()));
let ordered_ok = matches!(j.join_type(), JoinType::Inner | JoinType::Right)
&& Arc::ptr_eq(&right, j.right())
&& j.right().output_partitioning().partition_count() == 1;
let op = if !self.config.join {
Err("join disabled in config".into())
} else if node.output_ordering().is_some() && !ordered_ok {
Err("the join's output carries an ordering of a probe side split over partitions".into())
} else if Arc::ptr_eq(&left, &right) {
Err("both join inputs are one plan node".into())
} else {
translate::join_op(j)
};
let mut c = Candidate::new(op, left);
c.right = Some(right);
return Some(c);
}
None
}
fn statistics_rows(input: &Arc<dyn ExecutionPlan>) -> Precision<usize> {
match StatisticsContext::new().compute(input.as_ref(), &StatisticsArgs::new()) {
Ok(s) => s.num_rows,
Err(_) => Precision::Absent,
}
}
fn row_count(&self, input: &Arc<dyn ExecutionPlan>) -> (Option<usize>, String) {
match Self::statistics_rows(input) {
Precision::Exact(n) => (Some(n), format!("{n} (exact)")),
Precision::Inexact(n) if self.config.accept_inexact => (Some(n), format!("~{n} (inexact, accepted)")),
Precision::Inexact(n) => (None, format!("~{n} (an estimate; accept_inexact is off)")),
Precision::Absent => (None, "unknown".into()),
}
}
fn size_ok(&self, input: &Arc<dyn ExecutionPlan>) -> (bool, String, Option<usize>) {
let rows = Self::statistics_rows(input);
let min = self.config.min_rows;
match rows {
Precision::Exact(n) => (n >= min, format!("input rows {n} (exact) vs min_rows {min}"), Some(n)),
Precision::Inexact(n) if self.config.accept_inexact => {
(n >= min, format!("input rows ~{n} (inexact, accepted) vs min_rows {min}"), Some(n))
}
Precision::Inexact(n) => (false, format!("input rows ~{n} are an estimate (accept_inexact is off)"), None),
Precision::Absent => (
self.config.take_when_unknown,
format!("input row count unknown (take_when_unknown = {})", self.config.take_when_unknown),
None,
),
}
}
fn visit(&self, plan: u64, node: Arc<dyn ExecutionPlan>, top: usize) -> Result<Transformed<Arc<dyn ExecutionPlan>>> {
if node.downcast_ref::<MetalExec>().is_some() {
return Ok(Transformed::no(node));
}
let Some(c) = self.candidate(&node) else {
return Ok(Transformed::no(node));
};
let op = match c.op {
Ok(op) => op,
Err(why) => {
self.record(plan, &node, false, why);
return Ok(Transformed::no(node));
}
};
let at_top = Arc::as_ptr(&node) as *const () as usize == top;
let wrap = match node.output_partitioning() {
p if p.partition_count() <= 1 => None,
_ if at_top => None,
Partitioning::Hash(exprs, n) => Some(Partitioning::Hash(exprs.clone(), *n)),
Partitioning::RoundRobinBatch(n) | Partitioning::UnknownPartitioning(n) => {
Some(Partitioning::RoundRobinBatch(*n))
}
Partitioning::Range(_) => {
self.record(plan, &node, false, "range-partitioned output".into());
return Ok(Transformed::no(node));
}
};
let mut input = Arc::clone(&c.input);
while let Some(r) = input.downcast_ref::<RepartitionExec>() {
if !matches!(r.partitioning(), Partitioning::RoundRobinBatch(_)) {
break;
}
let next = Arc::clone(r.input());
input = next;
}
let mut inputs = vec![Arc::clone(&input)];
let (ok, mut size, rows) = match &c.right {
None => self.size_ok(&input),
Some(right) => {
inputs.push(Arc::clone(right));
let (lrows, ltext) = self.row_count(&input);
let (rrows, rtext) = self.row_count(right);
let rows_text = format!("left (build) rows {ltext}, right (probe) rows {rtext}");
match (lrows, rrows) {
(Some(l), Some(r)) => {
let min = self.config.min_rows;
let n = l.max(r);
let mut why = format!("{rows_text}; the larger vs min_rows {min}");
let mut ok = n >= min;
if ok && self.config.join_choice == JoinChoice::Measured {
match crate::choice::join_takes(&op, input.as_ref(), right.as_ref(), l as u64, r as u64) {
Ok(w) => why = format!("{why}; {w}"),
Err(w) => {
why = format!("{why}; {w}");
ok = false;
}
}
}
(ok, why, Some(l + r))
}
_ => (false, format!("{rows_text}: a join needs both row counts"), None),
}
}
};
if !ok {
self.record(plan, &node, false, size);
return Ok(Transformed::no(node));
}
if matches!(op, MetalOp::Aggregate { .. }) && self.config.aggregate_choice == AggregateChoice::Measured {
let Some(shape) = crate::choice::shape(&op, input.as_ref()) else {
self.record(plan, &node, false, "aggregate without a shape".into());
return Ok(Transformed::no(node));
};
if let Some(n) = self.config.table_rows.or(rows) {
match crate::choice::any_bucket(&shape, n as u64) {
Ok(why) => size = format!("{size}; {why}"),
Err(why) => {
self.record(plan, &node, false, format!("{size}; {why}"));
return Ok(Transformed::no(node));
}
}
}
}
let original = c.original.unwrap_or_else(|| Arc::clone(&node));
let mut wrap = wrap;
let mut out_parts = 1;
if c.right.is_some() && !at_top {
out_parts = node.output_partitioning().partition_count();
if !matches!(wrap, Some(Partitioning::Hash(..))) {
wrap = None;
}
}
let metal: Arc<dyn ExecutionPlan> = Arc::new(
MetalExec::new(
op,
inputs,
original,
(Arc::clone(&self.log), plan),
crate::exec::AggSettings {
choice: self.config.aggregate_choice,
table_rows: self.config.table_rows,
rows_hint: rows,
},
)
.with_output_partitions(out_parts),
);
let mut reason = size;
if out_parts > 1 {
reason.push_str(&format!("; output in {out_parts} partitions"));
}
if at_top && node.output_partitioning().partition_count() > 1 {
reason.push_str("; output kept at one partition (only projections above it)");
}
let metal = match c.reparent {
#[allow(deprecated)] Some(p) => {
reason.push_str("; projection kept above it");
p.with_new_children(vec![metal])?
}
None => metal,
};
let out: Arc<dyn ExecutionPlan> = match wrap {
None => metal,
Some(p) => {
if c.right.is_some() && matches!(p, Partitioning::Hash(..)) {
reason.push_str(&format!(
"; output re-partitioned to {p} where the parent requires it (to keep the parent's distribution)"
));
} else {
reason.push_str(&format!("; output re-partitioned to {p} to keep the parent's distribution"));
}
Arc::new(RepartitionExec::try_new(metal, p)?)
}
};
self.record(plan, &node, true, reason);
Ok(Transformed::yes(out))
}
}
#[allow(deprecated)] fn drop_hash_where_not_required(node: Arc<dyn ExecutionPlan>) -> Result<Transformed<Arc<dyn ExecutionPlan>>> {
let required = node.required_input_distribution();
let children: Vec<Arc<dyn ExecutionPlan>> = node.children().into_iter().cloned().collect();
let mut changed = false;
let mut new_children = Vec::with_capacity(children.len());
for (i, c) in children.into_iter().enumerate() {
let swap = match c.downcast_ref::<RepartitionExec>() {
Some(r) => {
let over_join = r
.input()
.downcast_ref::<MetalExec>()
.is_some_and(|m| matches!(m.op(), MetalOp::Join { .. }));
match r.partitioning() {
Partitioning::Hash(..)
if over_join && matches!(required.get(i), Some(Distribution::UnspecifiedDistribution)) =>
{
Some(Arc::clone(r.input()))
}
_ => None,
}
}
None => None,
};
match swap {
Some(s) => {
changed = true;
new_children.push(s);
}
None => new_children.push(c),
}
}
if !changed {
return Ok(Transformed::no(node));
}
Ok(Transformed::yes(node.with_new_children(new_children)?))
}
impl PhysicalOptimizerRule for ArrowMetalRule {
fn optimize(&self, plan: Arc<dyn ExecutionPlan>, config: &ConfigOptions) -> Result<Arc<dyn ExecutionPlan>> {
let mut top = Arc::clone(&plan);
while let Some(p) = top.downcast_ref::<ProjectionExec>() {
let next = Arc::clone(p.input());
top = next;
}
let top = Arc::as_ptr(&top) as *const () as usize;
let id = lock(&self.log).begin_plan();
let plan = plan.transform_down(|n| self.visit(id, n, top))?.data;
let relaxed = Arc::clone(&plan).transform_up(drop_hash_where_not_required)?;
if relaxed.transformed && SanityCheckPlan::new().optimize(Arc::clone(&relaxed.data), config).is_ok() {
return Ok(relaxed.data);
}
Ok(plan)
}
fn name(&self) -> &str {
"ArrowMetalRule"
}
fn schema_check(&self) -> bool {
true
}
}