Skip to main content

datafusion_arrowmetal/
exec.rs

1//! `MetalExec`: the `ExecutionPlan` that runs a replaced node on ArrowMetal.
2
3use std::fmt;
4use std::sync::{Arc, Mutex};
5
6use arrow::array::{Array, AsArray};
7use arrow::datatypes::{DataType, SchemaRef};
8use arrow::record_batch::RecordBatch;
9use datafusion::common::runtime::SpawnedTask;
10use datafusion::common::tree_node::{Transformed, TreeNode, TreeNodeRecursion};
11use datafusion::common::{internal_err, DataFusionError, Result};
12use datafusion::execution::memory_pool::{MemoryConsumer, MemoryReservation};
13use datafusion::execution::TaskContext;
14use datafusion::physical_expr::{EquivalenceProperties, PhysicalExpr};
15use datafusion::physical_plan::coalesce_partitions::CoalescePartitionsExec;
16use datafusion::physical_plan::execution_plan::{Boundedness, EmissionType};
17use datafusion::physical_plan::metrics::{Count, ExecutionPlanMetricsSet, Gauge, MetricBuilder, MetricsSet, Time};
18use datafusion::physical_plan::stream::RecordBatchStreamAdapter;
19use datafusion::physical_plan::{
20    execute_stream, DisplayAs, DisplayFormatType, ExecutionPlan, ExecutionPlanProperties, Partitioning,
21    PlanProperties, SendableRecordBatchStream,
22};
23use futures::stream::BoxStream;
24use futures::{FutureExt, StreamExt, TryStreamExt};
25
26use crate::rule::{lock, AggregateChoice, Decision, SharedLog};
27
28/// One sort key: an input column, a direction, and where its nulls go.
29#[derive(Debug, Clone, PartialEq)]
30#[non_exhaustive]
31pub struct SortKey {
32    /// The input column.
33    pub column: usize,
34    /// Descending order.
35    pub descending: bool,
36    /// Nulls before every value.
37    pub nulls_first: bool,
38}
39
40/// An aggregate function `MetalExec` computes.
41#[derive(Debug, Clone, Copy, PartialEq, Eq)]
42#[non_exhaustive]
43pub enum AggKind {
44    /// `sum`.
45    Sum,
46    /// `min`.
47    Min,
48    /// `max`.
49    Max,
50    /// `count(col)`: non-null values.
51    Count,
52    /// `count(*)`: rows.
53    CountAll,
54    /// `avg`.
55    Mean,
56}
57
58/// One aggregate of a [`MetalOp::Aggregate`].
59#[derive(Debug, Clone, PartialEq)]
60#[non_exhaustive]
61pub struct AggSpec {
62    /// The function.
63    pub kind: AggKind,
64    /// The argument column (`None` for `count(*)`).
65    pub column: Option<usize>,
66    /// The argument column is floating point (min/max keep DataFusion's NaN, signed-zero and
67    /// infinity semantics).
68    pub float: bool,
69    /// The type DataFusion's schema gives the result.
70    pub out_type: DataType,
71}
72
73/// Which rows of an equi-join come out (the `JoinType`s of DataFusion's `HashJoinExec` that
74/// `MetalExec` runs).
75#[derive(Debug, Clone, Copy, PartialEq, Eq)]
76#[non_exhaustive]
77pub enum JoinHow {
78    /// The pairs of rows whose keys are equal.
79    Inner,
80    /// The inner pairs, and every left (build-side) row without a match, with nulls on the right.
81    Left,
82    /// The inner pairs, and every right (probe-side) row without a match, with nulls on the left.
83    Right,
84}
85
86/// What a `MetalExec` computes, in terms of its input's column indices.
87#[derive(Debug, Clone, PartialEq)]
88#[non_exhaustive]
89pub enum MetalOp {
90    /// A sort (`fetch`: under a LIMIT).
91    Sort {
92        /// The ORDER BY keys, first to last.
93        keys: Vec<SortKey>,
94        /// The LIMIT, if any.
95        fetch: Option<usize>,
96    },
97    /// A GROUP BY over column keys (no aggregates: DISTINCT).
98    Aggregate {
99        /// The key columns.
100        keys: Vec<usize>,
101        /// The aggregates, in output order after the keys.
102        aggs: Vec<AggSpec>,
103    },
104    /// A filter, with the projection DataFusion's `FilterExec` carried.
105    Filter {
106        /// An ArrowMetal s-expression over columns named `c{i}`.
107        predicate: String,
108        /// The output columns, if the filter projects.
109        projection: Option<Vec<usize>>,
110    },
111    /// An equi-join over column keys (DataFusion's `HashJoinExec`), both inputs collected. The
112    /// output is DataFusion's: the left input's columns, then the right input's, through
113    /// `projection`.
114    Join {
115        /// Which unmatched rows are kept.
116        how: JoinHow,
117        /// The key columns of the left (build) input, in the join's `on` order.
118        left_keys: Vec<usize>,
119        /// The key columns of the right (probe) input, in the same order.
120        right_keys: Vec<usize>,
121        /// The output columns, as indices into the left input's columns followed by the right's.
122        projection: Vec<usize>,
123        /// The number of columns of the left input.
124        left_columns: usize,
125    },
126}
127
128/// Runs one [`MetalOp`] on ArrowMetal.
129///
130/// Collects its input (a join: both inputs), hands the batches to ArrowMetal's chunked import (no
131/// `concat_batches`; the columns of one call are imported at the same time, one thread each),
132/// runs the operation on the GPU through ArrowMetal's plan runner, and emits the result in
133/// `batch_size` slices as a single partition (a join: dealt out to as many partitions as the join
134/// it replaced, from one execution).
135///
136/// The replaced subtree (`original`) runs instead, over the batches already collected and the
137/// rest of the input, when:
138///
139/// * a replaced aggregate's run-time choice hands it back ([`AggregateChoice`]);
140/// * ArrowMetal returns an error, or the data holds values the GPU path cannot answer exactly
141///   (a runtime fallback);
142/// * the session's memory pool refuses the reservation for the input or the result (DataFusion's
143///   own operators can spill).
144///
145/// Each of these is recorded in the rule's report. The GPU call runs on a blocking thread and is
146/// not interrupted when the query is dropped: it finishes, and its result is discarded.
147pub struct MetalExec {
148    op: MetalOp,
149    /// What it reads: one input, or a join's left and right inputs.
150    inputs: Vec<Arc<dyn ExecutionPlan>>,
151    /// The subtree this node replaced, with `inputs` as its leaves: the runtime fallback.
152    original: Arc<dyn ExecutionPlan>,
153    props: Arc<PlanProperties>,
154    log: SharedLog,
155    /// The plan (of the rule's log) this node belongs to.
156    plan: u64,
157    metrics: ExecutionPlanMetricsSet,
158    /// For an aggregate: who runs it (decided at run time under `AggregateChoice::Measured`).
159    choice: AggregateChoice,
160    /// Look the run-time decision up at this row count instead of the input's.
161    table_rows: Option<usize>,
162    /// The input's row count as the plan's statistics give it (exact, or an accepted estimate).
163    rows_hint: Option<usize>,
164    /// A join's execution, shared by its output partitions.
165    join_slot: JoinSlot,
166}
167
168/// How a replaced aggregate is decided (see [`AggregateChoice`]).
169#[derive(Debug, Clone, Copy)]
170pub(crate) struct AggSettings {
171    pub choice: AggregateChoice,
172    /// Look the table up at this row count instead of the input's.
173    pub table_rows: Option<usize>,
174    /// The input's row count as the plan's statistics give it.
175    pub rows_hint: Option<usize>,
176}
177
178/// The per-phase metrics of one execution.
179struct PhaseMetrics {
180    input_time: Time,
181    import_time: Time,
182    kernel_time: Time,
183    export_time: Time,
184    probe_time: Time,
185    input_batches: Count,
186    handed_back: Count,
187    groups_estimate: Gauge,
188}
189
190impl PhaseMetrics {
191    fn add(&self, t: &crate::gpu::GpuTimes) {
192        self.import_time.add_duration(t.import);
193        self.kernel_time.add_duration(t.kernel);
194        self.export_time.add_duration(t.export);
195    }
196}
197
198/// `b` in slices of at most `batch_size` rows.
199fn slices(b: RecordBatch, batch_size: usize) -> Vec<RecordBatch> {
200    let mut v = Vec::new();
201    let mut off = 0;
202    while off < b.num_rows() {
203        let n = batch_size.min(b.num_rows() - off);
204        v.push(b.slice(off, n));
205        off += n;
206    }
207    v
208}
209
210/// The bytes a batch's columns reference: for a slice of a larger buffer, only the slice (so the
211/// slices of one table are not each counted at the table's size), without allocating.
212pub(crate) fn batch_bytes(b: &RecordBatch) -> usize {
213    b.columns().iter().map(|a| array_bytes(a.as_ref())).sum()
214}
215
216fn array_bytes(a: &dyn Array) -> usize {
217    let n = a.len();
218    let nulls = a.nulls().map_or(0, |_| n.div_ceil(8));
219    let values = match a.data_type() {
220        DataType::Boolean => n.div_ceil(8),
221        DataType::Utf8 => {
222            let o = a.as_string::<i32>().value_offsets();
223            (n + 1) * 4 + (o[n] - o[0]) as usize
224        }
225        DataType::LargeUtf8 => {
226            let o = a.as_string::<i64>().value_offsets();
227            (n + 1) * 8 + (o[n] - o[0]) as usize
228        }
229        DataType::Utf8View => {
230            let v = a.as_string_view();
231            n * 16 + v.views().iter().map(|w| *w as u32 as usize).filter(|&l| l > 12).sum::<usize>()
232        }
233        t => match t.primitive_width() {
234            Some(w) => n * w,
235            None => return a.get_array_memory_size(),
236        },
237    };
238    values + nulls
239}
240
241/// What collecting one or more input streams gave.
242struct Collected {
243    /// The batches of each stream, in order.
244    parts: Vec<Vec<RecordBatch>>,
245    /// The rest of each stream, where the memory pool refused a batch.
246    rest: Vec<Option<SendableRecordBatchStream>>,
247    /// The reservations holding `parts`.
248    held: Vec<MemoryReservation>,
249    /// Why the pool refused, if it did.
250    refused: Option<String>,
251    /// How many of the streams belong to each input, in order (one input but for a join).
252    per_input: Vec<usize>,
253}
254
255impl Collected {
256    fn rows(&self) -> usize {
257        self.parts.iter().flatten().map(|b| b.num_rows()).sum()
258    }
259}
260
261/// Rows of the prefix a replaced aggregate decides from under `AggregateChoice::Measured`: the
262/// first batches of each input partition, together at least this many rows (the probe samples at
263/// most a quarter of them).
264const PREFIX_ROWS: usize = 262_144;
265
266/// The largest sample the run-time choice draws. A range still not settled at this size hands the
267/// node back: at 5M and 10M rows with about rows / 2 groups, growing the sample to 32,768 rows took
268/// 0.35 to 0.53 ms, 4 % to 8 % of DataFusion's time for the query.
269const DECIDE_SAMPLE: usize = 8_192;
270
271/// The largest sample of the whole input that confirms a take decided from the prefix.
272const CONFIRM_SAMPLE: usize = 2_048;
273
274/// Drains `streams`, reserving each batch in the memory pool: every stream to its end (one task
275/// each), or with `quota`, each until it has given at least `quota` rows (polled together in
276/// this task; the rest of each stream is kept). A stream whose reservation is refused stops there
277/// and keeps the rest of its stream.
278async fn collect_streams(
279    streams: Vec<SendableRecordBatchStream>,
280    ctx: &TaskContext,
281    quota: Option<usize>,
282) -> Result<Collected> {
283    let drain = |mut s: SendableRecordBatchStream, r: MemoryReservation| async move {
284        let mut v = Vec::new();
285        let mut rows = 0usize;
286        while let Some(b) = s.next().await {
287            let b = b?;
288            rows += b.num_rows();
289            let refused = r.try_grow(batch_bytes(&b)).err().map(|e| e.to_string());
290            v.push(b);
291            if refused.is_some() || quota.is_some_and(|q| rows >= q) {
292                return Ok::<_, DataFusionError>((v, Some(s), r, refused));
293            }
294        }
295        Ok((v, None, r, None))
296    };
297    let reservation = || MemoryConsumer::new("MetalExec").register(ctx.memory_pool());
298    let results: Vec<_> = if quota.is_some() || streams.len() == 1 {
299        futures::future::join_all(streams.into_iter().map(|s| drain(s, reservation())))
300            .await
301            .into_iter()
302            .collect::<Result<_>>()?
303    } else {
304        let tasks: Vec<_> = streams.into_iter().map(|s| SpawnedTask::spawn(drain(s, reservation()))).collect();
305        let mut v = Vec::with_capacity(tasks.len());
306        for t in tasks {
307            v.push(t.join().await.map_err(|e| DataFusionError::External(Box::new(e)))??);
308        }
309        v
310    };
311    let mut out = Collected { parts: Vec::new(), rest: Vec::new(), held: Vec::new(), refused: None, per_input: Vec::new() };
312    for (v, rest, r, refused) in results {
313        out.parts.push(v);
314        out.rest.push(rest);
315        out.held.push(r);
316        if out.refused.is_none() {
317            out.refused = refused;
318        }
319    }
320    if out.parts.is_empty() {
321        out.parts.push(Vec::new());
322        out.rest.push(None);
323    }
324    out.per_input = vec![out.parts.len()];
325    Ok(out)
326}
327
328/// `c` with the rest of each of its streams collected (a stream refused by the pool stays open,
329/// and `refused` says why).
330async fn finish(mut c: Collected, ctx: &TaskContext) -> Result<Collected> {
331    let open: Vec<usize> = (0..c.rest.len()).filter(|&i| c.rest[i].is_some()).collect();
332    if open.is_empty() {
333        return Ok(c);
334    }
335    let streams: Vec<SendableRecordBatchStream> = open.iter().filter_map(|&i| c.rest[i].take()).collect();
336    let more = collect_streams(streams, ctx, None).await?;
337    for ((i, v), rest) in open.into_iter().zip(more.parts).zip(more.rest) {
338        c.parts[i].extend(v);
339        c.rest[i] = rest;
340    }
341    c.held.extend(more.held);
342    if c.refused.is_none() {
343        c.refused = more.refused;
344    }
345    Ok(c)
346}
347
348/// A leaf that replays collected batches, then the rest of their streams: the input of the
349/// replaced subtree when `MetalExec` hands a node back. One partition per collected stream; each
350/// partition can be executed once.
351/// One partition of a [`ReplayExec`]: the collected batches and the rest of their stream.
352type ReplaySlot = Option<(Vec<RecordBatch>, Option<SendableRecordBatchStream>)>;
353
354struct ReplayExec {
355    schema: SchemaRef,
356    slots: Mutex<Vec<ReplaySlot>>,
357    props: Arc<PlanProperties>,
358}
359
360impl fmt::Debug for ReplayExec {
361    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
362        f.debug_struct("ReplayExec").finish()
363    }
364}
365
366impl ReplayExec {
367    fn new(schema: SchemaRef, parts: Vec<Vec<RecordBatch>>, rest: Vec<Option<SendableRecordBatchStream>>) -> Self {
368        let n = parts.len();
369        let props = PlanProperties::new(
370            EquivalenceProperties::new(Arc::clone(&schema)),
371            Partitioning::UnknownPartitioning(n),
372            EmissionType::Incremental,
373            Boundedness::Bounded,
374        );
375        let slots = parts.into_iter().zip(rest).map(Some).collect();
376        Self { schema, slots: Mutex::new(slots), props: Arc::new(props) }
377    }
378}
379
380impl DisplayAs for ReplayExec {
381    fn fmt_as(&self, _t: DisplayFormatType, f: &mut fmt::Formatter) -> fmt::Result {
382        write!(f, "ReplayExec")
383    }
384}
385
386impl ExecutionPlan for ReplayExec {
387    fn name(&self) -> &str {
388        "ReplayExec"
389    }
390    fn properties(&self) -> &Arc<PlanProperties> {
391        &self.props
392    }
393    fn children(&self) -> Vec<&Arc<dyn ExecutionPlan>> {
394        vec![]
395    }
396    fn apply_expressions(
397        &self,
398        _f: &mut dyn FnMut(&Arc<dyn PhysicalExpr>) -> Result<TreeNodeRecursion>,
399    ) -> Result<TreeNodeRecursion> {
400        Ok(TreeNodeRecursion::Continue)
401    }
402    fn with_new_children(self: Arc<Self>, children: Vec<Arc<dyn ExecutionPlan>>) -> Result<Arc<dyn ExecutionPlan>> {
403        if children.is_empty() {
404            Ok(self)
405        } else {
406            internal_err!("ReplayExec has no children")
407        }
408    }
409    fn execute(&self, partition: usize, _ctx: Arc<TaskContext>) -> Result<SendableRecordBatchStream> {
410        let slot = self.slots.lock().unwrap_or_else(|p| p.into_inner()).get_mut(partition).and_then(Option::take);
411        let Some((batches, rest)) = slot else {
412            return internal_err!("ReplayExec partition {partition} executed twice or out of range");
413        };
414        let head = futures::stream::iter(batches.into_iter().map(Ok));
415        let s: BoxStream<'static, Result<RecordBatch>> = match rest {
416            Some(r) => Box::pin(head.chain(r)),
417            None => Box::pin(head),
418        };
419        Ok(Box::pin(RecordBatchStreamAdapter::new(Arc::clone(&self.schema), s)))
420    }
421}
422
423impl fmt::Debug for MetalExec {
424    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
425        f.debug_struct("MetalExec").field("op", &self.op).finish()
426    }
427}
428
429/// Everything one execution needs, moved into its future.
430struct Job {
431    op: MetalOp,
432    inputs: Vec<Arc<dyn ExecutionPlan>>,
433    original: Arc<dyn ExecutionPlan>,
434    out_schema: SchemaRef,
435    log: SharedLog,
436    plan: u64,
437    choice: AggregateChoice,
438    table_rows: Option<usize>,
439    rows_hint: Option<usize>,
440    mm: PhaseMetrics,
441    ctx: Arc<TaskContext>,
442}
443
444impl Job {
445    fn record(&self, d: Decision) {
446        lock(&self.log).push(self.plan, d);
447    }
448
449    /// The replaced subtree run by DataFusion over `c` (the collected batches, then the rest of
450    /// each stream), one partition per collected stream, coalesced to one stream. `c`'s
451    /// reservations are released first: DataFusion's operators reserve what they buffer as they
452    /// consume the replayed batches (and can spill), which they could not with this node still
453    /// holding the pool. A join's streams are its left input's partitions, then its right's
454    /// (`c.per_input`).
455    fn hand_back(&self, c: Collected) -> Result<BoxStream<'static, Result<RecordBatch>>> {
456        let plan = self.hand_back_plan(c)?;
457        Ok(Box::pin(execute_stream(plan, Arc::clone(&self.ctx))?))
458    }
459
460    /// The replaced subtree over `c` (see [`hand_back`](Self::hand_back)), not yet executed.
461    fn hand_back_plan(&self, c: Collected) -> Result<Arc<dyn ExecutionPlan>> {
462        drop(c.held);
463        let per_input = if c.per_input.len() == self.inputs.len() { c.per_input } else { vec![c.parts.len()] };
464        let (mut parts, mut rest) = (c.parts.into_iter(), c.rest.into_iter());
465        let mut plan = Arc::clone(&self.original);
466        for (input, n) in self.inputs.iter().zip(per_input) {
467            let p: Vec<Vec<RecordBatch>> = parts.by_ref().take(n).collect();
468            let r: Vec<Option<SendableRecordBatchStream>> = rest.by_ref().take(n).collect();
469            let leaf: Arc<dyn ExecutionPlan> = Arc::new(ReplayExec::new(input.schema(), p, r));
470            plan = replace_leaf(&plan, input, &leaf)?;
471        }
472        Ok(plan)
473    }
474
475    /// Runs the op on ArrowMetal over `c`; on an error, a panic, a data-dependent refusal or a
476    /// refused result reservation, hands the node back instead (recorded).
477    async fn run_gpu(self, c: Collected) -> Result<BoxStream<'static, Result<RecordBatch>>> {
478        let ctx = Arc::clone(&self.ctx);
479        match self.run_gpu_outcome(c).await? {
480            Outcome::Gpu(batches, held) => Ok(Box::pin(futures::stream::iter(batches.into_iter().map(Ok)).map(move |b| {
481                let _held = &held;
482                b
483            }))),
484            Outcome::Back(plan) => Ok(Box::pin(execute_stream(plan, ctx)?)),
485        }
486    }
487
488    /// [`run_gpu`](Self::run_gpu), with the result as `batch_size` slices and the reservation
489    /// holding them, or the plan DataFusion runs instead.
490    async fn run_gpu_outcome(self, c: Collected) -> Result<Outcome> {
491        let op = self.op.clone();
492        let s2 = Arc::clone(&self.out_schema);
493        let parts = c.parts;
494        let per_input = c.per_input.clone();
495        let (res, parts, times) = tokio::task::spawn_blocking(move || {
496            // The batches of each input (one input but for a join).
497            let mut groups: Vec<Vec<&RecordBatch>> = Vec::new();
498            let mut it = parts.iter();
499            for &n in &per_input {
500                groups.push(it.by_ref().take(n).flatten().collect());
501            }
502            let r = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| crate::gpu::run(&op, &groups, &s2)))
503                .unwrap_or_else(|p| {
504                    let msg = p
505                        .downcast_ref::<String>()
506                        .cloned()
507                        .or_else(|| p.downcast_ref::<&str>().map(|s| s.to_string()))
508                        .unwrap_or_default();
509                    Err(format!("panic in the GPU path: {msg}"))
510                });
511            (r, parts, crate::gpu::take_times())
512        })
513        .await
514        .map_err(|e| DataFusionError::External(Box::new(e)))?;
515        self.mm.add(&times);
516        let c = Collected { parts, ..c };
517        match res {
518            Ok(b) => {
519                let out = MemoryConsumer::new("MetalExec output").register(self.ctx.memory_pool());
520                if let Err(e) = out.try_grow(batch_bytes(&b)) {
521                    self.record(Decision::memory_hand_back(&self.op, "the result", &e.to_string()));
522                    self.mm.handed_back.add(1);
523                    return Ok(Outcome::Back(self.hand_back_plan(c)?));
524                }
525                // The input is released; the result's reservation is held until the stream ends.
526                drop(c);
527                let batch_size = self.ctx.session_config().batch_size();
528                Ok(Outcome::Gpu(slices(b, batch_size), out))
529            }
530            Err(msg) => {
531                self.record(Decision::runtime_fallback(&self.op, &msg));
532                Ok(Outcome::Back(self.hand_back_plan(c)?))
533            }
534        }
535    }
536
537    /// Sort and filter: the input coalesced into one stream, collected, run on ArrowMetal.
538    async fn run_sort_or_filter(self, stream: SendableRecordBatchStream) -> Result<BoxStream<'static, Result<RecordBatch>>> {
539        let t = std::time::Instant::now();
540        let c = collect_streams(vec![stream], &self.ctx, None).await?;
541        self.mm.input_time.add_duration(t.elapsed());
542        self.mm.input_batches.add(c.parts.iter().map(|p| p.len()).sum());
543        if let Some(e) = &c.refused {
544            self.record(Decision::memory_hand_back(&self.op, "the input", e));
545            self.mm.handed_back.add(1);
546            return self.hand_back(c);
547        }
548        self.run_gpu(c).await
549    }
550
551    /// An aggregate: the run-time choice, then ArrowMetal or the hand-back.
552    ///
553    /// * `ArrowMetal`: every input partition collected (concurrently, each kept as its own
554    ///   partition), then the GPU.
555    /// * `DataFusion`: handed back at once; the replaced subtree reads the input streams.
556    /// * `Measured`, with the input's row count known from the plan: the first batches of each
557    ///   partition (`PREFIX_ROWS` together) are collected and probed. A hand-back replays them
558    ///   and streams the rest of each partition into DataFusion's own operators. A take collects
559    ///   the rest and confirms the choice with a probe over the whole input (a prefix of data
560    ///   ordered by its keys shows fewer groups than the input holds) before the GPU runs.
561    /// * `Measured` without a row count: everything collected, probed, then either.
562    async fn run_aggregate(self, streams: Vec<SendableRecordBatchStream>) -> Result<BoxStream<'static, Result<RecordBatch>>> {
563        let MetalOp::Aggregate { keys, .. } = &self.op else {
564            return internal_err!("run_aggregate on {:?}", self.op);
565        };
566        let keys = keys.clone();
567        let t = std::time::Instant::now();
568        if self.choice == AggregateChoice::DataFusion {
569            let n = streams.len();
570            let c = Collected {
571                parts: vec![Vec::new(); n],
572                rest: streams.into_iter().map(Some).collect(),
573                held: Vec::new(),
574                refused: None,
575                per_input: vec![n],
576            };
577            self.record(Decision::runtime_choice(&self.op, false, "aggregate_choice is DataFusion".into(), self.rows_hint.unwrap_or(0), None));
578            self.mm.handed_back.add(1);
579            return self.hand_back(c);
580        }
581        let prefix = match (self.choice, self.rows_hint) {
582            (AggregateChoice::Measured, Some(_)) => Some(PREFIX_ROWS.div_ceil(streams.len().max(1))),
583            _ => None,
584        };
585        let mut c = collect_streams(streams, &self.ctx, prefix).await?;
586        if let Some(e) = &c.refused {
587            self.mm.input_time.add_duration(t.elapsed());
588            self.record(Decision::memory_hand_back(&self.op, "the input", e));
589            self.mm.handed_back.add(1);
590            return self.hand_back(c);
591        }
592        let mut from_prefix: Option<(String, crate::probe::GroupEstimate)> = None;
593        if let (Some(_), Some(total)) = (prefix, self.rows_hint) {
594            if c.rest.iter().any(Option::is_some) {
595                // Decide from the prefix.
596                self.mm.input_time.add_duration(t.elapsed());
597                let t = std::time::Instant::now();
598                let refs: Vec<&RecordBatch> = c.parts.iter().flatten().collect();
599                let (take, reason, estimate) = decide(&self.op, self.inputs[0].as_ref(), &refs, &keys, total, self.table_rows);
600                self.mm.probe_time.add_duration(t.elapsed());
601                if !take {
602                    if let Some(e) = &estimate {
603                        self.mm.groups_estimate.set(e.estimate as usize);
604                    }
605                    let reason = format!("{reason} (from the first {} rows of each partition)", prefix.unwrap_or(0));
606                    self.record(Decision::runtime_choice(&self.op, false, reason, total, estimate));
607                    self.mm.handed_back.add(1);
608                    return self.hand_back(c);
609                }
610                if let Some(e) = estimate {
611                    from_prefix = Some((reason, e));
612                }
613                let t = std::time::Instant::now();
614                c = finish(c, &self.ctx).await?;
615                self.mm.input_time.add_duration(t.elapsed());
616                if let Some(e) = &c.refused {
617                    self.record(Decision::memory_hand_back(&self.op, "the input", e));
618                    self.mm.handed_back.add(1);
619                    return self.hand_back(c);
620                }
621            } else {
622                self.mm.input_time.add_duration(t.elapsed());
623            }
624        } else {
625            self.mm.input_time.add_duration(t.elapsed());
626        }
627        self.mm.input_batches.add(c.parts.iter().map(|p| p.len()).sum());
628        let rows = c.rows();
629        let t = std::time::Instant::now();
630        let (on_gpu, reason, estimate) = match (self.choice, from_prefix) {
631            (AggregateChoice::Measured, Some((why, e))) => {
632                // Taken from the prefix: a small sample of the whole input must agree (its range
633                // must meet the prefix's). Data ordered by its keys shows fewer groups in a prefix
634                // than the input holds; a sample spread over the whole input sees them.
635                let refs: Vec<&RecordBatch> = c.parts.iter().flatten().collect();
636                let (agree, f) = confirm(&refs, &keys, rows, &e);
637                let check = format!(
638                    "a {}-row sample of the whole input puts it at {} to {} groups",
639                    f.sample_rows, f.low, f.high
640                );
641                if agree {
642                    (true, format!("{why} (from the first rows of each partition; {check})"), Some(e))
643                } else {
644                    (false, format!("{why} from the first rows of each partition, but {check}; left to DataFusion"), Some(f))
645                }
646            }
647            (AggregateChoice::Measured, None) => {
648                let refs: Vec<&RecordBatch> = c.parts.iter().flatten().collect();
649                decide(&self.op, self.inputs[0].as_ref(), &refs, &keys, rows, self.table_rows)
650            }
651            _ => (true, "aggregate_choice is ArrowMetal".to_string(), None),
652        };
653        self.mm.probe_time.add_duration(t.elapsed());
654        if let Some(e) = &estimate {
655            self.mm.groups_estimate.set(e.estimate as usize);
656        }
657        self.record(Decision::runtime_choice(&self.op, on_gpu, reason, rows, estimate));
658        if !on_gpu {
659            self.mm.handed_back.add(1);
660            return self.hand_back(c);
661        }
662        self.run_gpu(c).await
663    }
664
665    /// A join: every partition of both inputs collected, each in its own task (all at the same
666    /// time), then the GPU. `left` is the number of `streams` that belong to the left input.
667    async fn run_join(self, streams: Vec<SendableRecordBatchStream>, left: usize) -> Result<Outcome> {
668        let t = std::time::Instant::now();
669        let total = streams.len();
670        let mut c = collect_streams(streams, &self.ctx, None).await?;
671        c.per_input = vec![left, total - left];
672        self.mm.input_time.add_duration(t.elapsed());
673        self.mm.input_batches.add(c.parts.iter().map(|p| p.len()).sum());
674        if let Some(e) = &c.refused {
675            self.record(Decision::memory_hand_back(&self.op, "the input", e));
676            self.mm.handed_back.add(1);
677            return Ok(Outcome::Back(self.hand_back_plan(c)?));
678        }
679        self.run_gpu_outcome(c).await
680    }
681}
682
683/// What one execution of a `MetalExec` produced.
684enum Outcome {
685    /// The GPU's result in `batch_size` slices, and the reservation that holds it.
686    Gpu(Vec<RecordBatch>, MemoryReservation),
687    /// The replaced subtree, over the collected batches, for DataFusion to run instead.
688    Back(Arc<dyn ExecutionPlan>),
689}
690
691/// A join's one execution, shared by its output partitions: started by the first partition
692/// executed, and taken by each partition once.
693type SharedOutcome = futures::future::Shared<futures::future::BoxFuture<'static, std::result::Result<Arc<Outcome>, Arc<DataFusionError>>>>;
694
695/// The shared execution and how many output partitions have taken it.
696type JoinSlot = Arc<Mutex<Option<(SharedOutcome, usize)>>>;
697
698/// Output partition `p` of `n` of a shared join execution: every `n`-th GPU slice from the `p`-th,
699/// or partition `p` of the replaced subtree.
700fn partition_of(shared: SharedOutcome, p: usize, n: usize, ctx: Arc<TaskContext>) -> BoxStream<'static, Result<RecordBatch>> {
701    let s = futures::stream::once(shared).map(move |r| -> Result<BoxStream<'static, Result<RecordBatch>>> {
702        let out = r.map_err(DataFusionError::Shared)?;
703        match &*out {
704            Outcome::Gpu(batches, _) => {
705                let mine: Vec<RecordBatch> = batches.iter().skip(p).step_by(n).cloned().collect();
706                let held = Arc::clone(&out);
707                Ok(Box::pin(futures::stream::iter(mine.into_iter().map(Ok)).map(move |b| {
708                    let _held = &held;
709                    b
710                })))
711            }
712            // `fit_partitions` gave it `n` partitions.
713            Outcome::Back(plan) => Ok(Box::pin(plan.execute(p, Arc::clone(&ctx))?)),
714        }
715    });
716    Box::pin(s.try_flatten())
717}
718
719/// `plan` with `n` output partitions: as it is, coalesced to one, or split round-robin.
720fn fit_partitions(plan: Arc<dyn ExecutionPlan>, n: usize) -> Result<Arc<dyn ExecutionPlan>> {
721    let have = plan.output_partitioning().partition_count();
722    Ok(if have == n {
723        plan
724    } else if n == 1 {
725        Arc::new(CoalescePartitionsExec::new(plan))
726    } else {
727        Arc::new(datafusion::physical_plan::repartition::RepartitionExec::try_new(plan, Partitioning::RoundRobinBatch(n))?)
728    })
729}
730
731impl MetalExec {
732    /// `original` is the node (or node chain) being replaced; its equivalence properties (ordering,
733    /// constants) are kept, its partitioning becomes a single partition.
734    pub(crate) fn new(
735        op: MetalOp,
736        inputs: Vec<Arc<dyn ExecutionPlan>>,
737        original: Arc<dyn ExecutionPlan>,
738        (log, plan): (SharedLog, u64),
739        agg: AggSettings,
740    ) -> Self {
741        let AggSettings { choice, table_rows, rows_hint } = agg;
742        let mut eq = original.equivalence_properties().clone();
743        eq.clear_per_partition_constants();
744        let props = PlanProperties::new(
745            eq,
746            Partitioning::UnknownPartitioning(1),
747            EmissionType::Final,
748            Boundedness::Bounded,
749        );
750        Self {
751            op,
752            inputs,
753            original,
754            props: Arc::new(props),
755            log,
756            plan,
757            metrics: ExecutionPlanMetricsSet::new(),
758            choice,
759            table_rows,
760            rows_hint,
761            join_slot: Arc::new(Mutex::new(None)),
762        }
763    }
764
765    /// A join's `MetalExec` with `n` output partitions: one execution on the GPU, its result
766    /// dealt out to the partitions batch by batch, so the operators above consume it in parallel
767    /// as they did the `HashJoinExec`'s partitions.
768    pub(crate) fn with_output_partitions(mut self, n: usize) -> Self {
769        if matches!(self.op, MetalOp::Join { .. }) && n > 1 {
770            let props = PlanProperties::new(
771                self.props.eq_properties.clone(),
772                Partitioning::UnknownPartitioning(n),
773                EmissionType::Final,
774                Boundedness::Bounded,
775            );
776            self.props = Arc::new(props);
777        }
778        self
779    }
780
781    /// Everything one execution needs.
782    fn job(&self, ctx: &Arc<TaskContext>) -> Job {
783        Job {
784            op: self.op.clone(),
785            inputs: self.inputs.clone(),
786            original: Arc::clone(&self.original),
787            out_schema: self.schema(),
788            log: Arc::clone(&self.log),
789            plan: self.plan,
790            choice: self.choice,
791            table_rows: self.table_rows,
792            rows_hint: self.rows_hint,
793            mm: self.phase_metrics(),
794            ctx: Arc::clone(ctx),
795        }
796    }
797
798    /// The operation this node runs.
799    pub fn op(&self) -> &MetalOp {
800        &self.op
801    }
802
803    /// The input it reads (a join's left input).
804    pub fn input(&self) -> &Arc<dyn ExecutionPlan> {
805        &self.inputs[0]
806    }
807
808    /// Every input it reads: one, or a join's left and right inputs.
809    pub fn inputs(&self) -> &[Arc<dyn ExecutionPlan>] {
810        &self.inputs
811    }
812
813    /// The per-phase metrics (EXPLAIN ANALYZE shows them): waiting for and collecting the input
814    /// (includes the upstream operators' own work), the group-count probe, and the GPU call split
815    /// into import (the chunked copy into Metal buffers) / plan run / export.
816    fn phase_metrics(&self) -> PhaseMetrics {
817        let m = |name: &'static str| MetricBuilder::new(&self.metrics).subset_time(name, 0);
818        PhaseMetrics {
819            input_time: m("input_time"),
820            import_time: m("import_time"),
821            kernel_time: m("kernel_time"),
822            export_time: m("export_time"),
823            probe_time: m("probe_time"),
824            input_batches: MetricBuilder::new(&self.metrics).counter("input_batches", 0),
825            handed_back: MetricBuilder::new(&self.metrics).counter("handed_back", 0),
826            groups_estimate: MetricBuilder::new(&self.metrics).gauge("groups_estimate", 0),
827        }
828    }
829}
830
831/// Whether a sample of at most `CONFIRM_SAMPLE` rows of the whole input (`refs`, `rows` rows)
832/// agrees with the estimate `e` a take was decided from: their ranges must meet.
833pub(crate) fn confirm(
834    refs: &[&RecordBatch],
835    keys: &[usize],
836    rows: usize,
837    e: &crate::probe::GroupEstimate,
838) -> (bool, crate::probe::GroupEstimate) {
839    let f = crate::probe::estimate_up_to(refs, keys, None, Some(rows), CONFIRM_SAMPLE);
840    (f.low <= e.high && f.high >= e.low, f)
841}
842
843/// The run-time decision under `AggregateChoice::Measured`: (on ArrowMetal, why, the estimate).
844fn decide(
845    op: &MetalOp,
846    input: &dyn ExecutionPlan,
847    refs: &[&RecordBatch],
848    keys: &[usize],
849    rows: usize,
850    table_rows: Option<usize>,
851) -> (bool, String, Option<crate::probe::GroupEstimate>) {
852    let Some(shape) = crate::choice::shape(op, input) else {
853        return (false, "not an aggregate".into(), None);
854    };
855    let at = table_rows.unwrap_or(rows) as u64;
856    let settled = crate::choice::settled_for(&shape, at);
857    let e = crate::probe::estimate_up_to(refs, keys, Some(&settled), Some(rows), DECIDE_SAMPLE);
858    let how = if e.exact {
859        format!("{} groups, counted over all {} rows", e.estimate, e.rows)
860    } else {
861        format!("an estimated {} groups ({} to {}, Chao1 over a {}-row sample)", e.estimate, e.low, e.high, e.sample_rows)
862    };
863    let at_text = match table_rows {
864        Some(n) => format!("{} rows (table_rows; input {rows})", crate::choice::rows_text(n as u64)),
865        None => format!("{rows} rows"),
866    };
867    if !settled(e.low, e.high) {
868        return (
869            false,
870            format!("{how} at {at_text}: the range reaches group counts the table decides differently; left to DataFusion"),
871            Some(e),
872        );
873    }
874    let b = crate::choice::bucket(e.estimate, at);
875    match crate::choice::takes(&shape, b, at) {
876        Ok(why) => (true, format!("{how} at {at_text}: {why}"), Some(e)),
877        Err(why) => (false, format!("{how} at {at_text}: {why}; left to DataFusion"), Some(e)),
878    }
879}
880
881impl DisplayAs for MetalExec {
882    fn fmt_as(&self, _t: DisplayFormatType, f: &mut fmt::Formatter) -> fmt::Result {
883        let schema = self.inputs[0].schema();
884        let name = |i: usize| schema.field(i).name().clone();
885        match &self.op {
886            MetalOp::Join { how, left_keys, right_keys, projection, .. } => {
887                let right = self.inputs.get(1).map(|r| r.schema()).unwrap_or_else(|| Arc::clone(&schema));
888                let on: Vec<String> = left_keys
889                    .iter()
890                    .zip(right_keys)
891                    .map(|(&l, &r)| format!("({}, {})", name(l), right.field(r).name()))
892                    .collect();
893                write!(f, "MetalExec: join={how:?}, on=[{}], projection={projection:?}", on.join(", "))
894            }
895            MetalOp::Sort { keys, fetch } => {
896                let ks: Vec<String> = keys
897                    .iter()
898                    .map(|k| {
899                        format!(
900                            "{} {} NULLS {}",
901                            name(k.column),
902                            if k.descending { "DESC" } else { "ASC" },
903                            if k.nulls_first { "FIRST" } else { "LAST" }
904                        )
905                    })
906                    .collect();
907                write!(f, "MetalExec: sort=[{}]", ks.join(", "))?;
908                if let Some(n) = fetch {
909                    write!(f, ", fetch={n}")?;
910                }
911                Ok(())
912            }
913            MetalOp::Aggregate { keys, aggs } => {
914                let ks: Vec<String> = keys.iter().map(|&k| name(k)).collect();
915                let asx: Vec<String> = aggs
916                    .iter()
917                    .map(|a| format!("{:?}({})", a.kind, a.column.map(name).unwrap_or_else(|| "*".into())))
918                    .collect();
919                write!(f, "MetalExec: group_by=[{}], aggr=[{}]", ks.join(", "), asx.join(", "))
920            }
921            MetalOp::Filter { predicate, projection, .. } => {
922                write!(f, "MetalExec: filter={predicate}")?;
923                if let Some(p) = projection {
924                    write!(f, ", projection={p:?}")?;
925                }
926                Ok(())
927            }
928        }
929    }
930}
931
932/// `plan` with every occurrence of the node `old` (by pointer) replaced by `new`.
933fn replace_leaf(
934    plan: &Arc<dyn ExecutionPlan>,
935    old: &Arc<dyn ExecutionPlan>,
936    new: &Arc<dyn ExecutionPlan>,
937) -> Result<Arc<dyn ExecutionPlan>> {
938    if Arc::ptr_eq(plan, old) {
939        return Ok(Arc::clone(new));
940    }
941    Ok(Arc::clone(plan)
942        .transform_down(|n| {
943            if Arc::ptr_eq(&n, old) {
944                Ok(Transformed::new(Arc::clone(new), true, TreeNodeRecursion::Jump))
945            } else {
946                Ok(Transformed::no(n))
947            }
948        })?
949        .data)
950}
951
952impl ExecutionPlan for MetalExec {
953    fn name(&self) -> &str {
954        "MetalExec"
955    }
956
957    fn properties(&self) -> &Arc<PlanProperties> {
958        &self.props
959    }
960
961    fn children(&self) -> Vec<&Arc<dyn ExecutionPlan>> {
962        self.inputs.iter().collect()
963    }
964
965    fn benefits_from_input_partitioning(&self) -> Vec<bool> {
966        vec![false; self.inputs.len()]
967    }
968
969    fn apply_expressions(
970        &self,
971        _f: &mut dyn FnMut(&Arc<dyn PhysicalExpr>) -> Result<TreeNodeRecursion>,
972    ) -> Result<TreeNodeRecursion> {
973        Ok(TreeNodeRecursion::Continue)
974    }
975
976    fn with_new_children(
977        self: Arc<Self>,
978        children: Vec<Arc<dyn ExecutionPlan>>,
979    ) -> Result<Arc<dyn ExecutionPlan>> {
980        if children.len() != self.inputs.len() {
981            return internal_err!("MetalExec takes {} children, got {}", self.inputs.len(), children.len());
982        }
983        let mut original = Arc::clone(&self.original);
984        for (old, new) in self.inputs.iter().zip(&children) {
985            original = replace_leaf(&original, old, new)?;
986        }
987        Ok(Arc::new(MetalExec {
988            op: self.op.clone(),
989            inputs: children,
990            original,
991            props: Arc::clone(&self.props),
992            log: Arc::clone(&self.log),
993            plan: self.plan,
994            metrics: ExecutionPlanMetricsSet::new(),
995            choice: self.choice,
996            table_rows: self.table_rows,
997            rows_hint: self.rows_hint,
998            join_slot: Arc::new(Mutex::new(None)),
999        }))
1000    }
1001
1002    fn execute(&self, partition: usize, ctx: Arc<TaskContext>) -> Result<SendableRecordBatchStream> {
1003        let n = self.props.partitioning.partition_count();
1004        if partition >= n {
1005            return internal_err!("MetalExec has {n} partition(s), asked for {partition}");
1006        }
1007        let schema = self.schema();
1008        if matches!(self.op, MetalOp::Join { .. }) {
1009            // One execution for all output partitions: the first partition executed starts it,
1010            // and a new one starts once every partition has taken the last one.
1011            let mut slot = self.join_slot.lock().unwrap_or_else(|p| p.into_inner());
1012            let fresh = slot.as_ref().is_none_or(|(_, taken)| *taken >= n);
1013            if fresh {
1014                let job = self.job(&ctx);
1015                let mut streams = Vec::new();
1016                let mut left = 0;
1017                for (i, input) in self.inputs.iter().enumerate() {
1018                    for p in 0..input.output_partitioning().partition_count() {
1019                        streams.push(input.execute(p, Arc::clone(&ctx))?);
1020                    }
1021                    if i == 0 {
1022                        left = streams.len();
1023                    }
1024                }
1025                let fut: futures::future::BoxFuture<'static, std::result::Result<Arc<Outcome>, Arc<DataFusionError>>> =
1026                    Box::pin(async move {
1027                        let out = match job.run_join(streams, left).await {
1028                            Ok(Outcome::Back(plan)) => fit_partitions(plan, n).map(Outcome::Back),
1029                            other => other,
1030                        };
1031                        out.map(Arc::new).map_err(Arc::new)
1032                    });
1033                *slot = Some((fut.shared(), 0));
1034            }
1035            let Some((shared, taken)) = slot.as_mut() else {
1036                return internal_err!("MetalExec join slot empty");
1037            };
1038            *taken += 1;
1039            let shared = shared.clone();
1040            if *taken >= n {
1041                // Every partition has its handle: the result lives as long as their streams.
1042                *slot = None;
1043            }
1044            let s = partition_of(shared, partition, n, ctx);
1045            return Ok(Box::pin(RecordBatchStreamAdapter::new(schema, s)));
1046        }
1047        let job = self.job(&ctx);
1048        let fut: futures::future::BoxFuture<'static, Result<BoxStream<'static, Result<RecordBatch>>>> =
1049            if matches!(self.op, MetalOp::Aggregate { .. }) {
1050                // Every partition of the input, each its own stream.
1051                let input = &self.inputs[0];
1052                let mut streams = Vec::new();
1053                for p in 0..input.output_partitioning().partition_count() {
1054                    streams.push(input.execute(p, Arc::clone(&ctx))?);
1055                }
1056                Box::pin(job.run_aggregate(streams))
1057            } else {
1058                // Coalescing runs the input partitions concurrently, as DataFusion's own merge does.
1059                let input = &self.inputs[0];
1060                let source: Arc<dyn ExecutionPlan> = if input.output_partitioning().partition_count() > 1 {
1061                    Arc::new(CoalescePartitionsExec::new(Arc::clone(input)))
1062                } else {
1063                    Arc::clone(input)
1064                };
1065                let stream = source.execute(0, Arc::clone(&ctx))?;
1066                Box::pin(job.run_sort_or_filter(stream))
1067            };
1068        let s = futures::stream::once(fut).try_flatten();
1069        Ok(Box::pin(RecordBatchStreamAdapter::new(schema, s)))
1070    }
1071
1072    fn metrics(&self) -> Option<MetricsSet> {
1073        Some(self.metrics.clone_inner())
1074    }
1075}
1076
1077#[cfg(test)]
1078mod tests {
1079    use super::*;
1080    use arrow::array::Int64Array;
1081    use arrow::datatypes::{Field, Schema};
1082
1083    fn batches(keys: Vec<i64>) -> Vec<RecordBatch> {
1084        let schema = Arc::new(Schema::new(vec![Field::new("k", DataType::Int64, false)]));
1085        keys.chunks(8192)
1086            .map(|c| RecordBatch::try_new(schema.clone(), vec![Arc::new(Int64Array::from(c.to_vec()))]).unwrap())
1087            .collect()
1088    }
1089
1090    /// A prefix that repeats a few thousand keys, followed by keys seen once: the prefix's estimate
1091    /// is far below the input's groups, and the whole-input sample disagrees with it.
1092    #[test]
1093    fn a_prefix_that_under_counts_is_not_confirmed() {
1094        let n = 1_200_000usize;
1095        let head = 262_144usize;
1096        let keys: Vec<i64> = (0..n).map(|i| if i < head { (i % 20_000) as i64 } else { i as i64 }).collect();
1097        let all = batches(keys);
1098        let refs: Vec<&RecordBatch> = all.iter().collect();
1099        let prefix: Vec<&RecordBatch> = refs[..head / 8192].to_vec();
1100        let e = crate::probe::estimate_up_to(&prefix, &[0], None, Some(n), 8_192);
1101        assert!(e.high < 40_000, "{e:?}");
1102        let (agree, f) = confirm(&refs, &[0], n, &e);
1103        assert!(!agree, "prefix {e:?} whole {f:?}");
1104        // The same data in random order: the prefix's range meets the whole input's.
1105        let mut shuffled: Vec<i64> = (0..n).map(|i| if i < head { (i % 20_000) as i64 } else { i as i64 }).collect();
1106        let mut state = 0x1234_5678u64;
1107        for i in (1..n).rev() {
1108            state ^= state << 13;
1109            state ^= state >> 7;
1110            state ^= state << 17;
1111            shuffled.swap(i, (state % (i as u64 + 1)) as usize);
1112        }
1113        let all = batches(shuffled);
1114        let refs: Vec<&RecordBatch> = all.iter().collect();
1115        let prefix: Vec<&RecordBatch> = refs[..head / 8192].to_vec();
1116        let e = crate::probe::estimate_up_to(&prefix, &[0], None, Some(n), 8_192);
1117        let (agree, f) = confirm(&refs, &[0], n, &e);
1118        assert!(agree, "prefix {e:?} whole {f:?}");
1119    }
1120}