Skip to main content

datafusion_arrowmetal/
rule.rs

1//! `ArrowMetalRule`: the `PhysicalOptimizerRule` that swaps supported nodes for `MetalExec`.
2
3use std::collections::VecDeque;
4use std::fmt;
5use std::sync::{Arc, Mutex, MutexGuard};
6
7use datafusion::common::config::ConfigOptions;
8use datafusion::common::JoinType;
9use datafusion::common::stats::Precision;
10use datafusion::common::tree_node::{Transformed, TreeNode};
11use datafusion::common::Result;
12use datafusion::physical_optimizer::sanity_checker::SanityCheckPlan;
13use datafusion::physical_optimizer::PhysicalOptimizerRule;
14use datafusion::physical_plan::aggregates::{AggregateExec, AggregateMode};
15use datafusion::physical_plan::coalesce_partitions::CoalescePartitionsExec;
16use datafusion::physical_expr::expressions::Column;
17use datafusion::physical_plan::filter::FilterExec;
18use datafusion::physical_plan::joins::HashJoinExec;
19use datafusion::physical_plan::projection::ProjectionExec;
20use datafusion::physical_plan::repartition::RepartitionExec;
21use datafusion::physical_plan::sorts::sort::SortExec;
22use datafusion::physical_plan::sorts::sort_preserving_merge::SortPreservingMergeExec;
23use datafusion::physical_plan::{
24    displayable, Distribution, ExecutionPlan, ExecutionPlanProperties, Partitioning, StatisticsArgs,
25    StatisticsContext,
26};
27
28use crate::exec::{MetalExec, MetalOp};
29use crate::translate;
30
31/// When the rule takes a node.
32///
33/// [`Default`] is the measured take-list; [`ArrowMetalConfig::all`] takes every shape the rule can
34/// translate (what the differential grid and the benchmark use).
35///
36/// The fields are public to read and to set on a value (`let mut c = ArrowMetalConfig::all();
37/// c.min_rows = 0;`); outside this crate a config is built from [`Default`] or
38/// [`ArrowMetalConfig::all`] and the `with_*` methods.
39#[derive(Debug, Clone)]
40#[non_exhaustive]
41pub struct ArrowMetalConfig {
42    /// Take a node only when its input has at least this many rows.
43    pub min_rows: usize,
44    /// Use an inexact (estimated) row count as if it were exact. Default: no.
45    pub accept_inexact: bool,
46    /// Take a node whose input row count is unknown. Default: no (leave it).
47    pub take_when_unknown: bool,
48    /// Full sorts: `ORDER BY` without `LIMIT`.
49    pub sort: bool,
50    /// Top-k: `ORDER BY ... LIMIT` (a sort with a fetch).
51    pub topk: bool,
52    /// Aggregates (GROUP BY, DISTINCT). Who runs a replaced one is `aggregate_choice`.
53    pub aggregate: bool,
54    /// Filters: a `FilterExec` whose predicate translates.
55    pub filter: bool,
56    /// Who runs a replaced aggregate. Default: [`AggregateChoice::Measured`].
57    pub aggregate_choice: AggregateChoice,
58    /// Under [`AggregateChoice::Measured`], look the measured table up at this row count instead of
59    /// the input's, at plan time and at run time (tests and experiments; default `None`).
60    pub table_rows: Option<usize>,
61    /// Hash joins (`HashJoinExec`: inner, left, right on equal keys). Which ones is `join_choice`.
62    pub join: bool,
63    /// Which translatable joins are replaced. Default: [`JoinChoice::Measured`].
64    pub join_choice: JoinChoice,
65    /// The report keeps the decisions of the last this many plans the rule optimized (an EXPLAIN
66    /// and each execution plan count one each), with the run-time decisions of their nodes.
67    /// Default 64; 0 keeps one.
68    pub report_plans: usize,
69}
70
71/// Who runs an aggregate the rule replaced.
72///
73/// The rule sees the input's row count at plan time but not its group count, and the same
74/// aggregate SQL is faster on ArrowMetal at some group counts and slower at others. So a replaced
75/// aggregate is decided when it runs: `MetalExec` collects its input, and under `Measured` it
76/// estimates the group count from a sample of the key columns on the CPU and looks the shape up in
77/// the measured table (`src/agg_table.rs`). It runs on ArrowMetal only where the sweep measured it
78/// ahead; otherwise it hands the node back to DataFusion's own operators over the same batches.
79#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
80pub enum AggregateChoice {
81    /// Estimate the group count and look it up in the measured table (the default). At plan time
82    /// the rule replaces an aggregate only when the table takes its shape at some group count at
83    /// the input's exact row count.
84    #[default]
85    Measured,
86    /// Always run a replaced aggregate on ArrowMetal (the benchmark's crossover sweep, tests).
87    ArrowMetal,
88    /// Always hand a replaced aggregate back to DataFusion (measures the hand-back itself, tests).
89    DataFusion,
90}
91
92/// Which hash joins the rule replaces (decided at plan time, from both inputs' exact row counts).
93#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
94#[non_exhaustive]
95pub enum JoinChoice {
96    /// The joins the measured table (`src/join_table.rs`) takes: by join type, key type, the
97    /// build (left) input's rows and the probe (right) input's rows (the default).
98    #[default]
99    Measured,
100    /// Every join the rule can translate, from `min_rows` (the benchmark's sweep, tests).
101    ArrowMetal,
102}
103
104/// The default take-list, from the rule on/off measurement against DataFusion 55.1 on an M4 Max
105/// (16 partitions; MemTables of 8192-row batches and of one batch per partition; this crate with
106/// one totalOrder key per ORDER BY key and the chunked import;
107/// `datafusion/results/datafusion_rule_2026-09-29.csv`, rows `crate = now`; DataFusion alone /
108/// DataFusion with the rule, wall, best of 5, both layouts):
109///
110/// | shape | 1M | 10M | 50M | default |
111/// |---|---:|---:|---:|---|
112/// | ORDER BY int64 / Float64 / String / Float32 key, 3 columns | 5.9x - 12.8x | 19.2x - 28.7x | 18.9x - 27.2x | taken |
113/// | the same over DataFusion's Parquet scan (snappy, zstd) | | 13.6x, 10.8x | 16.5x, 13.8x | taken |
114/// | ORDER BY ... LIMIT 100 (int64, Float64 DESC, Float32 DESC) | 0.14x - 0.49x | 0.24x - 0.41x | 0.17x - 0.29x | left |
115/// | WHERE + whole-table sum/count (the filter is what is taken) | 0.56x - 0.59x | 0.62x - 0.72x | 0.74x | left |
116/// | Parquet WHERE + GROUP BY | | 0.52x - 0.53x | 0.54x - 0.55x | left |
117///
118/// Full sorts at smaller inputs (every key type, both layouts): 100k rows 1.34x - 2.63x, 250k
119/// 2.64x - 4.80x, 500k 4.32x - 7.49x; the 1x crossover is below 100k for every key, and 250k is
120/// the first measured size at which every key type is at or above 2.6x; hence `min_rows` 250,000.
121///
122/// Top-k is left: DataFusion's TopK answers LIMIT 100 over 50M rows in 4.6 - 6.5 ms; the rule's
123/// GPU top-k takes 20 - 29 ms there, of which collecting the input stream alone is 6.5 - 8 ms.
124///
125/// Aggregates (GROUP BY, DISTINCT) are decided per shape and group count
126/// ([`AggregateChoice::Measured`]): the same SQL is ahead of DataFusion at some group counts and
127/// behind at others, and the rule sees the row count but not the group count (DataFusion's
128/// MemTable and Parquet statistics carry no distinct counts). The rule replaces an aggregate when
129/// the measured table (`src/agg_table.rs`, generated by `scripts/groupby_table.py` from the sweep
130/// CSV it names) takes its shape at some group count at the input's exact row count; the
131/// `MetalExec` then estimates the group count from a sample of the keys and runs on ArrowMetal only
132/// at a bucket the table takes, handing the node back to DataFusion otherwise. Hash joins are
133/// replaced where the measured join table (`src/join_table.rs`, generated by
134/// `scripts/join_table.py` from the join CSV it names) takes their join type, key type and build
135/// and probe row counts ([`JoinChoice::Measured`]).
136impl Default for ArrowMetalConfig {
137    fn default() -> Self {
138        Self {
139            min_rows: 250_000,
140            accept_inexact: false,
141            take_when_unknown: false,
142            sort: true,
143            topk: false,
144            aggregate: true,
145            filter: false,
146            aggregate_choice: AggregateChoice::Measured,
147            table_rows: None,
148            join: true,
149            join_choice: JoinChoice::Measured,
150            report_plans: 64,
151        }
152    }
153}
154
155impl ArrowMetalConfig {
156    /// Every shape the rule can translate (sorts, top-k, aggregates, filters, joins), at the
157    /// default `min_rows`. Aggregates and joins are still decided by their measured tables unless
158    /// [`aggregate_choice`](Self::aggregate_choice) and [`join_choice`](Self::join_choice) say
159    /// otherwise.
160    pub fn all() -> Self {
161        Self { topk: true, aggregate: true, filter: true, ..Self::default() }
162    }
163
164    /// Sets [`min_rows`](Self::min_rows).
165    pub fn with_min_rows(mut self, rows: usize) -> Self {
166        self.min_rows = rows;
167        self
168    }
169    /// Sets [`accept_inexact`](Self::accept_inexact).
170    pub fn with_accept_inexact(mut self, on: bool) -> Self {
171        self.accept_inexact = on;
172        self
173    }
174    /// Sets [`take_when_unknown`](Self::take_when_unknown).
175    pub fn with_take_when_unknown(mut self, on: bool) -> Self {
176        self.take_when_unknown = on;
177        self
178    }
179    /// Sets [`sort`](Self::sort).
180    pub fn with_sort(mut self, on: bool) -> Self {
181        self.sort = on;
182        self
183    }
184    /// Sets [`topk`](Self::topk).
185    pub fn with_topk(mut self, on: bool) -> Self {
186        self.topk = on;
187        self
188    }
189    /// Sets [`aggregate`](Self::aggregate).
190    pub fn with_aggregate(mut self, on: bool) -> Self {
191        self.aggregate = on;
192        self
193    }
194    /// Sets [`filter`](Self::filter).
195    pub fn with_filter(mut self, on: bool) -> Self {
196        self.filter = on;
197        self
198    }
199    /// Sets [`aggregate_choice`](Self::aggregate_choice).
200    pub fn with_aggregate_choice(mut self, choice: AggregateChoice) -> Self {
201        self.aggregate_choice = choice;
202        self
203    }
204    /// Sets [`table_rows`](Self::table_rows).
205    pub fn with_table_rows(mut self, rows: Option<usize>) -> Self {
206        self.table_rows = rows;
207        self
208    }
209    /// Sets [`join`](Self::join).
210    pub fn with_join(mut self, on: bool) -> Self {
211        self.join = on;
212        self
213    }
214    /// Sets [`join_choice`](Self::join_choice).
215    pub fn with_join_choice(mut self, choice: JoinChoice) -> Self {
216        self.join_choice = choice;
217        self
218    }
219    /// Sets [`report_plans`](Self::report_plans).
220    pub fn with_report_plans(mut self, plans: usize) -> Self {
221        self.report_plans = plans;
222        self
223    }
224}
225
226/// One node the rule looked at, one replaced aggregate's run-time choice, or one runtime fallback.
227#[derive(Debug, Clone, PartialEq)]
228#[non_exhaustive]
229pub struct Decision {
230    /// The node, as DataFusion prints it on one line.
231    pub node: String,
232    /// At plan time: the node was replaced by a `MetalExec`. For a run-time choice: the aggregate
233    /// ran on ArrowMetal.
234    pub taken: bool,
235    /// Why it was taken or left.
236    pub reason: String,
237    /// True for a `MetalExec` that hit an ArrowMetal error at run time and ran DataFusion instead.
238    pub runtime_fallback: bool,
239    /// For a replaced aggregate's run-time choice: what it was decided from. `taken` is then
240    /// whether it ran on ArrowMetal.
241    pub groups: Option<GroupChoice>,
242}
243
244/// What a replaced aggregate's run-time choice was made from.
245#[derive(Debug, Clone, PartialEq)]
246#[non_exhaustive]
247pub struct GroupChoice {
248    /// Rows of the collected input.
249    pub rows: usize,
250    /// The group-count probe's answer (`None` when `aggregate_choice` forced the choice).
251    pub estimate: Option<crate::probe::GroupEstimate>,
252}
253
254impl Decision {
255    pub(crate) fn runtime_fallback(op: &MetalOp, msg: &str) -> Self {
256        let reason = if msg.starts_with(crate::gpu::DATA_DEPENDENT) {
257            format!("{msg}; ran the DataFusion plan instead")
258        } else {
259            format!("ArrowMetal error at run time, ran the DataFusion plan instead: {msg}")
260        };
261        Decision { node: format!("MetalExec {op:?}"), taken: false, reason, runtime_fallback: true, groups: None }
262    }
263
264    pub(crate) fn memory_hand_back(op: &MetalOp, what: &str, err: &str) -> Self {
265        Decision {
266            node: format!("MetalExec {op:?}"),
267            taken: false,
268            reason: format!("the memory pool refused the reservation for {what}, ran the DataFusion plan instead: {err}"),
269            runtime_fallback: true,
270            groups: None,
271        }
272    }
273
274    pub(crate) fn runtime_choice(
275        op: &MetalOp,
276        on_arrowmetal: bool,
277        reason: String,
278        rows: usize,
279        estimate: Option<crate::probe::GroupEstimate>,
280    ) -> Self {
281        let reason = format!("{}: {reason}", if on_arrowmetal { "ran on ArrowMetal" } else { "handed back to DataFusion" });
282        Decision {
283            node: format!("MetalExec {op:?}"),
284            taken: on_arrowmetal,
285            reason,
286            runtime_fallback: false,
287            groups: Some(GroupChoice { rows, estimate }),
288        }
289    }
290
291    /// A replaced aggregate's run-time choice (see [`AggregateChoice`]).
292    pub fn is_runtime_choice(&self) -> bool {
293        self.groups.is_some()
294    }
295
296    /// A runtime fallback (the `runtime_fallback` field) caused by the data, not an error.
297    pub fn is_data_dependent(&self) -> bool {
298        self.runtime_fallback && self.reason.starts_with(crate::gpu::DATA_DEPENDENT)
299    }
300}
301
302impl fmt::Display for Decision {
303    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
304        let tag = if self.runtime_fallback {
305            "FALLBACK"
306        } else if self.groups.is_some() {
307            if self.taken { "GPU" } else { "HANDBACK" }
308        } else if self.taken {
309            "TAKEN"
310        } else {
311            "LEFT"
312        };
313        write!(f, "{tag:8} {} -- {}", self.node, self.reason)
314    }
315}
316
317/// What the rule decided for the last `report_plans` plans since it was created or last cleared.
318#[derive(Debug, Clone, Default)]
319#[non_exhaustive]
320pub struct Report(Vec<Decision>);
321
322impl Report {
323    /// Every decision, oldest first.
324    pub fn decisions(&self) -> &[Decision] {
325        &self.0
326    }
327    /// Nodes the rule replaced at plan time.
328    pub fn taken(&self) -> impl Iterator<Item = &Decision> {
329        self.0.iter().filter(|d| d.taken && d.groups.is_none())
330    }
331    /// Nodes the rule left at plan time.
332    pub fn left(&self) -> impl Iterator<Item = &Decision> {
333        self.0.iter().filter(|d| !d.taken && !d.runtime_fallback && d.groups.is_none())
334    }
335    /// Replaced aggregates' run-time choices (ArrowMetal or handed back).
336    pub fn runtime_choices(&self) -> impl Iterator<Item = &Decision> {
337        self.0.iter().filter(|d| d.groups.is_some())
338    }
339    /// `MetalExec`s that ran DataFusion's plan instead after an ArrowMetal error or on data the
340    /// GPU path cannot answer exactly (see [`Decision::is_data_dependent`]).
341    pub fn runtime_fallbacks(&self) -> impl Iterator<Item = &Decision> {
342        self.0.iter().filter(|d| d.runtime_fallback)
343    }
344}
345
346impl fmt::Display for Report {
347    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
348        for d in &self.0 {
349            writeln!(f, "{d}")?;
350        }
351        Ok(())
352    }
353}
354
355/// The shared decision log: the decisions of the last `keep` plans, each tagged with its plan.
356#[derive(Debug)]
357pub(crate) struct Log {
358    entries: VecDeque<(u64, Decision)>,
359    plan: u64,
360    keep: u64,
361}
362
363impl Log {
364    fn new(keep: usize) -> Self {
365        Self { entries: VecDeque::new(), plan: 0, keep: keep.max(1) as u64 }
366    }
367
368    /// Starts a new plan: drops the decisions of plans outside the last `keep`.
369    fn begin_plan(&mut self) -> u64 {
370        self.plan += 1;
371        let first = self.plan.saturating_sub(self.keep - 1);
372        while self.entries.front().is_some_and(|(p, _)| *p < first) {
373            self.entries.pop_front();
374        }
375        self.plan
376    }
377
378    /// Records a decision of plan `plan` (dropped when that plan is no longer kept).
379    pub(crate) fn push(&mut self, plan: u64, d: Decision) {
380        if plan + self.keep > self.plan {
381            self.entries.push_back((plan, d));
382        }
383    }
384}
385
386pub(crate) type SharedLog = Arc<Mutex<Log>>;
387
388/// The log, also after a panic elsewhere poisoned its mutex (the log stays consistent: every
389/// update is a single push or pop).
390pub(crate) fn lock(log: &SharedLog) -> MutexGuard<'_, Log> {
391    log.lock().unwrap_or_else(|p| p.into_inner())
392}
393
394/// Replaces `SortExec` (+ its `SortPreservingMergeExec`), hash `AggregateExec` (a `Single` node, or
395/// a `Final`/`Partial` pair) and `FilterExec` with [`MetalExec`] when ArrowMetal gives the same
396/// answer and the input is big enough. Cloning shares the report.
397#[derive(Debug, Clone)]
398pub struct ArrowMetalRule {
399    config: ArrowMetalConfig,
400    log: SharedLog,
401}
402
403struct Candidate {
404    op: std::result::Result<MetalOp, String>,
405    /// The input `MetalExec` reads (the replaced chain's leaf input; a join's left input).
406    input: Arc<dyn ExecutionPlan>,
407    /// A join's right input.
408    right: Option<Arc<dyn ExecutionPlan>>,
409    /// The runtime fallback, when it is not the node itself (it must have `MetalExec`'s schema).
410    original: Option<Arc<dyn ExecutionPlan>>,
411    /// A node to put back above `MetalExec` (the projection of `SPM -> Projection -> SortExec`).
412    reparent: Option<Arc<dyn ExecutionPlan>>,
413}
414
415impl Candidate {
416    fn new(op: std::result::Result<MetalOp, String>, input: Arc<dyn ExecutionPlan>) -> Self {
417        Self { op, input, right: None, original: None, reparent: None }
418    }
419}
420
421/// `p` without the exchanges directly above its data: `RepartitionExec`s and
422/// `CoalescePartitionsExec`s. A join's `MetalExec` collects every partition of both inputs, so a
423/// re-partitioning below it only moves rows it reads anyway. (The fallback plan keeps them.)
424fn below_exchanges(p: &Arc<dyn ExecutionPlan>) -> Arc<dyn ExecutionPlan> {
425    let mut cur = Arc::clone(p);
426    loop {
427        let next = if let Some(r) = cur.downcast_ref::<RepartitionExec>() {
428            Arc::clone(r.input())
429        } else if let Some(c) = cur.downcast_ref::<CoalescePartitionsExec>() {
430            if c.fetch().is_some() {
431                return cur;
432            }
433            Arc::clone(c.input())
434        } else {
435            return cur;
436        };
437        cur = next;
438    }
439}
440
441/// `spm_expr` (written against the projection's output) mapped through `proj` onto the
442/// projection's input, or `None` when a merge key is not a plain column of the projection.
443fn map_through_projection(
444    spm_expr: &datafusion::physical_expr::LexOrdering,
445    proj: &ProjectionExec,
446) -> Option<Vec<(usize, arrow::compute::SortOptions)>> {
447    let mut out = Vec::new();
448    for s in spm_expr.iter() {
449        let c = s.expr.downcast_ref::<Column>()?;
450        let pe = proj.expr().get(c.index())?;
451        let inner = pe.expr.downcast_ref::<Column>()?;
452        out.push((inner.index(), s.options));
453    }
454    Some(out)
455}
456
457fn sort_keys(expr: &datafusion::physical_expr::LexOrdering) -> Option<Vec<(usize, arrow::compute::SortOptions)>> {
458    expr.iter().map(|s| s.expr.downcast_ref::<Column>().map(|c| (c.index(), s.options))).collect()
459}
460
461impl ArrowMetalRule {
462    /// A rule with this config and an empty report.
463    pub fn new(config: ArrowMetalConfig) -> Self {
464        let keep = config.report_plans;
465        Self { config, log: Arc::new(Mutex::new(Log::new(keep))) }
466    }
467
468    /// The config the rule was made with.
469    pub fn config(&self) -> &ArrowMetalConfig {
470        &self.config
471    }
472
473    /// A snapshot of the decisions of the last `report_plans` plans (plan time, run-time choices
474    /// and runtime fallbacks).
475    pub fn report(&self) -> Report {
476        Report(lock(&self.log).entries.iter().map(|(_, d)| d.clone()).collect())
477    }
478
479    /// Empties the report.
480    pub fn clear_report(&self) {
481        lock(&self.log).entries.clear();
482    }
483
484    fn record(&self, plan: u64, node: &Arc<dyn ExecutionPlan>, taken: bool, reason: String) {
485        let node = displayable(node.as_ref()).one_line().to_string().trim_end().to_string();
486        lock(&self.log).push(plan, Decision { node, taken, reason, runtime_fallback: false, groups: None });
487    }
488
489    /// Why a sort with this `fetch` is switched off in the config, if it is.
490    fn sort_disabled(&self, fetch: Option<usize>) -> Option<String> {
491        match fetch {
492            None if !self.config.sort => Some("sort disabled in config".into()),
493            Some(n) if !self.config.topk => Some(format!("top-k (sort with fetch {n}) disabled in config")),
494            _ => None,
495        }
496    }
497
498    /// The node this rule would replace, as an operation over an input, or `None` when the node is
499    /// not one it handles at all (a scan, a projection, ...).
500    fn candidate(&self, node: &Arc<dyn ExecutionPlan>) -> Option<Candidate> {
501        if let Some(spm) = node.downcast_ref::<SortPreservingMergeExec>() {
502            // `SPM -> SortExec`, or `SPM -> ProjectionExec -> SortExec`: DataFusion 55 plans the
503            // latter for any ORDER BY whose SELECT list reorders or computes columns (the
504            // projection stays above the per-partition sorts). The pair is replaced by one
505            // MetalExec sort, and the projection is put back above it.
506            let (sort_node, proj) = if spm.input().downcast_ref::<SortExec>().is_some() {
507                (Arc::clone(spm.input()), None)
508            } else {
509                let p = spm.input().downcast_ref::<ProjectionExec>()?;
510                p.input().downcast_ref::<SortExec>()?;
511                (Arc::clone(p.input()), Some(Arc::clone(spm.input())))
512            };
513            let sort = sort_node.downcast_ref::<SortExec>()?;
514            let fetch = match (spm.fetch(), sort.fetch()) {
515                (Some(a), Some(b)) => Some(a.min(b)),
516                (a, b) => a.or(b),
517            };
518            if let Some(why) = self.sort_disabled(fetch) {
519                return Some(Candidate::new(Err(why), Arc::clone(sort.input())));
520            }
521            let same_order = match &proj {
522                None => sort.expr() == spm.expr(),
523                Some(p) => {
524                    let p = p.downcast_ref::<ProjectionExec>()?;
525                    let mapped = map_through_projection(spm.expr(), p);
526                    mapped.is_some() && mapped == sort_keys(sort.expr())
527                }
528            };
529            if !same_order {
530                return Some(Candidate::new(
531                    Err("merge ordering differs from the sort's".into()),
532                    Arc::clone(sort.input()),
533                ));
534            }
535            let input = Arc::clone(sort.input());
536            let op = translate::sort_op(sort.expr(), &input.schema(), fetch);
537            let mut c = Candidate::new(op, input);
538            if proj.is_some() {
539                // The fallback is the merge over the sorts, without the projection (which stays).
540                c.original = Some(Arc::new(
541                    SortPreservingMergeExec::new(sort.expr().clone(), Arc::clone(&sort_node)).with_fetch(fetch),
542                ));
543                c.reparent = proj;
544            }
545            return Some(c);
546        }
547        if let Some(sort) = node.downcast_ref::<SortExec>() {
548            let input = Arc::clone(sort.input());
549            if let Some(why) = self.sort_disabled(sort.fetch()) {
550                return Some(Candidate::new(Err(why), input));
551            }
552            if sort.preserve_partitioning() && input.output_partitioning().partition_count() > 1 {
553                return Some(Candidate::new(Err("per-partition sort (preserve_partitioning) with no replaced merge above it".into()), input));
554            }
555            let op = translate::sort_op(sort.expr(), &input.schema(), sort.fetch());
556            return Some(Candidate::new(op, input));
557        }
558        if let Some(agg) = node.downcast_ref::<AggregateExec>() {
559            if !self.config.aggregate {
560                return Some(Candidate::new(Err("aggregate disabled in config".into()), Arc::clone(agg.input())));
561            }
562            if node.output_ordering().is_some() {
563                return Some(Candidate::new(
564                    Err("the aggregate's output carries an ordering (sorted input), which the GPU group-by does not keep".into()),
565                    Arc::clone(agg.input()),
566                ));
567            }
568            return Some(match agg.mode() {
569                AggregateMode::Single | AggregateMode::SinglePartitioned => {
570                    Candidate::new(translate::aggregate_op(agg), Arc::clone(agg.input()))
571                }
572                AggregateMode::Final | AggregateMode::FinalPartitioned => {
573                    // Walk down through the exchange to the Partial that feeds this Final; the pair
574                    // is replaced as one, reading the Partial's input.
575                    let mut cur = Arc::clone(agg.input());
576                    loop {
577                        if cur.downcast_ref::<RepartitionExec>().is_some()
578                            || cur.downcast_ref::<CoalescePartitionsExec>().is_some()
579                        {
580                            let next = Arc::clone(cur.children()[0]);
581                            cur = next;
582                            continue;
583                        }
584                        break;
585                    }
586                    match cur.downcast_ref::<AggregateExec>() {
587                        Some(p) if *p.mode() == AggregateMode::Partial && translate::same_aggregates(agg, p) => {
588                            Candidate::new(translate::aggregate_op(p), Arc::clone(p.input()))
589                        }
590                        _ => Candidate::new(Err("Final aggregate without a matching Partial below its exchange".into()), Arc::clone(agg.input())),
591                    }
592                }
593                AggregateMode::Partial => Candidate::new(Err("Partial aggregate whose Final was not replaced".into()), Arc::clone(agg.input())),
594                AggregateMode::PartialReduce => Candidate::new(Err("PartialReduce aggregate".into()), Arc::clone(agg.input())),
595            });
596        }
597        if let Some(f) = node.downcast_ref::<FilterExec>() {
598            let input = Arc::clone(f.input());
599            if !self.config.filter {
600                return Some(Candidate::new(Err("filter disabled in config".into()), input));
601            }
602            if node.fetch().is_some() {
603                return Some(Candidate::new(Err("filter with a fetch limit".into()), input));
604            }
605            if node.output_ordering().is_some() && input.output_partitioning().partition_count() > 1 {
606                return Some(Candidate::new(Err("order-preserving filter over several partitions".into()), input));
607            }
608            let projection = f.projection().as_ref().map(|p| p.iter().copied().collect::<Vec<usize>>());
609            let op = translate::filter_op(f.predicate(), &input.schema(), projection);
610            return Some(Candidate::new(op, input));
611        }
612        if let Some(j) = node.downcast_ref::<HashJoinExec>() {
613            let (left, right) = (below_exchanges(j.left()), below_exchanges(j.right()));
614            // DataFusion's hash join keeps the probe (right) side's order per partition for inner
615            // and right joins, and plans may rely on it (a sort pushed below the join). The plan
616            // runner's join keeps its left input's order, which is the probe side for both; so an
617            // ordered join is taken when its probe input is one partition read as it is.
618            let ordered_ok = matches!(j.join_type(), JoinType::Inner | JoinType::Right)
619                && Arc::ptr_eq(&right, j.right())
620                && j.right().output_partitioning().partition_count() == 1;
621            let op = if !self.config.join {
622                Err("join disabled in config".into())
623            } else if node.output_ordering().is_some() && !ordered_ok {
624                Err("the join's output carries an ordering of a probe side split over partitions".into())
625            } else if Arc::ptr_eq(&left, &right) {
626                Err("both join inputs are one plan node".into())
627            } else {
628                translate::join_op(j)
629            };
630            let mut c = Candidate::new(op, left);
631            c.right = Some(right);
632            return Some(c);
633        }
634        None
635    }
636
637    /// The input's row count as DataFusion's statistics give it.
638    fn statistics_rows(input: &Arc<dyn ExecutionPlan>) -> Precision<usize> {
639        match StatisticsContext::new().compute(input.as_ref(), &StatisticsArgs::new()) {
640            Ok(s) => s.num_rows,
641            Err(_) => Precision::Absent,
642        }
643    }
644
645    /// A join input's row count, when the rule may go by it (exact, or an accepted estimate), and
646    /// the phrase the report uses for it.
647    fn row_count(&self, input: &Arc<dyn ExecutionPlan>) -> (Option<usize>, String) {
648        match Self::statistics_rows(input) {
649            Precision::Exact(n) => (Some(n), format!("{n} (exact)")),
650            Precision::Inexact(n) if self.config.accept_inexact => (Some(n), format!("~{n} (inexact, accepted)")),
651            Precision::Inexact(n) => (None, format!("~{n} (an estimate; accept_inexact is off)")),
652            Precision::Absent => (None, "unknown".into()),
653        }
654    }
655
656    /// Whether the input is big enough, the phrase the report uses for its size, and the row count
657    /// it went by (None when unknown).
658    fn size_ok(&self, input: &Arc<dyn ExecutionPlan>) -> (bool, String, Option<usize>) {
659        let rows = Self::statistics_rows(input);
660        let min = self.config.min_rows;
661        match rows {
662            Precision::Exact(n) => (n >= min, format!("input rows {n} (exact) vs min_rows {min}"), Some(n)),
663            Precision::Inexact(n) if self.config.accept_inexact => {
664                (n >= min, format!("input rows ~{n} (inexact, accepted) vs min_rows {min}"), Some(n))
665            }
666            Precision::Inexact(n) => (false, format!("input rows ~{n} are an estimate (accept_inexact is off)"), None),
667            Precision::Absent => (
668                self.config.take_when_unknown,
669                format!("input row count unknown (take_when_unknown = {})", self.config.take_when_unknown),
670                None,
671            ),
672        }
673    }
674
675    /// `top` is the node under the root's chain of projections (by address): nothing above it
676    /// requires a distribution, so a replacement there keeps MetalExec's single output partition.
677    fn visit(&self, plan: u64, node: Arc<dyn ExecutionPlan>, top: usize) -> Result<Transformed<Arc<dyn ExecutionPlan>>> {
678        if node.downcast_ref::<MetalExec>().is_some() {
679            return Ok(Transformed::no(node));
680        }
681        let Some(c) = self.candidate(&node) else {
682            return Ok(Transformed::no(node));
683        };
684        let op = match c.op {
685            Ok(op) => op,
686            Err(why) => {
687                self.record(plan, &node, false, why);
688                return Ok(Transformed::no(node));
689            }
690        };
691        // Keep the replaced node's partition count, so every parent's distribution requirement
692        // (fixed by EnsureRequirements before this rule runs) still holds. At the top of the plan
693        // (only projections above) there is no such requirement: re-partitioning the output there
694        // would only hash every result row to split it, for `collect` to merge it again.
695        let at_top = Arc::as_ptr(&node) as *const () as usize == top;
696        let wrap = match node.output_partitioning() {
697            p if p.partition_count() <= 1 => None,
698            _ if at_top => None,
699            Partitioning::Hash(exprs, n) => Some(Partitioning::Hash(exprs.clone(), *n)),
700            Partitioning::RoundRobinBatch(n) | Partitioning::UnknownPartitioning(n) => {
701                Some(Partitioning::RoundRobinBatch(*n))
702            }
703            Partitioning::Range(_) => {
704                self.record(plan, &node, false, "range-partitioned output".into());
705                return Ok(Transformed::no(node));
706            }
707        };
708        // MetalExec coalesces its input, so a round-robin split right below it is pure overhead:
709        // read what feeds the split instead. (The fallback plan keeps the split.)
710        let mut input = Arc::clone(&c.input);
711        while let Some(r) = input.downcast_ref::<RepartitionExec>() {
712            if !matches!(r.partitioning(), Partitioning::RoundRobinBatch(_)) {
713                break;
714            }
715            let next = Arc::clone(r.input());
716            input = next;
717        }
718        let mut inputs = vec![Arc::clone(&input)];
719        let (ok, mut size, rows) = match &c.right {
720            None => self.size_ok(&input),
721            Some(right) => {
722                // A join: both inputs need a row count; `min_rows` is compared with the larger.
723                inputs.push(Arc::clone(right));
724                let (lrows, ltext) = self.row_count(&input);
725                let (rrows, rtext) = self.row_count(right);
726                let rows_text = format!("left (build) rows {ltext}, right (probe) rows {rtext}");
727                match (lrows, rrows) {
728                    (Some(l), Some(r)) => {
729                        let min = self.config.min_rows;
730                        let n = l.max(r);
731                        let mut why = format!("{rows_text}; the larger vs min_rows {min}");
732                        let mut ok = n >= min;
733                        if ok && self.config.join_choice == JoinChoice::Measured {
734                            match crate::choice::join_takes(&op, input.as_ref(), right.as_ref(), l as u64, r as u64) {
735                                Ok(w) => why = format!("{why}; {w}"),
736                                Err(w) => {
737                                    why = format!("{why}; {w}");
738                                    ok = false;
739                                }
740                            }
741                        }
742                        (ok, why, Some(l + r))
743                    }
744                    _ => (false, format!("{rows_text}: a join needs both row counts"), None),
745                }
746            }
747        };
748        if !ok {
749            self.record(plan, &node, false, size);
750            return Ok(Transformed::no(node));
751        }
752        // A replaced aggregate is decided at run time; under `Measured` it is replaced only when
753        // the measured table takes its shape at some group count at this row count.
754        if matches!(op, MetalOp::Aggregate { .. }) && self.config.aggregate_choice == AggregateChoice::Measured {
755            let Some(shape) = crate::choice::shape(&op, input.as_ref()) else {
756                self.record(plan, &node, false, "aggregate without a shape".into());
757                return Ok(Transformed::no(node));
758            };
759            if let Some(n) = self.config.table_rows.or(rows) {
760                match crate::choice::any_bucket(&shape, n as u64) {
761                    Ok(why) => size = format!("{size}; {why}"),
762                    Err(why) => {
763                        self.record(plan, &node, false, format!("{size}; {why}"));
764                        return Ok(Transformed::no(node));
765                    }
766                }
767            }
768        }
769        let original = c.original.unwrap_or_else(|| Arc::clone(&node));
770        let mut wrap = wrap;
771        let mut out_parts = 1;
772        if c.right.is_some() && !at_top {
773            // A join keeps the replaced node's partition count itself (one GPU execution dealt
774            // out to the partitions); only a hash partitioning still needs a re-partition above.
775            out_parts = node.output_partitioning().partition_count();
776            if !matches!(wrap, Some(Partitioning::Hash(..))) {
777                wrap = None;
778            }
779        }
780        let metal: Arc<dyn ExecutionPlan> = Arc::new(
781            MetalExec::new(
782                op,
783                inputs,
784                original,
785                (Arc::clone(&self.log), plan),
786                crate::exec::AggSettings {
787                    choice: self.config.aggregate_choice,
788                    table_rows: self.config.table_rows,
789                    rows_hint: rows,
790                },
791            )
792            .with_output_partitions(out_parts),
793        );
794        let mut reason = size;
795        if out_parts > 1 {
796            reason.push_str(&format!("; output in {out_parts} partitions"));
797        }
798        if at_top && node.output_partitioning().partition_count() > 1 {
799            reason.push_str("; output kept at one partition (only projections above it)");
800        }
801        let metal = match c.reparent {
802            #[allow(deprecated)] // `replace_children` is the 55 name; `with_new_children` still works
803            Some(p) => {
804                reason.push_str("; projection kept above it");
805                p.with_new_children(vec![metal])?
806            }
807            None => metal,
808        };
809        let out: Arc<dyn ExecutionPlan> = match wrap {
810            None => metal,
811            Some(p) => {
812                if c.right.is_some() && matches!(p, Partitioning::Hash(..)) {
813                    reason.push_str(&format!(
814                        "; output re-partitioned to {p} where the parent requires it (to keep the parent's distribution)"
815                    ));
816                } else {
817                    reason.push_str(&format!("; output re-partitioned to {p} to keep the parent's distribution"));
818                }
819                Arc::new(RepartitionExec::try_new(metal, p)?)
820            }
821        };
822        self.record(plan, &node, true, reason);
823        Ok(Transformed::yes(out))
824    }
825}
826
827/// A replaced hash-partitioned join gets a hash `RepartitionExec` above it, so a parent that
828/// requires that distribution still has it. Where the parent does not require a hash
829/// distribution of that input (a partial aggregate, a projection, a filter), hashing every joined
830/// row only to split it is dropped: the join's own output partitions feed the parent.
831#[allow(deprecated)] // `required_input_distribution` and `with_new_children`: the 55 names still work
832fn drop_hash_where_not_required(node: Arc<dyn ExecutionPlan>) -> Result<Transformed<Arc<dyn ExecutionPlan>>> {
833    let required = node.required_input_distribution();
834    let children: Vec<Arc<dyn ExecutionPlan>> = node.children().into_iter().cloned().collect();
835    let mut changed = false;
836    let mut new_children = Vec::with_capacity(children.len());
837    for (i, c) in children.into_iter().enumerate() {
838        let swap = match c.downcast_ref::<RepartitionExec>() {
839            Some(r) => {
840                let over_join = r
841                    .input()
842                    .downcast_ref::<MetalExec>()
843                    .is_some_and(|m| matches!(m.op(), MetalOp::Join { .. }));
844                match r.partitioning() {
845                    Partitioning::Hash(..)
846                        if over_join && matches!(required.get(i), Some(Distribution::UnspecifiedDistribution)) =>
847                    {
848                        Some(Arc::clone(r.input()))
849                    }
850                    _ => None,
851                }
852            }
853            None => None,
854        };
855        match swap {
856            Some(s) => {
857                changed = true;
858                new_children.push(s);
859            }
860            None => new_children.push(c),
861        }
862    }
863    if !changed {
864        return Ok(Transformed::no(node));
865    }
866    Ok(Transformed::yes(node.with_new_children(new_children)?))
867}
868
869impl PhysicalOptimizerRule for ArrowMetalRule {
870    fn optimize(&self, plan: Arc<dyn ExecutionPlan>, config: &ConfigOptions) -> Result<Arc<dyn ExecutionPlan>> {
871        let mut top = Arc::clone(&plan);
872        while let Some(p) = top.downcast_ref::<ProjectionExec>() {
873            let next = Arc::clone(p.input());
874            top = next;
875        }
876        let top = Arc::as_ptr(&top) as *const () as usize;
877        let id = lock(&self.log).begin_plan();
878        let plan = plan.transform_down(|n| self.visit(id, n, top))?.data;
879        // Dropped only when DataFusion's own check of every node's required input distribution
880        // still passes (a projection between a join and a hash-partitioned aggregate passes the
881        // join's partitioning through).
882        let relaxed = Arc::clone(&plan).transform_up(drop_hash_where_not_required)?;
883        if relaxed.transformed && SanityCheckPlan::new().optimize(Arc::clone(&relaxed.data), config).is_ok() {
884            return Ok(relaxed.data);
885        }
886        Ok(plan)
887    }
888
889    fn name(&self) -> &str {
890        "ArrowMetalRule"
891    }
892
893    fn schema_check(&self) -> bool {
894        true
895    }
896}